diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 159c2a45..443e5f0a 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -114,6 +114,7 @@ class PipelineFirstLayer(CustomMlxLayer): 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) + mx.eval(x) return self.original_layer(x, *args, **kwargs) @@ -142,6 +143,7 @@ class PipelineLastLayer(CustomMlxLayer): output = mx.distributed.send( output, (self.r + 1) % self.s, group=self.group ) + mx.async_eval(output) if cache is not None: cache.keys = mx.depends(cache.keys, output) # type: ignore[reportUnknownMemberType] @@ -248,6 +250,10 @@ def patch_pipeline_model[T](model: T, group: mx.distributed.Group) -> T: "cache", None ) + # Evaluate logits before all_gather to break the computation graph + # and prevent Metal command buffer timeouts with large batches + mx.eval(logits) + # Add dependency to last cache entry to ensure distributed ops are evaluated if cache is not None: cache[-1].state = mx.depends(cache[-1].state, logits) # type: ignore diff --git a/src/exo/worker/runner/batched_handler.py b/src/exo/worker/runner/batched_handler.py index 2cc00e41..361495d4 100644 --- a/src/exo/worker/runner/batched_handler.py +++ b/src/exo/worker/runner/batched_handler.py @@ -32,6 +32,7 @@ from exo.worker.engines.mlx.constants import MAX_TOKENS from exo.worker.engines.mlx.generator.generate import extract_top_logprobs from exo.worker.engines.mlx.utils_mlx import apply_chat_template from exo.worker.runner.bootstrap import logger +from exo.worker.runner.pipelined_generator import PipelinedGenerator, PipelinedResponse # Type alias for the finish_reason values TokenChunk accepts TokenFinishReason = Literal["stop", "length", "content_filter"] @@ -78,12 +79,14 @@ class BatchedInferenceHandler: tokenizer: TokenizerWrapper, model_id: ModelId, device_rank: int, + world_size: int = 1, max_batch_size: int = 8, ): self.model = model self.tokenizer = tokenizer self.model_id = model_id self.device_rank = device_rank + self.world_size = world_size self.max_batch_size = max_batch_size # GPT-OSS model detection @@ -98,8 +101,14 @@ class BatchedInferenceHandler: # Active batch generator and request tracking self.batch_generator: BatchGenerator | None = None + self.pipelined_generator: PipelinedGenerator | None = None self.uid_to_request: dict[int, ActiveRequest] = {} + # Use pipelined generator for multi-device pipeline parallelism + self.use_pipelined = world_size > 1 + if self.use_pipelined: + logger.info(f"Using PipelinedGenerator with {world_size} streams for pipeline overlap") + # EOS tokens for the model self.stop_tokens: set[int] = set() eos_ids: list[int] | None = getattr(tokenizer, "eos_token_ids", None) @@ -109,6 +118,8 @@ class BatchedInferenceHandler: @property def is_active(self) -> bool: """Check if there's an active batch being processed.""" + if self.use_pipelined: + return self.pipelined_generator is not None and self.pipelined_generator.has_active return self.batch_generator is not None and len(self.uid_to_request) > 0 @property @@ -157,7 +168,7 @@ class BatchedInferenceHandler: ) def flush(self) -> None: - """Start processing pending requests by adding them to the BatchGenerator.""" + """Start processing pending requests by adding them to the batch/pipelined generator.""" if not self.has_pending: return @@ -166,20 +177,7 @@ class BatchedInferenceHandler: requests_to_flush = self.pending[:available_slots] self.pending = self.pending[available_slots:] - # Create batch generator if not exists - if self.batch_generator is None: - logger.info(f"Creating new BatchGenerator for {len(requests_to_flush)} requests") - mx.reset_peak_memory() - self.batch_generator = BatchGenerator( - model=self.model, - max_tokens=MAX_TOKENS, - stop_tokens=self.stop_tokens if self.stop_tokens else None, - prefill_batch_size=1, - ) - else: - logger.info(f"Adding {len(requests_to_flush)} requests to existing BatchGenerator") - - # Prepare batch data - tokenize prompts since BatchGenerator expects token IDs + # Prepare batch data - tokenize prompts tokenized_prompts: list[list[int]] = [] max_tokens_list: list[int] = [] samplers: list[Callable[[mx.array], mx.array]] = [] @@ -192,15 +190,80 @@ class BatchedInferenceHandler: samplers.append(req.sampler) prompt_token_counts.append(len(tokens)) + if self.use_pipelined: + self._flush_pipelined(requests_to_flush, tokenized_prompts, max_tokens_list, samplers, prompt_token_counts) + else: + self._flush_batch(requests_to_flush, tokenized_prompts, max_tokens_list, samplers, prompt_token_counts) + + def _flush_pipelined( + self, + requests_to_flush: list[PendingRequest], + tokenized_prompts: list[list[int]], + max_tokens_list: list[int], + samplers: list[Callable[[mx.array], mx.array]], + prompt_token_counts: list[int], + ) -> None: + """Flush using PipelinedGenerator (multi-stream pipeline overlap).""" + if self.pipelined_generator is None: + logger.info(f"Creating PipelinedGenerator for {len(requests_to_flush)} requests ({self.world_size} streams)") + mx.reset_peak_memory() + self.pipelined_generator = PipelinedGenerator( + model=self.model, + world_size=self.world_size, + stop_tokens=self.stop_tokens if self.stop_tokens else None, + max_tokens=MAX_TOKENS, + ) + else: + logger.info(f"Adding {len(requests_to_flush)} requests to PipelinedGenerator") + + uids = self.pipelined_generator.insert( + prompts=tokenized_prompts, + max_tokens=max_tokens_list, + samplers=samplers, + ) + + for uid, req, prompt_tokens in zip(uids, requests_to_flush, prompt_token_counts, strict=True): + parser = None + if self.is_gpt_oss and self._gpt_oss_encoding is not None: + parser = StreamableParser(self._gpt_oss_encoding, role=Role.ASSISTANT) # pyright: ignore[reportAny] + self.uid_to_request[uid] = ActiveRequest( + command_id=req.task.command_id, + should_extract_logprobs=req.should_extract_logprobs, + top_k=req.top_k, + prompt_tokens=prompt_tokens, + gpt_oss_parser=parser, + ) + + logger.info(f"Flushed {len(requests_to_flush)} requests into pipelined generator (active={self.pipelined_generator.active_count}, uids={list(self.uid_to_request.keys())})") + + def _flush_batch( + self, + requests_to_flush: list[PendingRequest], + tokenized_prompts: list[list[int]], + max_tokens_list: list[int], + samplers: list[Callable[[mx.array], mx.array]], + prompt_token_counts: list[int], + ) -> None: + """Flush using BatchGenerator (single-stream, for non-pipeline instances).""" + if self.batch_generator is None: + logger.info(f"Creating new BatchGenerator for {len(requests_to_flush)} requests") + mx.reset_peak_memory() + self.batch_generator = BatchGenerator( + model=self.model, + max_tokens=MAX_TOKENS, + stop_tokens=self.stop_tokens if self.stop_tokens else None, + prefill_batch_size=1, + ) + else: + logger.info(f"Adding {len(requests_to_flush)} requests to existing BatchGenerator") + # Insert into batch generator - # Note: BatchGenerator.insert() accepts samplers param at runtime but pyright doesn't see it uids: list[int] = self.batch_generator.insert( # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] prompts=tokenized_prompts, max_tokens=max_tokens_list, samplers=samplers, # pyright: ignore[reportCallIssue] ) - # Track active requests for uid, req, prompt_tokens in zip(uids, requests_to_flush, prompt_token_counts, strict=True): # pyright: ignore[reportUnknownArgumentType] parser = None if self.is_gpt_oss and self._gpt_oss_encoding is not None: @@ -221,6 +284,10 @@ class BatchedInferenceHandler: Returns a generator of events for completed tokens across all active requests. """ + if self.use_pipelined: + yield from self._step_pipelined() + return + if self.batch_generator is None or not self.uid_to_request: return @@ -348,6 +415,121 @@ class BatchedInferenceHandler: for uid in completed_uids: del self.uid_to_request[uid] + def _step_pipelined(self) -> Generator[Event, None, None]: + """Process one generation step using the multi-stream PipelinedGenerator.""" + if self.pipelined_generator is None or not self.uid_to_request: + return + + logger.debug(f"PipelinedGenerator.next() called (active={self.pipelined_generator.active_count})") + responses: list[PipelinedResponse] = self.pipelined_generator.next() + logger.debug(f"PipelinedGenerator.next() returned {len(responses)} responses") + + completed_uids: list[int] = [] + + for response in responses: + uid = response.uid + if uid not in self.uid_to_request: + logger.warning(f"Received response for unknown uid: {uid}") + continue + + active_request = self.uid_to_request[uid] + active_request.tokens_generated += 1 + + resp_token: int = response.token + resp_finish_reason: str | None = response.finish_reason + resp_logprobs: mx.array = response.logprobs + + # Only emit events from device_rank 0 + if self.device_rank != 0: + if resp_finish_reason is not None: + completed_uids.append(uid) + continue + + # Decode token to text + token_text = self.tokenizer.decode([resp_token]) + if active_request.gpt_oss_parser is not None: + parser = active_request.gpt_oss_parser # pyright: ignore[reportAny] + parser.process(resp_token) # pyright: ignore[reportAny] + delta: str | None = parser.last_content_delta # pyright: ignore[reportAny] + channel: str = parser.current_channel # pyright: ignore[reportAny] + + if channel == "analysis": + active_request.reasoning_tokens += 1 + + prefix = "" + if channel == "analysis" and not active_request.gpt_oss_thinking: + active_request.gpt_oss_thinking = True + prefix = "" + elif channel != "analysis" and active_request.gpt_oss_thinking: + active_request.gpt_oss_thinking = False + prefix = "" + + if resp_finish_reason is not None and active_request.gpt_oss_thinking: + prefix = "" + active_request.gpt_oss_thinking = False + + effective_delta = delta or "" + token_text = prefix + effective_delta if (prefix or effective_delta) else "" + if not token_text and resp_finish_reason is None: + continue + + # Extract logprobs if requested + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None + if active_request.should_extract_logprobs: + logprob, top_logprobs = extract_top_logprobs( + logprobs_array=resp_logprobs, + selected_token=resp_token, + tokenizer=self.tokenizer, + top_k=active_request.top_k, + ) + + # Build stats for final token + stats: GenerationStats | None = None + finish_reason: TokenFinishReason | None = None + if resp_finish_reason is not None: + elapsed_time = time.perf_counter() - active_request.start_time + prompt_tps = active_request.prompt_tokens / max(elapsed_time, 0.001) + generation_tps = active_request.tokens_generated / max(elapsed_time, 0.001) + + peak_memory_bytes = 0 + if mx.metal.is_available(): + peak_memory_bytes = mx.metal.get_peak_memory() + + stats = GenerationStats( + prompt_tps=prompt_tps, + generation_tps=generation_tps, + prompt_tokens=active_request.prompt_tokens, + generation_tokens=active_request.tokens_generated, + reasoning_tokens=active_request.reasoning_tokens, + peak_memory_usage=Memory.from_bytes(peak_memory_bytes), + ) + + if resp_finish_reason == "stop": + finish_reason = "stop" + elif resp_finish_reason == "length": + finish_reason = "length" + else: + finish_reason = "stop" + + completed_uids.append(uid) + + yield ChunkGenerated( + command_id=active_request.command_id, + chunk=TokenChunk( + model=self.model_id, + text=token_text, + token_id=resp_token, + logprob=logprob, + top_logprobs=top_logprobs, + finish_reason=finish_reason, + stats=stats, + ), + ) + + for uid in completed_uids: + del self.uid_to_request[uid] + def emit_error(self, command_id: CommandId, error_message: str) -> Event: """Create an error event for a failed request.""" return ChunkGenerated( @@ -360,12 +542,15 @@ class BatchedInferenceHandler: ) def _close_generator(self) -> None: - """Close and clean up the batch generator.""" + """Close and clean up the batch/pipelined generator.""" if self.batch_generator is not None: self.batch_generator.close() # pyright: ignore[reportUnknownMemberType,reportAttributeAccessIssue] self.batch_generator = None - self.uid_to_request.clear() - logger.info("Batch generator closed") + if self.pipelined_generator is not None: + self.pipelined_generator.close() + self.pipelined_generator = None + self.uid_to_request.clear() + logger.info("Generator closed") def close(self) -> None: """Close the handler and clean up resources.""" diff --git a/src/exo/worker/runner/pipelined_generator.py b/src/exo/worker/runner/pipelined_generator.py new file mode 100644 index 00000000..725cf052 --- /dev/null +++ b/src/exo/worker/runner/pipelined_generator.py @@ -0,0 +1,329 @@ +"""Multi-stream pipelined batch generator for pipeline-parallel inference. + +When a model is split across N ranks (pipeline parallelism), each rank's GPU is idle +for (N-1)/N of each step while waiting for other ranks to compute their layers. + +This module fills the pipeline bubble by splitting sequences into N micro-batch groups +and processing each group on a different MLX stream. The GPU can overlap one stream's +network communication (send/recv/all_gather) with another stream's compute. +""" + +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false +# pyright: reportUnknownArgumentType=false, reportAny=false + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + +import mlx.core as mx +import mlx.nn as nn +from mlx_lm.models.cache import make_prompt_cache + + +@dataclass +class MicroBatch: + """State for one micro-batch group of sequences.""" + + uids: list[int] + y: mx.array # Last sampled tokens [batch] + logprobs: list[mx.array] # Logprobs for each sequence + max_tokens: list[int] + num_tokens: list[int] + cache: list[Any] # KV cache (list of layer caches) + samplers: list[Callable[[mx.array], mx.array]] + tokens: list[mx.array] # All tokens generated so far per sequence + + def __len__(self) -> int: + return len(self.uids) + + +@dataclass +class PipelinedResponse: + """Response from one generation step.""" + + uid: int + token: int + logprobs: mx.array + finish_reason: str | None + cache: list[Any] | None = None + + +@dataclass +class PendingPrompt: + """A prompt waiting to be prefilled.""" + + uid: int + tokens: list[int] + max_tokens: int + sampler: Callable[[mx.array], mx.array] + + +class PipelinedGenerator: + """ + Multi-stream batch generator that fills pipeline bubbles. + + Splits active sequences into `world_size` micro-batch groups, each processed + on its own MLX stream. During mx.eval(), the GPU overlaps network operations + on one stream with compute on another. + """ + + def __init__( + self, + model: nn.Module, + world_size: int, + stop_tokens: set[int] | None = None, + max_tokens: int = 4096, + ): + self.model = model + self.world_size = world_size + self.stop_tokens = stop_tokens or set() + self.max_tokens_default = max_tokens + + # Create one stream per pipeline stage + self.streams = [mx.new_stream(mx.default_device()) for _ in range(world_size)] + + # Micro-batch groups (one per stream) + self.micro_batches: list[MicroBatch | None] = [None] * world_size + + # Pending prompts to be inserted + self.pending_prompts: list[PendingPrompt] = [] + + # UID counter + self._next_uid = 0 + + @property + def active_count(self) -> int: + """Total number of active sequences across all micro-batches.""" + return sum(len(mb) for mb in self.micro_batches if mb is not None) + + @property + def has_active(self) -> bool: + return self.active_count > 0 + + def insert( + self, + prompts: list[list[int]], + max_tokens: list[int], + samplers: list[Callable[[mx.array], mx.array]], + ) -> list[int]: + """Queue prompts for processing. Returns assigned UIDs.""" + uids: list[int] = [] + for prompt, mt, sampler in zip(prompts, max_tokens, samplers, strict=True): + uid = self._next_uid + self._next_uid += 1 + self.pending_prompts.append( + PendingPrompt(uid=uid, tokens=prompt, max_tokens=mt, sampler=sampler) + ) + uids.append(uid) + return uids + + def _prefill_group(self, group_idx: int, prompts: list[PendingPrompt]) -> None: + """Prefill a group of prompts and create a MicroBatch.""" + if not prompts: + return + + stream = self.streams[group_idx] + + with mx.stream(stream): + # Create per-sequence caches + caches = [make_prompt_cache(self.model) for _ in prompts] + + # Tokenize and prefill each sequence + all_y: list[mx.array] = [] + all_logprobs: list[mx.array] = [] + all_samplers: list[Callable[[mx.array], mx.array]] = [] + all_tokens: list[mx.array] = [] + + for prompt_info, cache in zip(prompts, caches, strict=True): + tokens = mx.array(prompt_info.tokens) + # Run prefill (process all tokens except last) + if len(prompt_info.tokens) > 1: + self.model(tokens[:-1][None, :], cache=cache) + mx.eval([c.state for c in cache]) + + # Process last token to get first generation logits + last_token = tokens[-1:][None, :] + logits = self.model(last_token, cache=cache) + logits = logits[:, -1, :] + logprobs = logits - mx.logsumexp(logits, axis=-1, keepdims=True) + sampled = prompt_info.sampler(logprobs) + + all_y.append(sampled.squeeze(0)) + all_logprobs.append(logprobs.squeeze(0)) + all_samplers.append(prompt_info.sampler) + all_tokens.append(tokens) + + mx.eval(*all_y, *all_logprobs) + + # Create micro-batch + batch = MicroBatch( + uids=[p.uid for p in prompts], + y=mx.stack(all_y), + logprobs=all_logprobs, + max_tokens=[p.max_tokens for p in prompts], + num_tokens=[0] * len(prompts), + cache=caches, + samplers=all_samplers, + tokens=all_tokens, + ) + + if self.micro_batches[group_idx] is None: + self.micro_batches[group_idx] = batch + else: + # Extend existing micro-batch (would need cache merging - for now replace) + self.micro_batches[group_idx] = batch + + def _prefill_pending(self) -> None: + """Distribute pending prompts across micro-batch groups and prefill.""" + if not self.pending_prompts: + return + + # Distribute round-robin across groups + groups: list[list[PendingPrompt]] = [[] for _ in range(self.world_size)] + for i, prompt in enumerate(self.pending_prompts): + groups[i % self.world_size].append(prompt) + self.pending_prompts.clear() + + for group_idx, group_prompts in enumerate(groups): + if group_prompts: + self._prefill_group(group_idx, group_prompts) + + def _step_all(self) -> None: + """ + Run one generation step across all micro-batch groups on different streams. + + This is where pipeline overlap happens: each group's model forward pass + runs on its own stream, and mx.eval() allows the GPU to overlap network + ops (send/recv/all_gather) from one stream with compute from another. + """ + # Build computation graphs on each stream (lazy, no evaluation yet) + new_y_list: list[mx.array] = [] + new_logprobs_list: list[list[mx.array]] = [] + active_indices: list[int] = [] + + for i, mb in enumerate(self.micro_batches): + if mb is None or len(mb) == 0: + continue + active_indices.append(i) + + with mx.stream(self.streams[i]): + # Prepare input: last sampled tokens + input_tokens = mb.y[:, None] # [batch, 1] + + # Forward pass (lazy graph construction) + # For pipeline models, this includes send/recv/all_gather ops + logits = self.model(input_tokens, cache=mb.cache) + logits = logits[:, -1, :] # [batch, vocab] + + # Compute logprobs + logprobs = logits - mx.logsumexp(logits, axis=-1, keepdims=True) + + # Sample per-sequence + batch_size = len(mb) + if batch_size == 1: + sampled = mb.samplers[0](logprobs) + else: + samples = [] + for e in range(batch_size): + samples.append(mb.samplers[e](logprobs[e: e + 1])) + sampled = mx.concatenate(samples, axis=0) + + new_y_list.append(sampled) + new_logprobs_list.append([logprobs[e] for e in range(batch_size)]) + + if not active_indices: + return + + # Evaluate ALL streams together - this is where overlap happens! + # The GPU can execute stream0's all_gather while computing stream1's layers. + mx.eval(*new_y_list, *[lp for lps in new_logprobs_list for lp in lps]) + + # Update micro-batch states with results + for list_idx, group_idx in enumerate(active_indices): + mb = self.micro_batches[group_idx] + assert mb is not None + mb.y = new_y_list[list_idx] + mb.logprobs = new_logprobs_list[list_idx] + # Append sampled tokens to history + for e in range(len(mb)): + mb.tokens[e] = mx.concatenate([mb.tokens[e], mb.y[e: e + 1]]) + + def next(self) -> list[PipelinedResponse]: + """ + Run one generation step and return responses. + + Returns a PipelinedResponse for each active sequence (across all groups). + Finished sequences are removed from their micro-batch. + """ + # Prefill any pending prompts first + self._prefill_pending() + + if not self.has_active: + return [] + + # Run the multi-stream forward pass + self._step_all() + + # Collect responses and filter completed sequences + responses: list[PipelinedResponse] = [] + + for group_idx, mb in enumerate(self.micro_batches): + if mb is None or len(mb) == 0: + continue + + keep_idx: list[int] = [] + end_idx: list[int] = [] + + for e in range(len(mb)): + token = int(mb.y[e].item()) + uid = mb.uids[e] + num_tok = mb.num_tokens[e] + 1 + max_tok = mb.max_tokens[e] + mb.num_tokens[e] = num_tok + + if token in self.stop_tokens: + finish_reason = "stop" + end_idx.append(e) + elif num_tok >= max_tok: + finish_reason = "length" + end_idx.append(e) + else: + finish_reason = None + keep_idx.append(e) + + responses.append( + PipelinedResponse( + uid=uid, + token=token, + logprobs=mb.logprobs[e], + finish_reason=finish_reason, + ) + ) + + # Remove finished sequences + if end_idx: + if keep_idx: + # Filter the micro-batch to keep only active sequences + mb.uids = [mb.uids[i] for i in keep_idx] + mb.y = mb.y[mx.array(keep_idx)] + mb.logprobs = [mb.logprobs[i] for i in keep_idx] + mb.max_tokens = [mb.max_tokens[i] for i in keep_idx] + mb.num_tokens = [mb.num_tokens[i] for i in keep_idx] + mb.samplers = [mb.samplers[i] for i in keep_idx] + mb.tokens = [mb.tokens[i] for i in keep_idx] + # Cache filtering: trim batch dimension + for c in mb.cache: + if hasattr(c, "keys") and c.keys is not None: + c.keys = c.keys[mx.array(keep_idx)] + c.values = c.values[mx.array(keep_idx)] + else: + self.micro_batches[group_idx] = None + + return responses + + def close(self) -> None: + """Clean up resources.""" + self.micro_batches = [None] * self.world_size + self.pending_prompts.clear() diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 394c340e..a7925722 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -89,7 +89,7 @@ from exo.worker.runner.bootstrap import logger # Batching configuration BATCH_ENABLED = True -BATCH_MAX_SIZE = 128 +BATCH_MAX_SIZE = 64 def _should_use_serial_processing( @@ -303,10 +303,11 @@ def main( tokenizer=tokenizer, model_id=shard_metadata.model_card.model_id, device_rank=device_rank, + world_size=shard_metadata.world_size, max_batch_size=BATCH_MAX_SIZE, ) logger.info( - f"Batch handler initialized (max_batch_size={BATCH_MAX_SIZE})" + f"Batch handler initialized (max_batch_size={BATCH_MAX_SIZE}, world_size={shard_metadata.world_size})" ) elif ( ModelTask.TextToImage in shard_metadata.model_card.tasks