From bc0637406b8069f515f6cf6ac476e6f781ccdf13 Mon Sep 17 00:00:00 2001 From: Alex Cheema Date: Mon, 5 Oct 2026 18:30:55 +0100 Subject: [PATCH] fix(mlx): wait for the in-flight decode step before a prompt eval mlx-lm's BatchGenerator async-evaluates the batch's next decode step and then processes new prompts, and exo prefills a new request while the batch's last step is still running. For a tensor-parallel model with MLX_METAL_FAST_SYNCH, both evals queue their collectives on the one communication stream, while their GPU work spins waiting for it. When the prompt eval's spinning GPU work held the GPU while the step's last kernels still had to run (more likely with another model's runner on the same Mac), neither could finish: one rank's communication stream waited in Fence::wait for a fence update that its GPU never ran, and the other rank waited in all_reduce. Synchronize mlx-lm's generation stream before exo's prefill and before the batch processes new prompt tokens, so that no prompt eval overlaps a decode step. A request now waits for at most one decode step before its prefill; steady decoding is unchanged. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../worker/engines/mlx/generator/generate.py | 6 ++++ .../engines/mlx/patches/opt_batch_gen.py | 22 +++++++++++- .../test_prompt_waits_for_decode_step.py | 35 +++++++++++++++++++ 3 files changed, 62 insertions(+), 1 deletion(-) create mode 100644 src/exo/worker/tests/unittests/test_mlx/test_prompt_waits_for_decode_step.py diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py index 2e3d051251..c45e3751dd 100644 --- a/src/exo/worker/engines/mlx/generator/generate.py +++ b/src/exo/worker/engines/mlx/generator/generate.py @@ -6,6 +6,9 @@ from typing import Callable, Generator, cast, get_args import mlx.core as mx +from mlx_lm.generate import ( + generation_stream as mlx_lm_generation_stream, +) from mlx_lm.generate import ( maybe_quantize_kv_cache, stream_generate, @@ -326,6 +329,9 @@ def combined_progress_callback(processed: int, total: int) -> None: set_pipeline_prefill(model, is_prefill=True) + # The batch's next decode step may still be running: let it finish before this prefill's + # collectives, as for the batch's own prompts (see opt_batch_gen._prompt_after_decode_step) + mx.synchronize(mlx_lm_generation_stream) mx_barrier(group) logger.info("Starting prefill") diff --git a/src/exo/worker/engines/mlx/patches/opt_batch_gen.py b/src/exo/worker/engines/mlx/patches/opt_batch_gen.py index e7bff41966..0fad4c6d02 100644 --- a/src/exo/worker/engines/mlx/patches/opt_batch_gen.py +++ b/src/exo/worker/engines/mlx/patches/opt_batch_gen.py @@ -1,8 +1,9 @@ +from collections.abc import Callable from dataclasses import dataclass, field from typing import cast import mlx.core as mx -from mlx_lm.generate import GenerationBatch +from mlx_lm.generate import GenerationBatch, PromptProcessingBatch, generation_stream _PRECOMPUTE_TOP_K = 20 @@ -132,5 +133,24 @@ def _patched_step(self: GenerationBatch) -> tuple[list[int], list[mx.array]]: return token_list, current_lp +_original_prompt = cast( + Callable[[PromptProcessingBatch, list[list[int]]], None], + PromptProcessingBatch.prompt, +) + + +def _prompt_after_decode_step( + self: PromptProcessingBatch, tokens: list[list[int]] +) -> None: + """BatchGenerator._next async-evaluates the batch's next decode step and then processes new + prompts. Wait for that step first. With MLX_METAL_FAST_SYNCH, a tensor-parallel prompt eval + started while the step's collectives are still running can deadlock: its GPU work spins + waiting for the communication stream, which is still waiting for the step's GPU work.""" + if tokens: + mx.synchronize(generation_stream) + _original_prompt(self, tokens) + + def apply_batch_gen_patch() -> None: GenerationBatch._step = _patched_step + PromptProcessingBatch.prompt = _prompt_after_decode_step diff --git a/src/exo/worker/tests/unittests/test_mlx/test_prompt_waits_for_decode_step.py b/src/exo/worker/tests/unittests/test_mlx/test_prompt_waits_for_decode_step.py new file mode 100644 index 0000000000..35dd8f171a --- /dev/null +++ b/src/exo/worker/tests/unittests/test_mlx/test_prompt_waits_for_decode_step.py @@ -0,0 +1,35 @@ +from typing import cast + +import mlx.core as mx +import pytest +from mlx_lm.generate import PromptProcessingBatch + +from exo.worker.engines.mlx.patches import opt_batch_gen + + +def test_prompt_waits_for_the_decode_step_before_new_tokens( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[str] = [] + + def synchronize(stream: object = None) -> None: + calls.append("wait") + + def original_prompt(batch: PromptProcessingBatch, tokens: list[list[int]]) -> None: + calls.append(f"prompt {len(tokens)}") + + monkeypatch.setattr(mx, "synchronize", synchronize) + monkeypatch.setattr(opt_batch_gen, "_original_prompt", original_prompt) + batch = PromptProcessingBatch.__new__(PromptProcessingBatch) + + opt_batch_gen._prompt_after_decode_step(batch, [[1, 2, 3]]) # pyright: ignore[reportPrivateUsage] + # mlx-lm calls it every step, mostly with nothing to process: no need to wait then + opt_batch_gen._prompt_after_decode_step(batch, []) # pyright: ignore[reportPrivateUsage] + + assert calls == ["wait", "prompt 1", "prompt 0"] + + +def test_batch_gen_patch_installs_it() -> None: + opt_batch_gen.apply_batch_gen_patch() + installed = cast(object, PromptProcessingBatch.prompt) + assert installed is opt_batch_gen._prompt_after_decode_step # pyright: ignore[reportPrivateUsage]