diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 7397cd8e..e3c23e03 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -152,6 +152,18 @@ class PipelineLastLayer(CustomMlxLayer): if cache is not None: cache.keys = mx.depends(cache.keys, sent) # type: ignore[reportUnknownMemberType] + if self.r == self.s - 1: + mx.eval(logits) + logger.info(f"[patch_pipeline_model] send logits to rank 0, shape={logits.shape}") + sent = mx.distributed.send(logits, dst=0, group=group) + mx.eval(sent) + logger.info(f"[patch_pipeline_model] send done") + elif rank == 0: + logger.info(f"[patch_pipeline_model] recv_like logits from rank {self.s - 1}, shape={logits.shape}") + logits = mx.distributed.recv_like(logits, src=self.s - 1, group=group) + mx.eval(logits) + logger.info(f"[patch_pipeline_model] recv_like done") + return output @@ -262,21 +274,6 @@ def patch_pipeline_model[T](model: T, group: mx.distributed.Group) -> T: if cache is not None: cache[-1].state = mx.depends(cache[-1].state, logits) # type: ignore - # Use send/recv instead of all_gather to get logits to rank 0 - # (JACCL recv picks up all_gather data instead of send data) - if world_size > 1: - if rank == world_size - 1: - mx.eval(logits) - logger.info(f"[patch_pipeline_model] send logits to rank 0, shape={logits.shape}") - sent = mx.distributed.send(logits, dst=0, group=group) - mx.eval(sent) - logger.info(f"[patch_pipeline_model] send done") - elif rank == 0: - logger.info(f"[patch_pipeline_model] recv_like logits from rank {world_size - 1}, shape={logits.shape}") - logits = mx.distributed.recv_like(logits, src=world_size - 1, group=group) - mx.eval(logits) - logger.info(f"[patch_pipeline_model] recv_like done") - return logits cls.__call__ = patched_call