← back to Exo
fix: DeepSeek V3.2 warmup crash and tool calling + add catalog cards (#1769)
fc1ae901116c82309f021cac504c22463325d082 · 2026-03-25 19:20:35 +0300 · vskiwi
## Summary
DeepSeek V3.2 (`DeepseekV32ForCausalLM`) is already supported by exo's
inference engine (architecture whitelisted in `model_cards.py`, DSML
encoding added in #1548), but **doesn't work out of the box** due to two
bugs:
### Bug 1: `warmup_inference` passes empty model ID
`warmup_inference()` in `generate.py` accepts `model_id: ModelId` as a
parameter but creates `TextGenerationTaskParams(model=ModelId(""), ...)`
instead of using it. Since `_needs_dsml_encoding()` checks
`"deepseek-v3.2" in task_params.model.lower()`, the empty string never
matches → falls back to `tokenizer.apply_chat_template()` →
**ValueError** because V3.2 has no Jinja chat template.
**Fix:** `model=ModelId("")` → `model=model_id` (one line).
### Bug 2: `_needs_dsml_encoding` limited to tool calling
`_needs_dsml_encoding()` returns `True` only when `task_params.tools` is
present or tool messages exist in `chat_template_messages`. For warmup
and regular chat requests without tools → `return False` → Jinja
fallback → **ValueError**.
Unlike V3.1 (which has a `.jinja` chat template file that transformers
picks up automatically), V3.2 **has no Jinja template at all** — it uses
Python-based DSML encoding for all message types.
**Fix:** For V3.2, always return `True` — DSML encoding handles all
message types.
### Catalog cards
Added inference model cards for:
- `mlx-community/DeepSeek-V3.2-8bit`
- `mlx-community/DeepSeek-V3.2-4bit`
Parameters taken from model `config.json` on HuggingFace, storage sizes
from HF API. Capabilities include `thinking_toggle` (related: #1456).
## Notes
- The model ID string matching approach (`"deepseek-v3.2" in
model.lower()`) is acknowledged tech debt — see #1371 for the planned
architecture-based approach.
## Test plan
- [x] Start exo with DeepSeek V3.2 model → warmup should complete
without crash
- [x] Send a regular chat message (no tools) → should get a response
- [x] Send a chat message with tools → should work as before
- [x] V3.2 cards should appear in the dashboard model catalog
---------
Co-authored-by: user <user@m1.note>
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
Co-authored-by: Evan <evanev7@gmail.com>
Files touched
M justfileM pyproject.tomlA resources/inference_model_cards/mlx-community--DeepSeek-V3.2-4bit.tomlA resources/inference_model_cards/mlx-community--DeepSeek-V3.2-8bit.tomlM rust/exo_pyo3_bindings/exo_pyo3_bindings.pyiM src/exo/api/adapters/chat_completions.pyM src/exo/utils/info_gatherer/info_gatherer.pyM src/exo/worker/engines/mlx/auto_parallel.pyM src/exo/worker/engines/mlx/dsml_encoding.pyM src/exo/worker/engines/mlx/generator/batch_generate.pyM src/exo/worker/engines/mlx/utils_mlx.pyM src/exo/worker/runner/llm_inference/batch_generator.pyM src/exo/worker/runner/llm_inference/model_output_parsers.pyM src/exo/worker/runner/runner_supervisor.pyD src/exo/worker/tests/unittests/test_mlx/test_batch_vs_generate.pyM src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.pyM src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.pyA src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.pyM tests/run_exo_on.shM uv.lock
Diff
commit fc1ae901116c82309f021cac504c22463325d082
Author: vskiwi <141816715+vskiwi@users.noreply.github.com>
Date: Wed Mar 25 19:20:35 2026 +0300
fix: DeepSeek V3.2 warmup crash and tool calling + add catalog cards (#1769)
## Summary
DeepSeek V3.2 (`DeepseekV32ForCausalLM`) is already supported by exo's
inference engine (architecture whitelisted in `model_cards.py`, DSML
encoding added in #1548), but **doesn't work out of the box** due to two
bugs:
### Bug 1: `warmup_inference` passes empty model ID
`warmup_inference()` in `generate.py` accepts `model_id: ModelId` as a
parameter but creates `TextGenerationTaskParams(model=ModelId(""), ...)`
instead of using it. Since `_needs_dsml_encoding()` checks
`"deepseek-v3.2" in task_params.model.lower()`, the empty string never
matches → falls back to `tokenizer.apply_chat_template()` →
**ValueError** because V3.2 has no Jinja chat template.
**Fix:** `model=ModelId("")` → `model=model_id` (one line).
### Bug 2: `_needs_dsml_encoding` limited to tool calling
`_needs_dsml_encoding()` returns `True` only when `task_params.tools` is
present or tool messages exist in `chat_template_messages`. For warmup
and regular chat requests without tools → `return False` → Jinja
fallback → **ValueError**.
Unlike V3.1 (which has a `.jinja` chat template file that transformers
picks up automatically), V3.2 **has no Jinja template at all** — it uses
Python-based DSML encoding for all message types.
**Fix:** For V3.2, always return `True` — DSML encoding handles all
message types.
### Catalog cards
Added inference model cards for:
- `mlx-community/DeepSeek-V3.2-8bit`
- `mlx-community/DeepSeek-V3.2-4bit`
Parameters taken from model `config.json` on HuggingFace, storage sizes
from HF API. Capabilities include `thinking_toggle` (related: #1456).
## Notes
- The model ID string matching approach (`"deepseek-v3.2" in
model.lower()`) is acknowledged tech debt — see #1371 for the planned
architecture-based approach.
## Test plan
- [x] Start exo with DeepSeek V3.2 model → warmup should complete
without crash
- [x] Send a regular chat message (no tools) → should get a response
- [x] Send a chat message with tools → should work as before
- [x] V3.2 cards should appear in the dashboard model catalog
---------
Co-authored-by: user <user@m1.note>
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
Co-authored-by: Evan <evanev7@gmail.com>
---
justfile | 3 +
pyproject.toml | 2 +-
.../mlx-community--DeepSeek-V3.2-4bit.toml | 13 +
.../mlx-community--DeepSeek-V3.2-8bit.toml | 13 +
rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi | 2 +-
src/exo/api/adapters/chat_completions.py | 1 +
src/exo/utils/info_gatherer/info_gatherer.py | 11 +-
src/exo/worker/engines/mlx/auto_parallel.py | 2 +-
src/exo/worker/engines/mlx/dsml_encoding.py | 23 +-
.../worker/engines/mlx/generator/batch_generate.py | 5 +-
src/exo/worker/engines/mlx/utils_mlx.py | 21 +-
.../worker/runner/llm_inference/batch_generator.py | 23 +-
.../runner/llm_inference/model_output_parsers.py | 93 +++--
src/exo/worker/runner/runner_supervisor.py | 72 ++--
.../unittests/test_mlx/test_batch_vs_generate.py | 389 ---------------------
.../test_mlx/test_prefix_cache_architectures.py | 2 +-
.../tests/unittests/test_runner/test_dsml_e2e.py | 68 ++++
.../test_runner/test_finish_reason_sse.py | 332 ++++++++++++++++++
tests/run_exo_on.sh | 15 +-
uv.lock | 4 +-
20 files changed, 582 insertions(+), 512 deletions(-)
diff --git a/justfile b/justfile
index c3bab3c8..f7686b50 100644
--- a/justfile
+++ b/justfile
@@ -1,5 +1,8 @@
export NIX_CONFIG := "extra-experimental-features = nix-command flakes"
+default: lint fmt
+all: lint fmt check
+
fmt:
treefmt || nix fmt
diff --git a/pyproject.toml b/pyproject.toml
index b8feb209..9b7cce4a 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -61,7 +61,7 @@ members = ["rust/exo_pyo3_bindings", "bench"]
[tool.uv.sources]
exo_pyo3_bindings = { workspace = true }
mlx = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git", branch = "address-rdma-gpu-locks", marker = "sys_platform == 'darwin'" }
-mlx-lm = { git = "https://github.com/rltakashige/mlx-lm", branch = "fix/float32-logprobs" }
+mlx-lm = { git = "https://github.com/rltakashige/mlx-lm", branch = "leo/fix-deepseek-v32-indexer" }
# Uncomment to use local mlx/mlx-lm development versions:
# mlx = { path = "/Users/Shared/mlx", editable=true }
# mlx-lm = { path = "/Users/Shared/mlx-lm", editable=true }
diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.2-4bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.2-4bit.toml
new file mode 100644
index 00000000..0ca9bb22
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.2-4bit.toml
@@ -0,0 +1,13 @@
+model_id = "mlx-community/DeepSeek-V3.2-4bit"
+n_layers = 61
+hidden_size = 7168
+num_key_value_heads = 128
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "deepseek"
+quantization = "4bit"
+base_model = "DeepSeek V3.2"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 378086226621
diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.2-8bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.2-8bit.toml
new file mode 100644
index 00000000..a6d9d628
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.2-8bit.toml
@@ -0,0 +1,13 @@
+model_id = "mlx-community/DeepSeek-V3.2-8bit"
+n_layers = 61
+hidden_size = 7168
+num_key_value_heads = 128
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "deepseek"
+quantization = "8bit"
+base_model = "DeepSeek V3.2"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 755957120916
diff --git a/rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi b/rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi
index dc189ba4..bfd8978a 100644
--- a/rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi
+++ b/rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi
@@ -42,7 +42,7 @@ class MessageTooLargeError(builtins.Exception):
@typing.final
class NetworkingHandle:
- def __new__(cls, identity: Keypair, bootstrap_peers: list[builtins.str], listen_port: builtins.int) -> NetworkingHandle: ...
+ def __new__(cls, identity: Keypair, bootstrap_peers: typing.Sequence[builtins.str], listen_port: builtins.int) -> NetworkingHandle: ...
async def gossipsub_subscribe(self, topic: builtins.str) -> builtins.bool:
r"""
Subscribe to a `GossipSub` topic.
diff --git a/src/exo/api/adapters/chat_completions.py b/src/exo/api/adapters/chat_completions.py
index f0aad021..38939151 100644
--- a/src/exo/api/adapters/chat_completions.py
+++ b/src/exo/api/adapters/chat_completions.py
@@ -221,6 +221,7 @@ async def generate_chat_stream(
if chunk.stats is not None:
yield f": generation_stats {chunk.stats.model_dump_json()}\n\n"
yield "data: [DONE]\n\n"
+ return
async def collect_chat_response(
diff --git a/src/exo/utils/info_gatherer/info_gatherer.py b/src/exo/utils/info_gatherer/info_gatherer.py
index 52276e64..de9f4e97 100644
--- a/src/exo/utils/info_gatherer/info_gatherer.py
+++ b/src/exo/utils/info_gatherer/info_gatherer.py
@@ -10,7 +10,6 @@ from typing import Self, cast
import anyio
from anyio import fail_after, open_process, to_thread
from anyio.streams.buffered import BufferedByteReceiveStream
-from anyio.streams.text import TextReceiveStream
from loguru import logger
from pydantic import ValidationError
@@ -590,11 +589,15 @@ class InfoGatherer:
if not p.stdout:
logger.critical("MacMon closed stdout")
return
- stream = TextReceiveStream(BufferedByteReceiveStream(p.stdout))
+ stream = BufferedByteReceiveStream(p.stdout)
while True:
with fail_after(read_timeout):
- text = await stream.receive()
- await self.info_sender.send(MacmonMetrics.from_raw_json(text))
+ data = await stream.receive_until(
+ delimiter=b"\n", max_bytes=8 * 1024
+ )
+ text = data.decode("utf-8", errors="replace").strip()
+ metrics = MacmonMetrics.from_raw_json(text)
+ await self.info_sender.send(metrics)
except TimeoutError:
logger.warning(
f"MacMon produced no output for {read_timeout}s, restarting"
diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py
index c1169a5e..0085e2f5 100644
--- a/src/exo/worker/engines/mlx/auto_parallel.py
+++ b/src/exo/worker/engines/mlx/auto_parallel.py
@@ -57,8 +57,8 @@ from mlx_lm.models.step3p5 import Model as Step35Model
from mlx_lm.models.step3p5 import Step3p5MLP as Step35MLP
from mlx_lm.models.step3p5 import Step3p5Model as Step35InnerModel
-from exo.shared.logging import logger
from exo.shared.types.worker.shards import PipelineShardMetadata
+from exo.worker.runner.bootstrap import logger
if TYPE_CHECKING:
from mlx_lm.models.cache import Cache
diff --git a/src/exo/worker/engines/mlx/dsml_encoding.py b/src/exo/worker/engines/mlx/dsml_encoding.py
index 8005ad49..9d1dfdd2 100644
--- a/src/exo/worker/engines/mlx/dsml_encoding.py
+++ b/src/exo/worker/engines/mlx/dsml_encoding.py
@@ -15,7 +15,28 @@ USER_TOKEN = "<\uff5cUser\uff5c>"
ASSISTANT_TOKEN = "<\uff5cAssistant\uff5c>"
TOOL_CALLS_START = f"<{DSML_TOKEN}function_calls>"
TOOL_CALLS_END = f"</{DSML_TOKEN}function_calls>"
-encode_messages = deepseek_v32.encode_messages
+_ORPHAN_THINK_END = ASSISTANT_TOKEN + THINKING_END
+_FIXED_THINK_BLOCK = ASSISTANT_TOKEN + THINKING_START + "\n" + THINKING_END
+
+
+def encode_messages(
+ messages: list[dict[str, Any]],
+ thinking_mode: str = "thinking",
+ context: list[dict[str, Any]] | None = None,
+ drop_thinking: bool = True,
+ add_default_bos_token: bool = True,
+ tools: Any = None, # pyright: ignore[reportAny]
+) -> str:
+ prompt: str = deepseek_v32.encode_messages(
+ messages,
+ thinking_mode=thinking_mode,
+ context=context,
+ drop_thinking=drop_thinking,
+ add_default_bos_token=add_default_bos_token,
+ tools=tools,
+ )
+ return prompt.replace(_ORPHAN_THINK_END, _FIXED_THINK_BLOCK)
+
_INVOKE_PATTERN = re.compile(
rf"<{re.escape(DSML_TOKEN)}invoke\s+name=\"([^\"]+)\">"
diff --git a/src/exo/worker/engines/mlx/generator/batch_generate.py b/src/exo/worker/engines/mlx/generator/batch_generate.py
index d9410c18..3b026051 100644
--- a/src/exo/worker/engines/mlx/generator/batch_generate.py
+++ b/src/exo/worker/engines/mlx/generator/batch_generate.py
@@ -393,9 +393,8 @@ class ExoBatchGenerator:
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
):
self.kv_prefix_cache.update_kv_cache(
matched_index,
diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py
index e34b4a43..c8fc862c 100644
--- a/src/exo/worker/engines/mlx/utils_mlx.py
+++ b/src/exo/worker/engines/mlx/utils_mlx.py
@@ -486,16 +486,7 @@ def _patch_lossy_chat_template(template: str) -> str | None:
def _needs_dsml_encoding(task_params: TextGenerationTaskParams) -> bool:
- if "deepseek-v3.2" not in task_params.model.lower():
- return False
- # Use DSML encoding when tools are provided or tool results are in the conversation
- if task_params.tools:
- return True
- if task_params.chat_template_messages:
- return any(
- msg.get("role") == "tool" for msg in task_params.chat_template_messages
- )
- return False
+ return "deepseek-v3.2" in task_params.model.lower()
def apply_chat_template(
@@ -514,8 +505,6 @@ def apply_chat_template(
if task_params.chat_template_messages is not None:
# Use pre-formatted messages that preserve tool_calls, thinking, etc.
formatted_messages = list(task_params.chat_template_messages)
- for msg in formatted_messages:
- _normalize_tool_calls(msg)
else:
# Add system message (instructions) if present
if task_params.instructions:
@@ -541,7 +530,10 @@ def apply_chat_template(
prompt = encode_messages(
messages=formatted_messages,
- thinking_mode="thinking" if task_params.enable_thinking else "chat",
+ # Only use chat mode if enable thinking is explicitly Fakse.
+ thinking_mode="chat"
+ if task_params.enable_thinking is False
+ else "thinking",
tools=task_params.tools,
)
if partial_assistant_content:
@@ -549,6 +541,9 @@ def apply_chat_template(
logger.info(prompt)
return prompt
+ for msg in formatted_messages:
+ _normalize_tool_calls(msg)
+
extra_kwargs: dict[str, Any] = {}
if task_params.enable_thinking is not None:
# Qwen3 and GLM use "enable_thinking"; DeepSeek uses "thinking".
diff --git a/src/exo/worker/runner/llm_inference/batch_generator.py b/src/exo/worker/runner/llm_inference/batch_generator.py
index 33bbb3a6..5998ef08 100644
--- a/src/exo/worker/runner/llm_inference/batch_generator.py
+++ b/src/exo/worker/runner/llm_inference/batch_generator.py
@@ -195,21 +195,29 @@ class SequentialGenerator(InferenceGenerator):
assert self._active is not None
task, mlx_gen, queue, output_generator = self._active
- response = None
+ output: list[
+ tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
+ ] = []
try:
- queue.push(next(mlx_gen))
- response = next(output_generator)
+ response = next(mlx_gen)
+ queue.push(response)
+ # drain potentially many responses every time
+ while (parsed := next(output_generator, None)) is not None:
+ output.append((task.task_id, parsed))
+
except (StopIteration, PrefillCancelled):
- response = Finished()
+ output.append((task.task_id, Finished()))
self._active = None
if self._queue:
self._start_next()
+
except Exception as e:
self._send_error(task, e)
self._active = None
raise
+
return itertools.chain(
- [] if response is None else [(task.task_id, response)],
+ output,
map(lambda task: (task, Cancelled()), self._cancelled_tasks),
)
@@ -428,11 +436,10 @@ class BatchGenerator(InferenceGenerator):
task, queue, output_generator = self._active_tasks[uid]
queue.push(response)
# If a generator fails to parse for some reason and returns early, we should not crash
- parsed = next(output_generator, None)
-
- if parsed is not None:
+ while (parsed := next(output_generator, None)) is not None:
output.append((task.task_id, parsed))
+ # check if original response was terminal and append a Finished()
if response.finish_reason is not None:
output.append((task.task_id, Finished()))
del self._active_tasks[uid]
diff --git a/src/exo/worker/runner/llm_inference/model_output_parsers.py b/src/exo/worker/runner/llm_inference/model_output_parsers.py
index b5729697..1242909a 100644
--- a/src/exo/worker/runner/llm_inference/model_output_parsers.py
+++ b/src/exo/worker/runner/llm_inference/model_output_parsers.py
@@ -159,11 +159,42 @@ def parse_deepseek_v32(
# Text accumulated during a tool call block
tool_call_text = ""
+ def _try_parse_tool_call(
+ text: str, response: GenerationResponse
+ ) -> ToolCallResponse | GenerationResponse:
+ parsed = parse_dsml_output(text)
+ if parsed is not None:
+ return ToolCallResponse(
+ tool_calls=parsed, usage=response.usage, stats=response.stats
+ )
+ logger.warning(f"DSML tool call parsing failed for: {text}")
+ return response.model_copy(update={"text": text})
+
for response in responses:
if response is None:
yield None
continue
+ if response.finish_reason is not None:
+ yield from pending_buffer
+ pending_buffer.clear()
+ if in_tool_call:
+ tool_call_text += response.text
+ yield (
+ _try_parse_tool_call(tool_call_text, response)
+ if TOOL_CALLS_END in tool_call_text
+ else response.model_copy(update={"text": tool_call_text})
+ )
+ elif TOOL_CALLS_START in response.text and TOOL_CALLS_END in response.text:
+ dsml_start = response.text.index(TOOL_CALLS_START)
+ before = response.text[:dsml_start]
+ if before:
+ yield response.model_copy(update={"text": before})
+ yield _try_parse_tool_call(response.text[dsml_start:], response)
+ else:
+ yield response
+ break
+
# ── Handle thinking tags ──
if not thinking and THINKING_START in response.text:
thinking = True
@@ -191,28 +222,7 @@ def parse_deepseek_v32(
if in_tool_call:
tool_call_text += response.text
if TOOL_CALLS_END in tool_call_text:
- # Parse the accumulated DSML block
- parsed = parse_dsml_output(tool_call_text)
- if parsed is not None:
- logger.info(f"parsed DSML tool calls: {parsed}")
- yield ToolCallResponse(
- tool_calls=parsed,
- usage=response.usage,
- stats=response.stats,
- )
- else:
- logger.warning(
- f"DSML tool call parsing failed for: {tool_call_text}"
- )
- yield response.model_copy(update={"text": tool_call_text})
- in_tool_call = False
- tool_call_text = ""
- continue
-
- # EOS reached before end marker — yield buffered text as-is
- if response.finish_reason is not None:
- logger.info("DSML tool call parsing interrupted by EOS")
- yield response.model_copy(update={"text": tool_call_text})
+ yield _try_parse_tool_call(tool_call_text, response)
in_tool_call = False
tool_call_text = ""
continue
@@ -228,33 +238,22 @@ def parse_deepseek_v32(
if pre_text:
# Flush pending buffer tokens that contributed text before the marker
for buf_resp in pending_buffer:
- if pre_text:
- chunk = buf_resp.text
- if len(chunk) <= len(pre_text):
- yield buf_resp
- pre_text = pre_text[len(chunk) :]
- else:
- yield buf_resp.model_copy(update={"text": pre_text})
- pre_text = ""
+ if not pre_text:
+ break
+ chunk = buf_resp.text
+ if len(chunk) <= len(pre_text):
+ yield buf_resp
+ pre_text = pre_text[len(chunk) :]
+ else:
+ yield buf_resp.model_copy(update={"text": pre_text})
+ pre_text = ""
pending_buffer = []
tool_call_text = accumulated[start_idx:]
accumulated = ""
# Check if the end marker is already present (entire tool call in one token)
if TOOL_CALLS_END in tool_call_text:
- parsed = parse_dsml_output(tool_call_text)
- if parsed is not None:
- logger.info(f"parsed DSML tool calls: {parsed}")
- yield ToolCallResponse(
- tool_calls=parsed,
- usage=response.usage,
- stats=response.stats,
- )
- else:
- logger.warning(
- f"DSML tool call parsing failed for: {tool_call_text}"
- )
- yield response.model_copy(update={"text": tool_call_text})
+ yield _try_parse_tool_call(tool_call_text, response)
tool_call_text = ""
else:
in_tool_call = True
@@ -267,15 +266,13 @@ def parse_deepseek_v32(
continue
# No partial match — flush all pending tokens and the current one
- for buf_resp in pending_buffer:
- yield buf_resp
- pending_buffer = []
+ yield from pending_buffer
+ pending_buffer.clear()
accumulated = ""
yield response
# Flush any remaining pending buffer at generator end
- for buf_resp in pending_buffer:
- yield buf_resp
+ yield from pending_buffer
def _could_be_dsml_prefix(text: str) -> bool:
diff --git a/src/exo/worker/runner/runner_supervisor.py b/src/exo/worker/runner/runner_supervisor.py
index b0882f39..c4b4dc5a 100644
--- a/src/exo/worker/runner/runner_supervisor.py
+++ b/src/exo/worker/runner/runner_supervisor.py
@@ -110,39 +110,45 @@ class RunnerSupervisor:
async def run(self):
self.runner_process.start()
- async with self._tg as tg:
- tg.start_soon(self._watch_runner)
- tg.start_soon(self._forward_events)
+ try:
+ async with self._tg as tg:
+ tg.start_soon(self._watch_runner)
+ tg.start_soon(self._forward_events)
+ finally:
+ logger.info("Runner supervisor shutting down")
+ if not self._cancel_watch_runner.cancel_called:
+ self._cancel_watch_runner.cancel()
+ with contextlib.suppress(ClosedResourceError):
+ self._ev_recv.close()
+ with contextlib.suppress(ClosedResourceError):
+ self._task_sender.close()
+ with contextlib.suppress(ClosedResourceError):
+ self._event_sender.close()
+ with contextlib.suppress(ClosedResourceError):
+ self._cancel_sender.send(CANCEL_ALL_TASKS)
+ with contextlib.suppress(ClosedResourceError):
+ self._cancel_sender.close()
- def shutdown(self):
- logger.info("Runner supervisor shutting down")
- self._tg.cancel_tasks()
- if not self._cancel_watch_runner.cancel_called:
- self._cancel_watch_runner.cancel()
- with contextlib.suppress(ClosedResourceError):
- self._ev_recv.close()
- with contextlib.suppress(ClosedResourceError):
- self._task_sender.close()
- with contextlib.suppress(ClosedResourceError):
- self._event_sender.close()
- with contextlib.suppress(ClosedResourceError):
- self._cancel_sender.send(CANCEL_ALL_TASKS)
- with contextlib.suppress(ClosedResourceError):
- self._cancel_sender.close()
- self.runner_process.join(5)
- if not self.runner_process.is_alive():
- logger.info("Runner process succesfully terminated")
- return
+ await to_thread.run_sync(self.runner_process.join, 5)
- # This is overkill but it's not technically bad, just unnecessary.
- logger.warning("Runner process didn't shutdown succesfully, terminating")
- self.runner_process.terminate()
- self.runner_process.join(1)
- if not self.runner_process.is_alive():
- return
+ if self.runner_process.is_alive():
+ logger.warning(
+ "Runner process didn't shutdown succesfully, terminating"
+ )
+ self.runner_process.terminate()
+ self.runner_process.join(timeout=5)
+ # This is overkill but it's not technically bad, just unnecessary.
+ if self.runner_process.is_alive():
+ logger.critical("Runner process didn't respond to SIGTERM, killing")
+ self.runner_process.kill()
+ self.runner_process.join(timeout=5)
+ else:
+ logger.info("Runner process succesfully terminated")
+
+ self.runner_process.close()
- logger.critical("Runner process didn't respond to SIGTERM, killing")
- self.runner_process.kill()
+ def shutdown(self):
+ self._tg.cancel_tasks()
async def start_task(self, task: Task):
if task.task_id in self.pending:
@@ -218,12 +224,6 @@ class RunnerSupervisor:
for tid in self.pending:
self.pending[tid].set()
- def __del__(self) -> None:
- if self.runner_process.is_alive():
- logger.critical("RunnerSupervisor was not stopped cleanly.")
- with contextlib.suppress(ValueError):
- self.runner_process.kill()
-
async def _watch_runner(self) -> None:
with self._cancel_watch_runner:
while True:
diff --git a/src/exo/worker/tests/unittests/test_mlx/test_batch_vs_generate.py b/src/exo/worker/tests/unittests/test_mlx/test_batch_vs_generate.py
deleted file mode 100644
index 5d1a792b..00000000
--- a/src/exo/worker/tests/unittests/test_mlx/test_batch_vs_generate.py
+++ /dev/null
@@ -1,389 +0,0 @@
-import copy
-import gc
-import json
-import shutil
-import tempfile
-from pathlib import Path
-from typing import Any, cast
-
-import mlx.core as mx
-import pytest
-from mlx_lm.tokenizer_utils import TokenizerWrapper
-
-from exo.shared.types.common import ModelId
-from exo.shared.types.mlx import KVCacheType, Model
-from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams
-from exo.worker.engines.mlx.cache import CacheSnapshot, KVPrefixCache, cache_length
-from exo.worker.engines.mlx.generator.batch_generate import ExoBatchGenerator
-from exo.worker.engines.mlx.generator.generate import mlx_generate
-from exo.worker.engines.mlx.utils_mlx import (
- apply_chat_template,
- load_tokenizer_for_model_id,
-)
-
-from .test_prefix_cache_architectures import (
- ARCHITECTURES,
- ArchSpec,
- _arch_available, # pyright: ignore[reportPrivateUsage]
- _build_model, # pyright: ignore[reportPrivateUsage]
- _copy_tokenizer, # pyright: ignore[reportPrivateUsage]
- _find_snapshot, # pyright: ignore[reportPrivateUsage]
- _reduce_config, # pyright: ignore[reportPrivateUsage]
-)
-
-
-def _make_task(
- content: str = "Hello, what is 2+2?",
- max_tokens: int = 10,
- seed: int = 42,
-) -> TextGenerationTaskParams:
- return TextGenerationTaskParams(
- model=ModelId("test"),
- input=[InputMessage(role="user", content=content)],
- max_output_tokens=max_tokens,
- temperature=0.7,
- seed=seed,
- )
-
-
-# ── Helpers ──────────────────────────────────────────────────────────────── #
-
-
-def _collect_mlx_generate(
- model: Model,
- tokenizer: TokenizerWrapper,
- task: TextGenerationTaskParams,
- kv_prefix_cache: KVPrefixCache | None,
-) -> list[int]:
- """Run mlx_generate and collect output token IDs."""
- prompt = apply_chat_template(tokenizer=tokenizer, task_params=task)
- tokens: list[int] = []
- for resp in mlx_generate(
- model=model,
- tokenizer=tokenizer,
- task=task,
- prompt=prompt,
- kv_prefix_cache=kv_prefix_cache,
- group=None,
- ):
- tokens.append(resp.token)
- if resp.finish_reason is not None:
- break
- return tokens
-
-
-def _collect_batch_generate(
- model: Model,
- tokenizer: TokenizerWrapper,
- task_params: TextGenerationTaskParams,
- kv_prefix_cache: KVPrefixCache | None,
-) -> list[int]:
- """Run ExoBatchGenerator and collect raw output token IDs"""
- exo_gen = ExoBatchGenerator(
- model=model,
- tokenizer=tokenizer,
- group=None,
- kv_prefix_cache=kv_prefix_cache,
- )
-
- prompt = apply_chat_template(tokenizer=tokenizer, task_params=task_params)
- exo_gen.submit(task_params=task_params, prompt=prompt)
-
- tokens: list[int] = []
- while exo_gen.has_work:
- results = exo_gen.step()
- for _uid, response in results:
- tokens.append(response.token)
-
- exo_gen.close()
- return tokens
-
-
-def _assert_state_equal(sa: object, sb: object, label: str) -> None:
- """Compare two state items, handling both plain arrays and tuples of arrays (CacheList)."""
- if isinstance(sa, tuple):
- assert isinstance(sb, tuple), f"{label}: type mismatch"
- for k, (arr_a, arr_b) in enumerate(
- zip(
- cast(tuple[mx.array, ...], sa),
- cast(tuple[mx.array, ...], sb),
- strict=True,
- )
- ):
- a_f = mx.array(arr_a).astype(mx.float32)
- b_f = mx.array(arr_b).astype(mx.float32)
- if a_f.size == 0:
- assert b_f.size == 0, f"{label}[{k}]: size mismatch"
- continue
- diff = float(mx.max(mx.abs(a_f - b_f)).item())
- assert diff == 0.0, f"{label}[{k}]: max diff {diff}"
- else:
- sa_f = mx.array(cast(mx.array, sa)).astype(mx.float32)
- sb_f = mx.array(cast(mx.array, sb)).astype(mx.float32)
- if sa_f.size == 0:
- assert sb_f.size == 0, f"{label}: size mismatch"
- return
- diff = float(mx.max(mx.abs(sa_f - sb_f)).item())
- assert diff == 0.0, f"{label}: max diff {diff}"
-
-
-def _compare_cache_arrays(
- cache_a: KVCacheType,
- cache_b: KVCacheType,
- label: str = "",
-) -> None:
- """Assert two KV caches have identical array values."""
- assert len(cache_a) == len(cache_b), (
- f"{label}Cache layer count: {len(cache_a)} vs {len(cache_b)}"
- )
- for i, (a, b) in enumerate(zip(cache_a, cache_b, strict=True)):
- assert type(a) is type(b), (
- f"{label}Layer {i}: type {type(a).__name__} vs {type(b).__name__}"
- )
- states_a = a.state
- states_b = b.state
- assert len(states_a) == len(states_b), (
- f"{label}Layer {i}: state count {len(states_a)} vs {len(states_b)}"
- )
- for j, (sa, sb) in enumerate(zip(states_a, states_b, strict=True)):
- if sa is None and sb is None:
- continue
- assert sa is not None and sb is not None, (
- f"{label}Layer {i}, state {j}: one is None"
- )
- _assert_state_equal(sa, sb, f"{label}Layer {i}, state {j}")
-
-
-def _safe_state(cache: object) -> list[object]:
- """Safely access .state on a cache object. Returns [] if uninitialized."""
- # RotatingKVCache.state crashes when keys is None (uninitialized)
- if getattr(cache, "keys", _SENTINEL) is None:
- return []
- try:
- return list(cache.state) # type: ignore[union-attr]
- except (AttributeError, TypeError):
- return []
-
-
-_SENTINEL = object()
-
-
-def _compare_snapshots(
- snaps_a: list[CacheSnapshot] | None,
- snaps_b: list[CacheSnapshot] | None,
- label: str = "",
-) -> None:
- """Assert two snapshot lists are identical."""
- if snaps_a is None:
- assert snaps_b is None, f"{label}One side has snapshots, other doesn't"
- return
- assert snaps_b is not None, f"{label}One side has snapshots, other doesn't"
- assert len(snaps_a) == len(snaps_b), (
- f"{label}Snapshot count: {len(snaps_a)} vs {len(snaps_b)}"
- )
- for k, (sa, sb) in enumerate(zip(snaps_a, snaps_b, strict=True)):
- assert sa.token_count == sb.token_count, (
- f"{label}Snapshot {k} token_count: {sa.token_count} vs {sb.token_count}"
- )
- for layer_i, (s1, s2) in enumerate(zip(sa.states, sb.states, strict=True)):
- if s1 is None and s2 is None:
- continue
- assert s1 is not None and s2 is not None, (
- f"{label}Snapshot {k}, layer {layer_i}: one state is None"
- )
- state_a = _safe_state(s1)
- state_b = _safe_state(s2)
- if not state_a and not state_b:
- continue
- assert len(state_a) == len(state_b), (
- f"{label}Snapshot {k}, layer {layer_i}: state length mismatch"
- )
- for st_j, (arr_a, arr_b) in enumerate(zip(state_a, state_b, strict=True)):
- if arr_a is None and arr_b is None:
- continue
- assert arr_a is not None and arr_b is not None
- _assert_state_equal(
- arr_a,
- arr_b,
- f"{label}Snapshot {k}, layer {layer_i}, state {st_j}",
- )
-
-
-# ── Test class ────────────────────────────────────────────────────────────── #
-
-
-@pytest.mark.slow
-class TestBatchVsGenerate:
- """Verify BatchGenerator matches mlx_generate for output tokens and prefix cache."""
-
- @pytest.fixture(autouse=True)
- def _cleanup(self):
- yield
- mx.clear_cache()
- gc.collect()
-
- @pytest.mark.parametrize(
- "spec",
- ARCHITECTURES,
- ids=[a.name for a in ARCHITECTURES],
- )
- def test_same_output_and_cache(self, spec: ArchSpec) -> None:
- if not _arch_available(spec):
- pytest.skip(f"Model {spec.hub_name} not cached locally")
-
- snapshot = _find_snapshot(spec.hub_name)
- assert snapshot is not None
-
- tmpdir = Path(tempfile.mkdtemp(prefix=f"exo_batchtest_{spec.name}_"))
- try:
- # Build reduced config
- with open(snapshot / "config.json") as f:
- cfg = cast(dict[str, Any], json.load(f))
- reduced = _reduce_config(copy.deepcopy(cfg))
- (tmpdir / "config.json").write_text(json.dumps(reduced))
-
- # Copy tokenizer
- tok_src = snapshot
- if spec.tokenizer_hub is not None:
- alt = _find_snapshot(spec.tokenizer_hub)
- if alt is not None:
- tok_src = alt
- _copy_tokenizer(tok_src, tmpdir)
-
- # Load tokenizer, build model with random weights
- model_id = ModelId(f"mlx-community/{spec.hub_name}")
- tokenizer = load_tokenizer_for_model_id(model_id, tmpdir)
- mx.random.seed(0)
- model = _build_model(spec.module, reduced)
-
- task = _make_task()
-
- # ── Run mlx_generate path ──
- # Seed is set inside mlx_generate/ExoBatchGenerator.submit from task.seed
- kv_mlx = KVPrefixCache(None)
- mlx_tokens = _collect_mlx_generate(model, tokenizer, task, kv_mlx)
-
- # ── Run batch generator path ──
- kv_batch = KVPrefixCache(None)
- batch_tokens = _collect_batch_generate(model, tokenizer, task, kv_batch)
-
- # ── Compare output tokens ──
- assert len(mlx_tokens) > 0, "mlx_generate produced no tokens"
- assert len(batch_tokens) > 0, "BatchGenerator produced no tokens"
- assert mlx_tokens == batch_tokens, (
- f"[{spec.name}] Token mismatch:\n"
- f" mlx_generate: {mlx_tokens}\n"
- f" BatchGenerator: {batch_tokens}"
- )
-
- # ── Compare prefix cache KV arrays ──
- assert len(kv_mlx.caches) == 1, "mlx_generate didn't save to prefix cache"
- assert len(kv_batch.caches) == 1, (
- "BatchGenerator didn't save to prefix cache"
- )
-
- _compare_cache_arrays(
- kv_mlx.caches[0],
- kv_batch.caches[0],
- label=f"[{spec.name}] ",
- )
-
- # ── Compare cache lengths ──
- mlx_len = cache_length(kv_mlx.caches[0])
- batch_len = cache_length(kv_batch.caches[0])
- assert mlx_len == batch_len, (
- f"[{spec.name}] Cache length: mlx={mlx_len} vs batch={batch_len}"
- )
-
- # ── Compare snapshots ──
- _compare_snapshots(
- kv_mlx._snapshots[0], # pyright: ignore[reportPrivateUsage]
- kv_batch._snapshots[0], # pyright: ignore[reportPrivateUsage]
- label=f"[{spec.name}] ",
- )
-
- finally:
- shutil.rmtree(tmpdir, ignore_errors=True)
-
- @pytest.mark.parametrize(
- "spec",
- ARCHITECTURES,
- ids=[a.name for a in ARCHITECTURES],
- )
- def test_concurrent_batch_completes(self, spec: ArchSpec) -> None:
- """Two requests processed concurrently must both complete without
- crashing and produce non-empty output.
-
- Note: batch decode logits are NOT bit-exact with sequential because
- Metal's matmul kernel picks different reduction tiling for B=1 vs B=2
- when L=1 (decode step). This introduces sub-ULP float16 diffs in
- gate_proj/down_proj/lm_head which swiglu amplifies by |up_values|.
- With random weights these accumulate into argmax flips; with trained
- weights the diffs are absorbed and output matches exactly (verified
- with real Llama-3.2-1B-Instruct-4bit weights).
- """
- if not _arch_available(spec):
- pytest.skip(f"Model {spec.hub_name} not cached locally")
-
- snapshot = _find_snapshot(spec.hub_name)
- assert snapshot is not None
-
- tmpdir = Path(tempfile.mkdtemp(prefix=f"exo_concurrent_{spec.name}_"))
- try:
- with open(snapshot / "config.json") as f:
- cfg = cast(dict[str, Any], json.load(f))
- reduced = _reduce_config(copy.deepcopy(cfg))
- (tmpdir / "config.json").write_text(json.dumps(reduced))
-
- tok_src = snapshot
- if spec.tokenizer_hub is not None:
- alt = _find_snapshot(spec.tokenizer_hub)
- if alt is not None:
- tok_src = alt
- _copy_tokenizer(tok_src, tmpdir)
-
- model_id = ModelId(f"mlx-community/{spec.hub_name}")
- tokenizer = load_tokenizer_for_model_id(model_id, tmpdir)
- mx.random.seed(0)
- model = _build_model(spec.module, reduced)
-
- # Two different prompts → different prompt lengths.
- task_a = _make_task(content="Hello, what is 2+2?", seed=42)
- task_a = task_a.model_copy(update={"temperature": 0.0})
- task_b = _make_task(
- content="Write a short poem about the ocean and the sky.",
- seed=99,
- )
- task_b = task_b.model_copy(update={"temperature": 0.0})
-
- # ── Concurrent: submit both to one ExoBatchGenerator ──
- exo_gen = ExoBatchGenerator(
- model=model,
- tokenizer=tokenizer,
- group=None,
- kv_prefix_cache=None,
- )
-
- prompt_a = apply_chat_template(tokenizer=tokenizer, task_params=task_a)
- prompt_b = apply_chat_template(tokenizer=tokenizer, task_params=task_b)
- uid_a = exo_gen.submit(task_params=task_a, prompt=prompt_a)
- uid_b = exo_gen.submit(task_params=task_b, prompt=prompt_b)
-
- batch_tokens: dict[int, list[int]] = {uid_a: [], uid_b: []}
- finished: set[int] = set()
- while exo_gen.has_work:
- results = exo_gen.step()
- for uid, response in results:
- batch_tokens[uid].append(response.token)
- if response.finish_reason is not None:
- finished.add(uid)
-
- exo_gen.close()
-
- # ── Verify both completed ──
- assert len(batch_tokens[uid_a]) > 0, "No tokens for task A"
- assert len(batch_tokens[uid_b]) > 0, "No tokens for task B"
- assert uid_a in finished, "Task A never finished"
- assert uid_b in finished, "Task B never finished"
- finally:
- shutil.rmtree(tmpdir, ignore_errors=True)
diff --git a/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py b/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py
index 944e8290..609ea867 100644
--- a/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py
+++ b/src/exo/worker/tests/unittests/test_mlx/test_prefix_cache_architectures.py
@@ -190,7 +190,7 @@ ARCHITECTURES: list[ArchSpec] = [
def _arch_available(spec: ArchSpec) -> bool:
snap = _find_snapshot(spec.hub_name)
- if snap is None:
+ if snap is None or not (snap / "config.json").exists():
return False
if spec.tokenizer_hub is not None:
return _find_snapshot(spec.tokenizer_hub) is not None
diff --git a/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py b/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py
index a5502167..26c11fe5 100644
--- a/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py
+++ b/src/exo/worker/tests/unittests/test_runner/test_dsml_e2e.py
@@ -2,6 +2,7 @@ import json
from collections.abc import Generator
from typing import Any
+from exo.shared.types.common import ModelId
from exo.shared.types.worker.runner_response import (
GenerationResponse,
ToolCallResponse,
@@ -965,3 +966,70 @@ class TestE2EFullRoundTrip:
assert "sunny" in final_text.lower()
assert "5°C" in final_text
assert "12°C" in final_text
+
+
+class TestMultiTurnThinkingPrompt:
+ def test_no_orphan_think_end_in_multiturn(self):
+ messages: list[dict[str, Any]] = [
+ {"role": "user", "content": "Hi!"},
+ {"role": "assistant", "content": "Hello! How can I help you today?"},
+ {"role": "user", "content": "Tell me about Paris."},
+ ]
+ prompt = encode_messages(messages, thinking_mode="thinking")
+ assistant_token = "<\uff5cAssistant\uff5c>"
+ parts = prompt.split(assistant_token)
+ for part in parts[1:]:
+ assert not part.startswith(THINKING_END), (
+ f"Orphan </think> without <think> after <Assistant>: ...{assistant_token}{part[:50]}"
+ )
+
+
+class TestApplyChatTemplateWithToolCalls:
+ def test_dsml_encoding_with_tool_calls_in_history(self):
+ from exo.shared.types.text_generation import (
+ InputMessage,
+ TextGenerationTaskParams,
+ )
+ from exo.worker.engines.mlx.utils_mlx import apply_chat_template
+
+ chat_template_messages: list[dict[str, Any]] = [
+ {"role": "system", "content": "You are a helpful assistant."},
+ {"role": "user", "content": "What's the weather?"},
+ {
+ "role": "assistant",
+ "content": "",
+ "tool_calls": [
+ {
+ "id": "call_1",
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "arguments": '{"city": "Tokyo"}',
+ },
+ }
+ ],
+ },
+ {"role": "tool", "content": "Sunny, 25°C"},
+ {"role": "user", "content": "Thanks!"},
+ ]
+
+ from unittest.mock import MagicMock
+
+ tokenizer = MagicMock()
+ tokenizer.has_thinking = True
+ tokenizer.think_start = "<think>"
+ tokenizer.think_end = "</think>"
+
+ params = TextGenerationTaskParams(
+ model=ModelId("mlx-community/DeepSeek-V3.2-8bit"),
+ input=[InputMessage(role="user", content="Thanks!")],
+ instructions="You are a helpful assistant.",
+ enable_thinking=True,
+ chat_template_messages=chat_template_messages,
+ tools=_WEATHER_TOOLS,
+ )
+
+ prompt = apply_chat_template(tokenizer, params)
+ assert "get_weather" in prompt
+ assert "Tokyo" in prompt
+ assert "Sunny" in prompt
diff --git a/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py b/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py
new file mode 100644
index 00000000..907ccddb
--- /dev/null
+++ b/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py
@@ -0,0 +1,332 @@
+from collections.abc import Generator
+from typing import Any
+
+from exo.shared.types.worker.runner_response import (
+ FinishReason,
+ GenerationResponse,
+ ToolCallResponse,
+)
+from exo.worker.engines.mlx.dsml_encoding import (
+ DSML_TOKEN,
+ THINKING_END,
+ THINKING_START,
+ TOOL_CALLS_END,
+ TOOL_CALLS_START,
+)
+from exo.worker.runner.llm_inference.model_output_parsers import (
+ parse_deepseek_v32,
+ parse_thinking_models,
+ parse_tool_calls,
+)
+from exo.worker.runner.llm_inference.tool_parsers import make_mlx_parser
+
+
+def _make_response(
+ text: str, token: int, finish_reason: FinishReason | None = None
+) -> GenerationResponse:
+ return GenerationResponse(
+ text=text, token=token, finish_reason=finish_reason, usage=None
+ )
+
+
+def _queue_source(
+ tokens: list[GenerationResponse],
+) -> Generator[GenerationResponse | None]:
+ for token in tokens:
+ yield token
+ yield None
+ while True:
+ yield None
+
+
+def _step_until_finish(
+ parser_gen: Generator[GenerationResponse | ToolCallResponse | None],
+ max_steps: int = 200,
+) -> list[GenerationResponse | ToolCallResponse]:
+ results: list[GenerationResponse | ToolCallResponse] = []
+ for _ in range(max_steps):
+ try:
+ result = next(parser_gen)
+ except StopIteration:
+ break
+ if result is None:
+ continue
+ results.append(result)
+ if isinstance(result, GenerationResponse) and result.finish_reason is not None:
+ return results
+ if isinstance(result, ToolCallResponse):
+ return results
+ return results
+
+
+def _got_finish(results: list[GenerationResponse | ToolCallResponse]) -> bool:
+ for r in results:
+ if isinstance(r, ToolCallResponse):
+ return True
+ if r.finish_reason is not None:
+ return True
+ return False
+
+
+# ── parse_deepseek_v32 ──────────────────────────────────────────
+
+
+class TestDeepSeekV32FinishReason:
+ def test_finish_reason_with_buffered_dsml_prefix(self):
+ tokens = [
+ _make_response("Hello! The answer is x", 0),
+ _make_response("<", 1),
+ _make_response("", 2, finish_reason="stop"),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ assert _got_finish(results)
+ full_text = "".join(
+ r.text for r in results if isinstance(r, GenerationResponse)
+ )
+ assert "Hello" in full_text
+ assert "<" in full_text
+
+ def test_finish_reason_completes_tool_call_block(self):
+ tokens = [
+ _make_response(TOOL_CALLS_START, 0),
+ _make_response("\n", 1),
+ _make_response(f'<{DSML_TOKEN}invoke name="get_weather">\n', 2),
+ _make_response(
+ f'<{DSML_TOKEN}parameter name="city" string="true">Tokyo</{DSML_TOKEN}parameter>\n',
+ 3,
+ ),
+ _make_response(f"</{DSML_TOKEN}invoke>\n", 4),
+ _make_response(TOOL_CALLS_END, 5, finish_reason="stop"),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
+ assert len(tool_results) == 1
+ assert tool_results[0].tool_calls[0].name == "get_weather"
+
+ def test_finish_reason_mid_tool_call_before_close(self):
+ tokens = [
+ _make_response(TOOL_CALLS_START, 0),
+ _make_response("\n", 1),
+ _make_response(
+ f'<{DSML_TOKEN}invoke name="get_weather">\n', 2, finish_reason="stop"
+ ),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ assert _got_finish(results)
+
+ def test_finish_reason_single_token_complete_dsml_block(self):
+ dsml_block = (
+ f"{TOOL_CALLS_START}\n"
+ f'<{DSML_TOKEN}invoke name="get_weather">\n'
+ f'<{DSML_TOKEN}parameter name="city" string="true">Tokyo</{DSML_TOKEN}parameter>\n'
+ f"</{DSML_TOKEN}invoke>\n"
+ f"{TOOL_CALLS_END}"
+ )
+ tokens = [_make_response(dsml_block, 0, finish_reason="stop")]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
+ assert len(tool_results) == 1
+ assert tool_results[0].tool_calls[0].name == "get_weather"
+
+ def test_finish_reason_during_thinking(self):
+ tokens = [
+ _make_response(THINKING_START, 0),
+ _make_response("I need to think about this", 1),
+ _make_response(" carefully before responding", 2, finish_reason="stop"),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ assert _got_finish(results)
+
+ def test_finish_reason_after_thinking_then_tool_call(self):
+ tokens = [
+ _make_response(THINKING_START, 0),
+ _make_response("Let me check the weather.", 1),
+ _make_response(THINKING_END, 2),
+ _make_response("\n\n", 3),
+ _make_response(TOOL_CALLS_START, 4),
+ _make_response("\n", 5),
+ _make_response(f'<{DSML_TOKEN}invoke name="get_weather">\n', 6),
+ _make_response(
+ f'<{DSML_TOKEN}parameter name="city" string="true">NYC</{DSML_TOKEN}parameter>\n',
+ 7,
+ ),
+ _make_response(f"</{DSML_TOKEN}invoke>\n", 8),
+ _make_response(TOOL_CALLS_END, 9, finish_reason="stop"),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
+ assert len(tool_results) == 1
+ assert tool_results[0].tool_calls[0].name == "get_weather"
+
+ def test_finish_reason_normal_text_no_buffering(self):
+ tokens = [
+ _make_response("Hello", 0),
+ _make_response(" world", 1),
+ _make_response("!", 2, finish_reason="stop"),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ assert _got_finish(results)
+ full_text = "".join(
+ r.text for r in results if isinstance(r, GenerationResponse)
+ )
+ assert full_text == "Hello world!"
+
+ def test_finish_reason_multiple_buffered_prefix_tokens(self):
+ tokens = [
+ _make_response("text ", 0),
+ _make_response("<", 1),
+ _make_response("not a tag", 2),
+ _make_response(" more<", 3),
+ _make_response("", 4, finish_reason="stop"),
+ ]
+ results = _step_until_finish(parse_deepseek_v32(_queue_source(tokens)))
+ assert _got_finish(results)
+
+
+# ── parse_thinking_models ────────────────────────────────────────
+
+
+class TestThinkingModelsFinishReason:
+ def test_finish_reason_during_thinking(self):
+ tokens = [
+ _make_response("<think>", 0),
+ _make_response("reasoning here", 1),
+ _make_response("more reasoning", 2, finish_reason="stop"),
+ ]
+ results = _step_until_finish(
+ parse_thinking_models(
+ _queue_source(tokens),
+ think_start="<think>",
+ think_end="</think>",
+ starts_in_thinking=False,
+ )
+ )
+ assert _got_finish(results)
+ last_gen = [
+ r
+ for r in results
+ if isinstance(r, GenerationResponse) and r.finish_reason is not None
+ ]
+ assert len(last_gen) == 1
+ assert last_gen[0].is_thinking is False
+
+ def test_finish_reason_after_thinking(self):
+ tokens = [
+ _make_response("<think>", 0),
+ _make_response("hmm", 1),
+ _make_response("</think>", 2),
+ _make_response("The answer is 42.", 3, finish_reason="stop"),
+ ]
+ results = _step_until_finish(
+ parse_thinking_models(
+ _queue_source(tokens),
+ think_start="<think>",
+ think_end="</think>",
+ starts_in_thinking=False,
+ )
+ )
+ assert _got_finish(results)
+
+ def test_finish_reason_starts_in_thinking(self):
+ tokens = [
+ _make_response("still thinking", 0),
+ _make_response("</think>", 1),
+ _make_response("done", 2, finish_reason="stop"),
+ ]
+ results = _step_until_finish(
+ parse_thinking_models(
+ _queue_source(tokens),
+ think_start="<think>",
+ think_end="</think>",
+ starts_in_thinking=True,
+ )
+ )
+ assert _got_finish(results)
+
+
+# ── parse_tool_calls (generic) ──────────────────────────────────
+
+
+def _dummy_parser_fn(text: str) -> dict[str, Any]:
+ return {"name": "test_fn", "arguments": {"arg": text}}
+
+
+_dummy_parser = make_mlx_parser("<tool_call>", "</tool_call>", _dummy_parser_fn)
+
+
+class TestGenericToolCallsFinishReason:
+ def test_finish_reason_after_complete_tool_call(self):
+ tokens = [
+ _make_response("<tool_call>", 0),
+ _make_response("body", 1),
+ _make_response("</tool_call>", 2),
+ _make_response("extra text", 3, finish_reason="stop"),
+ ]
+ results = _step_until_finish(
+ parse_tool_calls(
+ _queue_source(tokens),
+ _dummy_parser,
+ tools=None,
+ )
+ )
+ tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
+ assert len(tool_results) == 1
+
+ def test_finish_reason_mid_tool_call_unclosed(self):
+ tokens = [
+ _make_response("<tool_call>", 0),
+ _make_response("partial content", 1, finish_reason="stop"),
+ ]
+ results = _step_until_finish(
+ parse_tool_calls(
+ _queue_source(tokens),
+ _dummy_parser,
+ tools=None,
+ )
+ )
+ assert _got_finish(results)
+
+ def test_finish_reason_no_tool_calls(self):
+ tokens = [
+ _make_response("Just", 0),
+ _make_response(" a", 1),
+ _make_response(" normal", 2),
+ _make_response(" response.", 3, finish_reason="stop"),
+ ]
+ results = _step_until_finish(
+ parse_tool_calls(
+ _queue_source(tokens),
+ _dummy_parser,
+ tools=None,
+ )
+ )
+ assert _got_finish(results)
+
+
+# ── Double parser chain (parse_thinking_models → parse_deepseek_v32) ──
+
+
+class TestBatchGeneratorSingleNext:
+ def test_finish_reason_with_buffered_tokens_drain_loop(self):
+ from exo.worker.runner.llm_inference.batch_generator import GeneratorQueue
+
+ queue: GeneratorQueue[GenerationResponse] = GeneratorQueue()
+ parser = parse_deepseek_v32(queue.gen())
+
+ tokens = [
+ _make_response("Hello ", 0),
+ _make_response(" `<", 1),
+ _make_response("", 2, finish_reason="stop"),
+ ]
+
+ collected: list[GenerationResponse | ToolCallResponse] = []
+ for token in tokens:
+ queue.push(token)
+ while (parsed := next(parser, None)) is not None:
+ collected.append(parsed)
+ if token.finish_reason is not None:
+ break
+
+ assert _got_finish(collected), (
+ f"No finish_reason in collected: {[(type(r).__name__, getattr(r, 'finish_reason', None) if isinstance(r, GenerationResponse) else 'tool') for r in collected]}"
+ )
diff --git a/tests/run_exo_on.sh b/tests/run_exo_on.sh
index 3cbc3bc0..12db1103 100755
--- a/tests/run_exo_on.sh
+++ b/tests/run_exo_on.sh
@@ -11,10 +11,17 @@ set -euo pipefail
exit 1
}
+upstream=$(git rev-parse --abbrev-ref --symbolic-full-name "@{u}" 2>/dev/null) || {
+ echo "No upstream"
+ exit 1
+}
commit=$(git rev-parse HEAD)
-git fetch -q origin
-git branch -r --contains "$commit" | grep -qE '^\s*origin/' || {
- echo "Not pushed to origin"
+remote=${upstream%%/*}
+remote_installable=$(git remote get-url "$remote" | sed -E "s#^(git@github.com:|https://github\.com/)([^/]+)/([^/]+)(\.git)?\$#github:\2/\3/$commit#")
+
+git fetch -q "$remote"
+git branch -r --contains "$commit" | grep -qE "^[[:space:]]*$remote/" || {
+ echo "Not pushed to $remote"
exit 1
}
@@ -35,7 +42,7 @@ i=0
for host; do
colour=${colours[i++ % 4]}
ssh -T -o BatchMode=yes -o ServerAliveInterval=30 "$host@$host" \
- "EXO_LIBP2P_NAMESPACE=$commit /nix/var/nix/profiles/default/bin/nix run github:exo-explore/exo/$commit" |&
+ "EXO_LIBP2P_NAMESPACE=$commit /nix/var/nix/profiles/default/bin/nix run $remote_installable" 2>&1 |
awk -v p="${colour}[${host}]${reset}" '{ print p $0; fflush() }' &
done
diff --git a/uv.lock b/uv.lock
index 431458fc..f56bc5f4 100644
--- a/uv.lock
+++ b/uv.lock
@@ -524,7 +524,7 @@ requires-dist = [
{ name = "mflux", specifier = "==0.17.2" },
{ name = "mlx", marker = "sys_platform == 'darwin'", git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks" },
{ name = "mlx", extras = ["cpu"], marker = "sys_platform == 'linux'", specifier = "==0.30.6" },
- { name = "mlx-lm", git = "https://github.com/rltakashige/mlx-lm?branch=fix%2Ffloat32-logprobs" },
+ { name = "mlx-lm", git = "https://github.com/rltakashige/mlx-lm?branch=leo%2Ffix-deepseek-v32-indexer" },
{ name = "msgspec", specifier = ">=0.19.0" },
{ name = "openai-harmony", specifier = ">=0.0.8" },
{ name = "psutil", specifier = ">=7.0.0" },
@@ -1446,7 +1446,7 @@ wheels = [
[[package]]
name = "mlx-lm"
version = "0.31.2"
-source = { git = "https://github.com/rltakashige/mlx-lm?branch=fix%2Ffloat32-logprobs#8e94256220f954949133e036980951681e353945" }
+source = { git = "https://github.com/rltakashige/mlx-lm?branch=leo%2Ffix-deepseek-v32-indexer#d388ff77858fec3b5d2e3b1d9502a7e2878b8109" }
dependencies = [
{ name = "jinja2", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" },
{ name = "mlx", version = "0.31.2.dev20260324+e5e64331", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#e5e64331830d9b04ae9082b843073f9c1fa7705e" }, marker = "sys_platform == 'darwin'" },
← 565ed41c Fix occasional warmup bugs by using mlx_generate (#1794)
·
back to Exo
·
Support image generation cancellation (#1774) 6de14cfe →