Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions src/exo/worker/engines/mlx/generator/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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")

Expand Down
22 changes: 21 additions & 1 deletion src/exo/worker/engines/mlx/patches/opt_batch_gen.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
Original file line number Diff line number Diff line change
@@ -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]
Loading