From 4abdaaf74bda83ea2fd449de561f71ee06515e92 Mon Sep 17 00:00:00 2001 From: rltakashige Date: Tue, 10 Feb 2026 11:53:23 +0000 Subject: [PATCH] Address GPU timeouts (#1429) ## Motivation For large prompts and/or slow machines, users are running into GPU timeout errors very often. ## Changes Only during prefill, we eval distributed operations. We don't do this during decode to maintain decode performance. Raise the prefill step size to 8192 because now we can (we see a speedup here). We also now see a 2x speedup in pipeline parallel prefill by disabling an unnecessary all_gather during prefill. ## Why It Works GPU timeout errors happen in the Metal backend when GPU operations take too long without making progress. By isolating distributed operations, we can allow them to run without any timeouts. ## Test Plan ### Manual Testing Doesn't GPU timeout on 100k tokens on Minimax anymore. Also tested on Kimi. ### Automated Testing Needs more exo bench, but I think this is a good step in the right direction. --- .mlx_typings/mlx/core/__init__.pyi | 2 +- src/exo/worker/engines/mlx/auto_parallel.py | 23 ++++++++++++++++--- .../worker/engines/mlx/generator/generate.py | 13 ++++++++--- 3 files changed, 31 insertions(+), 7 deletions(-) diff --git a/.mlx_typings/mlx/core/__init__.pyi b/.mlx_typings/mlx/core/__init__.pyi index 025b4ab2..cabcbfd5 100644 --- a/.mlx_typings/mlx/core/__init__.pyi +++ b/.mlx_typings/mlx/core/__init__.pyi @@ -2366,7 +2366,7 @@ class custom_function: def default_device() -> Device: """Get the default device.""" -def default_stream(device: Device) -> Stream: +def default_stream(device: Device | DeviceType) -> Stream: """Get the device's default stream.""" def degrees(a: array, /, *, stream: Stream | Device | None = ...) -> array: diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 04e47b2c..91572a70 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -121,10 +121,15 @@ class PipelineFirstLayer(CustomMlxLayer): super().__init__(original_layer) self.r: int = r self.group = group + self.is_prefill: bool = False def __call__(self, x: mx.array, *args: object, **kwargs: object) -> mx.array: if self.r != 0: x = mx.distributed.recv_like(x, (self.r - 1), group=self.group) + if self.is_prefill: + # We want to avoid GPU timeout errors by evalling the distributed operation + # so that it stays on CPU, which does not have a timeout. + mx.eval(x) return self.original_layer(x, *args, **kwargs) @@ -141,6 +146,7 @@ class PipelineLastLayer(CustomMlxLayer): self.s: int = s self.group = group self.original_layer_signature = signature(self.original_layer.__call__) + self.is_prefill: bool = False def __call__(self, x: mx.array, *args: object, **kwargs: object) -> mx.array: cache = self.original_layer_signature.bind_partial( @@ -155,14 +161,25 @@ class PipelineLastLayer(CustomMlxLayer): ) if cache is not None: cache.keys = mx.depends(cache.keys, output) # type: ignore[reportUnknownMemberType] + if self.is_prefill: + mx.eval(output) + if cache is not None: + mx.eval(cache.keys) # type: ignore - output = mx.distributed.all_gather(output, group=self.group)[ - -output.shape[0] : - ] # type :ignore + if not self.is_prefill: + output = mx.distributed.all_gather(output, group=self.group)[ + -output.shape[0] : + ] return output +def set_pipeline_prefill(model: nn.Module, is_prefill: bool) -> None: + for layer in model.layers: # type: ignore + if isinstance(layer, (PipelineFirstLayer, PipelineLastLayer)): + layer.is_prefill = is_prefill + + def _inner_model(model: nn.Module) -> nn.Module: inner = getattr(model, "model", None) if isinstance(inner, nn.Module): diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py index 7063f8de..98c61fc4 100644 --- a/src/exo/worker/engines/mlx/generator/generate.py +++ b/src/exo/worker/engines/mlx/generator/generate.py @@ -24,6 +24,7 @@ from exo.shared.types.worker.runner_response import ( GenerationResponse, ) from exo.worker.engines.mlx import Model +from exo.worker.engines.mlx.auto_parallel import set_pipeline_prefill from exo.worker.engines.mlx.cache import ( CacheSnapshot, KVPrefixCache, @@ -83,6 +84,8 @@ def prefill( if has_ssm: snapshots.append(snapshot_ssm_states(cache)) + set_pipeline_prefill(model, is_prefill=True) + # Use max_tokens=1 because max_tokens=0 does not work. # We just throw away the generated token - we only care about filling the cache for _ in stream_generate( @@ -92,13 +95,15 @@ def prefill( max_tokens=1, sampler=sampler, prompt_cache=cache, - prefill_step_size=2048, + prefill_step_size=8192, kv_group_size=KV_GROUP_SIZE, kv_bits=KV_BITS, prompt_progress_callback=progress_callback, ): break # Stop after first iteration - cache is now filled + set_pipeline_prefill(model, is_prefill=False) + # stream_generate added 1 extra generated token to the cache, so we should trim it. # Because of needing to roll back arrays cache, we will generate on 2 tokens so trim 1 more. pre_gen = deepcopy(snapshots[-2]) if has_ssm else None @@ -325,6 +330,9 @@ def mlx_generate( reasoning_tokens = 0 think_start = tokenizer.think_start think_end = tokenizer.think_end + + mx_barrier(group) + for completion_tokens, out in enumerate( stream_generate( model=model, @@ -334,8 +342,7 @@ def mlx_generate( sampler=sampler, logits_processors=logits_processors, prompt_cache=caches, - # TODO: Dynamically change prefill step size to be the maximum possible without timing out. - prefill_step_size=2048, + prefill_step_size=1, kv_group_size=KV_GROUP_SIZE, kv_bits=KV_BITS, ),