← back to Exo
Fix occasional warmup bugs by using mlx_generate (#1794)
565ed41c132db6861bfc3ca0b0e1ef63d4e92755 · 2026-03-25 15:53:54 +0000 · rltakashige
## Motivation
Warmup occasionally had issues; e.g. #1748 and #1793 because we were
using a standard stream_generate, all of which are issues that are
resolved in the wrapper function mlx_generate.
Files touched
M src/exo/worker/engines/mlx/generator/generate.py
Diff
commit 565ed41c132db6861bfc3ca0b0e1ef63d4e92755
Author: rltakashige <rl.takashige@gmail.com>
Date: Wed Mar 25 15:53:54 2026 +0000
Fix occasional warmup bugs by using mlx_generate (#1794)
## Motivation
Warmup occasionally had issues; e.g. #1748 and #1793 because we were
using a standard stream_generate, all of which are issues that are
resolved in the wrapper function mlx_generate.
---
src/exo/worker/engines/mlx/generator/generate.py | 61 ++++++++++--------------
1 file changed, 26 insertions(+), 35 deletions(-)
diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py
index 82b06ef5..19a9d455 100644
--- a/src/exo/worker/engines/mlx/generator/generate.py
+++ b/src/exo/worker/engines/mlx/generator/generate.py
@@ -313,55 +313,46 @@ def warmup_inference(
model_id: ModelId,
) -> int:
logger.info(f"warming up inference for instance: {model_id}")
- t = time.monotonic()
content = "Prompt to warm up the inference engine. Repeat this."
+ warmup_task_params = TextGenerationTaskParams(
+ model=model_id,
+ input=[InputMessage(role="user", content=content)],
+ max_output_tokens=50,
+ temperature=0.0,
+ )
+
warmup_prompt = apply_chat_template(
tokenizer=tokenizer,
- task_params=TextGenerationTaskParams(
- model=ModelId(""),
- input=[InputMessage(role="user", content=content)],
- ),
+ task_params=warmup_task_params,
)
tokens_generated = 0
- cache = make_kv_cache(
- model=model,
- )
-
- # Use a default sampler for warmup
- sampler = make_sampler(temp=0.0)
-
mx_barrier(group)
logger.info("Generating warmup tokens")
- try:
- # for slow warmups, pipeline prefill=True tends to be more likely to succeed within the 5s gpu timeout window
- # as we don't block on the last all gather.
- set_pipeline_prefill(model, is_prefill=True)
- for _r in stream_generate(
- model=model,
- tokenizer=tokenizer,
- prompt=warmup_prompt,
- max_tokens=50,
- sampler=sampler,
- prompt_cache=cache,
- prefill_step_size=2048,
- kv_group_size=KV_GROUP_SIZE,
- kv_bits=KV_BITS,
- ):
- tokens_generated += 1
- finally:
- set_pipeline_prefill(model, is_prefill=False)
- mx_barrier(group)
+ t = time.monotonic()
+
+ for _r in mlx_generate(
+ model=model,
+ tokenizer=tokenizer,
+ task=warmup_task_params,
+ prompt=warmup_prompt,
+ kv_prefix_cache=None,
+ group=group,
+ ):
+ tokens_generated += 1
- logger.info(f"warmed up by generating {tokens_generated} tokens")
check_for_cancel_every = min(
math.ceil(tokens_generated / min(time.monotonic() - t, 0.001)), 100
)
+
+ mx_barrier(group)
+
+ logger.info(f"warmed up by generating {tokens_generated} tokens")
if group is not None:
check_for_cancel_every = int(
mx.max(
@@ -654,9 +645,9 @@ def mlx_generate(
if len(all_prompt_tokens) > 0
else 0.0
)
- if (
- matched_index is not None
- and hit_ratio >= _MIN_PREFIX_HIT_RATIO_TO_UPDATE
+ if matched_index is not None and (
+ prefix_hit_length > 1000
+ or hit_ratio >= _MIN_PREFIX_HIT_RATIO_TO_UPDATE
):
kv_prefix_cache.update_kv_cache(
matched_index,
← 2da740c3 Feat/static peer discovery (#1690)
·
back to Exo
·
fix: DeepSeek V3.2 warmup crash and tool calling + add catal fc1ae901 →