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, ),