From c55cbf6739664bf1d3d5b0cd652f9b36638997dd Mon Sep 17 00:00:00 2001 From: rltakashige Date: Tue, 27 Jan 2026 15:29:06 +0000 Subject: [PATCH] Add mlx lm style tensor sharding for Minimax (#1299) ## Motivation Broken right now. We'll potentially add a better one later ## Changes ## Why It Works ## Test Plan ### Manual Testing Used for evals without any issue. ### Automated Testing --- src/exo/worker/engines/mlx/auto_parallel.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index f2ecb698..89eb4185 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -622,6 +622,7 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy): on_timeout: TimeoutCallback | None, ) -> nn.Module: model = cast(MiniMaxModel, model) + rank = self.group.rank() for layer in model.layers: eval_with_timeout( layer.parameters(), timeout_seconds / len(model.layers), on_timeout @@ -631,6 +632,16 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy): layer.self_attn.k_proj = self.all_to_sharded_linear(layer.self_attn.k_proj) layer.self_attn.v_proj = self.all_to_sharded_linear(layer.self_attn.v_proj) layer.self_attn.o_proj = self.sharded_to_all_linear(layer.self_attn.o_proj) + + # Shard qk_norm weights if present (must match sharded head count) + if getattr(layer.self_attn, "use_qk_norm", False): + layer.self_attn.q_norm.weight = layer.self_attn.q_norm.weight.split( # type: ignore + self.N, axis=-1 + )[rank] + layer.self_attn.k_norm.weight = layer.self_attn.k_norm.weight.split( # type: ignore + self.N, axis=-1 + )[rank] + layer.self_attn.num_attention_heads //= self.N layer.self_attn.num_key_value_heads //= self.N