diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 69ea2984..5347dee2 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -294,7 +294,9 @@ def patch_pipeline_model[T](model: T, group: mx.distributed.Group) -> T: # 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 + last = cache[-1] # type: ignore + dep_cache = last[0] if hasattr(last, "caches") else last # type: ignore + dep_cache.keys = mx.depends(dep_cache.keys, logits) # type: ignore return logits @@ -320,7 +322,9 @@ def patch_tensor_model[T](model: T) -> T: # Add dependency to last cache entry to ensure distributed ops are evaluated if cache is not None and len(cache) > 0: # pyright: ignore[reportAny] - cache[-1].state = mx.depends(cache[-1].state, logits) # pyright: ignore[reportAny,reportUnknownMemberType] + last = cache[-1] # pyright: ignore[reportAny] + dep_cache = last[0] if hasattr(last, "caches") else last # pyright: ignore[reportAny] + dep_cache.keys = mx.depends(dep_cache.keys, logits) # pyright: ignore[reportAny,reportUnknownMemberType] return logits