Compare commits

...
Author SHA1 Message Date
Ryuichi Leo Takashige 325ec6136a Fix TP=2 2026-04-21 07:49:26 +01:00
Ryuichi Leo Takashige 726680b141 Add Kimi K2.6 2026-04-21 07:49:16 +01:00
3 changed files with 596 additions and 71 deletions

No files matched your search

@@ -0,0 +1,32 @@
model_id = "moonshotai/Kimi-K2.6"
n_layers = 61
hidden_size = 7168
num_key_value_heads = 64
supports_tensor = true
tasks = ["TextGeneration"]
family = "kimi"
quantization = ""
base_model = "Kimi K2.6"
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
context_length = 262144
[storage_size]
in_bytes = 595148192736
[vision]
image_token_id = 163605
model_type = "kimi_vl"
processor_repo = "moonshotai/Kimi-K2.6"
# Source: https://huggingface.co/moonshotai/Kimi-K2.6
[sampling_defaults]
temperature = 1.0
top_p = 0.95
min_p = 0.01
# Source: https://huggingface.co/moonshotai/Kimi-K2.6
[sampling_defaults.non_thinking]
temperature = 0.6
top_p = 0.95
min_p = 0.01
+557 -64
View File
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Literal, Protocol, cast
import mlx.core as mx
import mlx.nn as nn
from mlx.nn.layers.distributed import (
ShardedToAllLinear,
shard_inplace,
shard_linear,
sum_gradients,
@@ -58,6 +59,7 @@ from mlx_lm.models.qwen3_vl import Model as Qwen3VLModel
from mlx_lm.models.step3p5 import Model as Step35Model
from mlx_lm.models.step3p5 import Step3p5MLP as Step35MLP
from mlx_lm.models.step3p5 import Step3p5Model as Step35InnerModel
from mlx_lm.models.switch_layers import QuantizedSwitchLinear, SwitchLinear
from exo.shared.types.worker.shards import PipelineShardMetadata
from exo.worker.runner.bootstrap import logger
@@ -65,6 +67,67 @@ from exo.worker.runner.bootstrap import logger
if TYPE_CHECKING:
from mlx_lm.models.cache import Cache
def _fp32_reducing_sharded_to_all_call(
self: ShardedToAllLinear, x: mx.array
) -> mx.array:
weight = cast(mx.array, self["weight"])
group = cast(mx.distributed.Group, self.group)
y = mx.matmul(x, weight.T, output_dtype=mx.float32) # pyright: ignore[reportCallIssue]
y = mx.distributed.all_sum(y, group=group)
if "bias" in self:
y = y + cast(mx.array, self["bias"]).astype(mx.float32)
return y.astype(x.dtype)
ShardedToAllLinear.__call__ = _fp32_reducing_sharded_to_all_call
from mlx.nn.layers.distributed import AllToShardedLinear # noqa: E402
def _splitk_override_for_unsharded(M: int, N_full: int, K: int) -> int:
"""Return the override value that forces a per-rank matmul to match the
unsharded kernel's K-reduction shape.
The updated mlx heuristic returns 0 when splitk wouldn't dispatch for the
unsharded shape; the caller must then set the override to -1 so the
per-rank matmul also skips splitk. If the heuristic returns a positive n,
set the override to n so the per-rank matmul uses the same partition
count as the unsharded call.
"""
n = mx.compute_splitk_partitions(M, N_full, K)
return n if n > 0 else -1
def _splitk_override_all_to_sharded_call(
self: AllToShardedLinear, x: mx.array
) -> mx.array:
"""All-to-sharded matmul that matches the unsharded kernel's K-reduction.
mx.eval is forced inside the override block because the override is
thread-local at kernel-dispatch time; without the eval the matmul is
deferred past the clear and the dispatch falls back to the per-rank
heuristic.
"""
x = sum_gradients(self.group)(x) # pyright: ignore[reportAttributeAccessIssue]
weight = cast(mx.array, self["weight"])
per_rank_N, K = weight.shape
N_full = per_rank_N * self.group.size() # pyright: ignore[reportAttributeAccessIssue]
M = x.shape[-2] if x.ndim >= 2 else 1
mx.set_splitk_partitions_override(_splitk_override_for_unsharded(M, N_full, K))
try:
if "bias" in self:
y = mx.addmm(cast(mx.array, self["bias"]), x, weight.T)
else:
y = mx.matmul(x, weight.T)
finally:
mx.set_splitk_partitions_override(0)
return y
AllToShardedLinear.__call__ = _splitk_override_all_to_sharded_call
LayerLoadedCallback = Callable[[int, int], None] # (layers_loaded, total_layers)
@@ -637,13 +700,13 @@ class LlamaShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_proj = self.all_to_sharded_linear(layer.self_attn.q_proj)
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)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.n_heads //= self.N
if layer.self_attn.n_kv_heads is not None:
layer.self_attn.n_kv_heads //= self.N
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
mx.eval(layer)
if on_layer_loaded is not None:
@@ -699,7 +762,7 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_b_proj
)
layer.self_attn.o_proj = self.sharded_to_all_linear(layer.self_attn.o_proj)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.num_heads //= self.N
# Logic from upstream mlx
@@ -716,24 +779,24 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
# Shard the MLP
if isinstance(layer.mlp, (DeepseekV3MLP, DeepseekV32MLP)):
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
# Shard the MoE.
# Shard the MoE with column-sharded down_proj for bit-exactness.
else:
if getattr(layer.mlp, "shared_experts", None) is not None:
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.gate_proj
)
self.sharded_to_all_linear_in_place(
layer.mlp.shared_experts.down_proj
)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.up_proj
)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.down_proj
)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.gate_proj)
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.up_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.down_proj)
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
layer.mlp.sharding_group = self.group
@@ -744,19 +807,440 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
return model
class NShardedLinear(nn.Module):
def __init__(self, in_dims: int, out_dims: int, bias: bool, group: mx.distributed.Group):
super().__init__()
N = group.size()
self.group = group
self.weight = mx.zeros((out_dims // N, in_dims))
if bias:
self.bias = mx.zeros((out_dims // N,))
def __call__(self, x: mx.array) -> mx.array:
x_full = _all_gather_last(x, self.group)
weight = cast(mx.array, self["weight"])
M = x_full.shape[-2] if x_full.ndim >= 2 else 1
per_rank_N, K = weight.shape
N_full = per_rank_N * self.group.size()
mx.set_splitk_partitions_override(_splitk_override_for_unsharded(M, N_full, K))
try:
if "bias" in self:
y = mx.addmm(cast(mx.array, self["bias"]), x_full, weight.T)
else:
y = mx.matmul(x_full, weight.T)
finally:
mx.set_splitk_partitions_override(0)
return _all_gather_last(y, self.group)
@classmethod
def from_linear(cls, linear: nn.Linear, group: mx.distributed.Group) -> "NShardedLinear":
out_dims, in_dims = linear.weight.shape # pyright: ignore[reportAttributeAccessIssue]
N = group.size()
rank = group.rank()
per_rank = out_dims // N
instance = cls(in_dims, out_dims, hasattr(linear, "bias"), group)
new_weight = cast(mx.array, linear["weight"])[rank * per_rank : (rank + 1) * per_rank]
instance.update({"weight": new_weight})
if hasattr(linear, "bias"):
new_bias = cast(mx.array, linear["bias"])[rank * per_rank : (rank + 1) * per_rank]
instance.update({"bias": new_bias})
return instance
def _all_gather_last(x: mx.array, group: mx.distributed.Group) -> mx.array:
"""all_gather over the last axis.
``mx.distributed.all_gather`` concatenates along axis 0 and its internal
``ensure_row_contiguous`` forces a strided memcpy if the input view isn't
contiguous. Naively wrapping with ``moveaxis`` produces two strided full-
tensor memcpys per call (inside all_gather + at the next consumer). We
sidestep that by (1) flattening all leading axes so the tensor becomes 2D
``(prefix, last_shard)``, (2) forcing a single contiguous transpose to
``(last_shard, prefix)``, (3) running the contiguous all_gather (no
internal copy), and (4) transposing+reshaping back to
``(*leading, last_shard * N)``. One explicit memcpy instead of two
strided ones.
"""
leading = x.shape[:-1]
last = x.shape[-1]
x2 = x.reshape(-1, last)
xt = mx.contiguous(x2.T)
g = mx.distributed.all_gather(xt, group=group)
return mx.contiguous(g.T).reshape(*leading, last * group.size())
class ShardedInputNorm(CustomMlxLayer):
def __init__(self, norm: _LayerCallable, group: mx.distributed.Group):
super().__init__(norm)
self.group = group
def __call__(self, x: mx.array) -> mx.array:
return cast(mx.array, self.original_layer(_all_gather_last(x, self.group)))
class ShardedEmbedding(CustomMlxLayer):
def __init__(self, embed: _LayerCallable, group: mx.distributed.Group):
super().__init__(embed)
self.group = group
def __call__(self, ids: mx.array) -> mx.array:
y = cast(mx.array, self.original_layer(ids))
N = self.group.size()
per_rank = y.shape[-1] // N
rank = self.group.rank()
return y[..., rank * per_rank : (rank + 1) * per_rank]
def _wrap_block_entry_norms(layer: nn.Module, group: mx.distributed.Group) -> None:
children = layer.children() if hasattr(layer, "children") else {}
to_wrap: list[str] = []
for name, child in (children.items() if isinstance(children, dict) else children): # pyright: ignore[reportGeneralTypeIssues]
if isinstance(child, (nn.RMSNorm, nn.LayerNorm)):
to_wrap.append(name)
for name in to_wrap:
orig = getattr(layer, name)
setattr(layer, name, ShardedInputNorm(orig, group))
def _switch_mlp_activation(switch_mlp: object, x_gate: mx.array, x_up: mx.array) -> mx.array:
activation = getattr(switch_mlp, "activation", None)
if activation is not None:
return cast(mx.array, activation(x_up, x_gate))
return nn.silu(x_gate) * x_up
def _switch_mlp_n_sharded_sharded_out(
switch_mlp: object,
x: mx.array,
indices: mx.array,
scores: mx.array,
group: mx.distributed.Group,
) -> mx.array:
"""Run a SwitchGLU-shaped MoE expert block with column-sharded down_proj
and return the output STILL SHARDED on the hidden dim (H/N per rank).
gate_proj/up_proj are all-to-sharded (output dim split on expert_hidden).
down_proj is all-to-sharded (output dim split on model_dim). The
intermediate is all_gathered to full before down_proj so each rank runs
the full K reduction over intermediate_full. The per-rank down_proj
output is weighted by the routing scores and summed across the top-k
dimension, but NOT gathered to full — the caller is expected to combine
sharded outputs and issue a single final all_gather.
"""
from mlx_lm.models.switch_layers import _gather_sort, _scatter_unsort
x_exp = mx.expand_dims(x, (-2, -3))
do_sort = indices.size >= 64
idx = indices
inv_order = None
if do_sort:
x_exp, idx, inv_order = _gather_sort(x_exp, indices)
gp = switch_mlp.gate_proj # pyright: ignore[reportAttributeAccessIssue]
up = switch_mlp.up_proj # pyright: ignore[reportAttributeAccessIssue]
dp = switch_mlp.down_proj # pyright: ignore[reportAttributeAccessIssue]
x_up = mx.gather_mm(
x_exp,
cast(mx.array, up["weight"]).swapaxes(-1, -2),
rhs_indices=idx,
sorted_indices=do_sort,
)
if "bias" in up:
x_up = x_up + mx.expand_dims(cast(mx.array, up["bias"])[idx], -2)
x_gate = mx.gather_mm(
x_exp,
cast(mx.array, gp["weight"]).swapaxes(-1, -2),
rhs_indices=idx,
sorted_indices=do_sort,
)
if "bias" in gp:
x_gate = x_gate + mx.expand_dims(cast(mx.array, gp["bias"])[idx], -2)
hidden_shard = _switch_mlp_activation(switch_mlp, x_gate, x_up)
hidden_full = _all_gather_last(hidden_shard, group)
out_shard = mx.gather_mm(
hidden_full,
cast(mx.array, dp["weight"]).swapaxes(-1, -2),
rhs_indices=idx,
sorted_indices=do_sort,
)
if "bias" in dp:
out_shard = out_shard + mx.expand_dims(cast(mx.array, dp["bias"])[idx], -2)
if do_sort:
out_shard = _scatter_unsort(out_shard, inv_order, indices.shape)
out_shard = out_shard.squeeze(-2)
return (out_shard * scores[..., None]).sum(axis=-2).astype(out_shard.dtype)
def _switch_mlp_n_sharded(
switch_mlp: object,
x: mx.array,
indices: mx.array,
scores: mx.array,
group: mx.distributed.Group,
) -> mx.array:
out_shard = _switch_mlp_n_sharded_sharded_out(
switch_mlp, x, indices, scores, group
)
return _all_gather_last(out_shard, group)
def _switch_fc_n_sharded(
switch_mlp: object,
x: mx.array,
indices: mx.array,
group: mx.distributed.Group,
activation: Callable[[mx.array], mx.array],
) -> mx.array:
"""SwitchMLP variant (NemotronH): fc1 -> activation -> fc2, both SwitchLinears.
Both fc1 and fc2 are column-sharded (output dim). fc1 output is
all_gathered before fc2, and fc2 output is all_gathered afterward.
"""
from mlx_lm.models.switch_layers import _gather_sort, _scatter_unsort
x_exp = mx.expand_dims(x, (-2, -3))
do_sort = indices.size >= 64
idx = indices
inv_order = None
if do_sort:
x_exp, idx, inv_order = _gather_sort(x_exp, indices)
fc1 = switch_mlp.fc1 # pyright: ignore[reportAttributeAccessIssue]
fc2 = switch_mlp.fc2 # pyright: ignore[reportAttributeAccessIssue]
h_shard = mx.gather_mm(
x_exp,
cast(mx.array, fc1["weight"]).swapaxes(-1, -2),
rhs_indices=idx,
sorted_indices=do_sort,
)
h_shard = activation(h_shard)
h_full = _all_gather_last(h_shard, group)
out_shard = mx.gather_mm(
h_full,
cast(mx.array, fc2["weight"]).swapaxes(-1, -2),
rhs_indices=idx,
sorted_indices=do_sort,
)
out_full = _all_gather_last(out_shard, group)
if do_sort:
out_full = _scatter_unsort(out_full, inv_order, indices.shape)
return out_full.squeeze(-2)
def _matmul_with_unsharded_splitk(
x: mx.array, weight: mx.array, per_rank_N: int, K: int, group: mx.distributed.Group
) -> mx.array:
"""Per-rank matmul with splitk override forced to the unsharded count."""
N_full = per_rank_N * group.size()
M = x.shape[-2] if x.ndim >= 2 else 1
mx.set_splitk_partitions_override(_splitk_override_for_unsharded(M, N_full, K))
try:
y = mx.matmul(x, weight.T)
finally:
mx.set_splitk_partitions_override(0)
return y
def _mlp_n_sharded_sharded_out(
mlp: object, x: mx.array, group: mx.distributed.Group
) -> mx.array:
"""Like _mlp_n_sharded but returns the output still sharded on H/N."""
up = mlp.up_proj # pyright: ignore[reportAttributeAccessIssue]
dp = mlp.down_proj # pyright: ignore[reportAttributeAccessIssue]
up_w = cast(mx.array, up["weight"])
dp_w = cast(mx.array, dp["weight"])
x_up = _matmul_with_unsharded_splitk(x, up_w, up_w.shape[0], up_w.shape[1], group)
if hasattr(mlp, "gate_proj"):
gp = mlp.gate_proj # pyright: ignore[reportAttributeAccessIssue]
gp_w = cast(mx.array, gp["weight"])
x_gate = _matmul_with_unsharded_splitk(
x, gp_w, gp_w.shape[0], gp_w.shape[1], group
)
hidden_shard = nn.silu(x_gate) * x_up
else:
hidden_shard = nn.relu2(x_up)
hidden_full = _all_gather_last(hidden_shard, group)
return _matmul_with_unsharded_splitk(
hidden_full, dp_w, dp_w.shape[0], dp_w.shape[1], group
)
def _mlp_n_sharded(mlp: object, x: mx.array, group: mx.distributed.Group) -> mx.array:
"""SwiGLU MLP (gate_proj/up_proj/down_proj) with N-sharded down_proj.
Falls back to up_proj + activation + down_proj for MLPs that don't have a
gate_proj (e.g. NemotronHMLP). The sharded linears here were sharded
via ``shard_inplace`` so they are still plain ``nn.Linear`` instances;
we run their matmuls directly with the splitk override so each per-rank
matmul produces a bf16 output that's bit-exact per column to what the
unsharded kernel would produce.
"""
out_shard = _mlp_n_sharded_sharded_out(mlp, x, group)
return _all_gather_last(out_shard, group)
class ShardedMoE(CustomMlxLayer):
"""Wraps any MoE layer with distributed sum_gradients / all_sum."""
"""Wraps a MoE block to use N-sharded (column) down_proj + all_gather.
Each sharded down_proj output element is bit-exact per column to the
unsharded tp=1 kernel's output (every (m, n) goes through the full K-fold
the unsharded kernel uses). all_gather is a pure bf16 byte-shuffle with
no rounding introduced. The MoE's internal routing (gate + softmax +
argpartition) is unchanged.
Dispatches by inspecting the wrapped MoE's attributes to handle the
different MoE block APIs across model families.
"""
def __init__(self, layer: _LayerCallable):
super().__init__(layer)
self.sharding_group: mx.distributed.Group | None = None
def __call__(self, x: mx.array) -> mx.array:
if self.sharding_group is not None:
x = sum_gradients(self.sharding_group)(x)
y = self.original_layer.__call__(x)
if self.sharding_group is not None:
y = mx.distributed.all_sum(y, group=self.sharding_group)
if self.sharding_group is None:
return cast(mx.array, self.original_layer.__call__(x))
moe = self.original_layer
if hasattr(moe, "switch_mlp") and hasattr(moe, "shared_expert") and hasattr(
moe, "shared_expert_gate"
):
return self._qwen_style(x)
if hasattr(moe, "switch_mlp") and hasattr(moe.switch_mlp, "fc1"):
return self._nemotron_h_style(x)
if hasattr(moe, "switch_mlp") and hasattr(moe, "shared_experts"):
return self._deepseek_style(x)
if hasattr(moe, "switch_mlp") and hasattr(moe, "e_score_correction_bias"):
return self._minimax_style(x)
if hasattr(moe, "switch_mlp") and hasattr(moe, "share_expert"):
return self._nemotronh_style(x)
if hasattr(moe, "switch_mlp"):
return self._generic_switch_mlp_style(x)
if hasattr(moe, "experts") and hasattr(moe, "router"):
return self._gpt_oss_style(x)
x = sum_gradients(self.sharding_group)(x)
return cast(mx.array, moe.__call__(x))
def _route_softmax_topk(self, moe: object, x: mx.array) -> tuple[mx.array, mx.array]:
gates = moe.gate(x) # pyright: ignore[reportAttributeAccessIssue]
gates = mx.softmax(gates, axis=-1, precise=True)
k = getattr(moe, "top_k", None) or moe.num_experts_per_tok # pyright: ignore[reportAttributeAccessIssue]
inds = mx.argpartition(gates, kth=-k, axis=-1)[..., -k:]
scores = mx.take_along_axis(gates, inds, axis=-1)
if getattr(moe, "norm_topk_prob", False):
scores = scores / scores.sum(axis=-1, keepdims=True)
return inds, scores
def _qwen_style(self, x: mx.array) -> mx.array:
assert self.sharding_group is not None
x = sum_gradients(self.sharding_group)(x)
moe = self.original_layer
inds, scores = self._route_softmax_topk(moe, x)
y_shard = _switch_mlp_n_sharded_sharded_out(
moe.switch_mlp, x, inds, scores, self.sharding_group # pyright: ignore[reportAttributeAccessIssue]
)
shared_shard = _mlp_n_sharded_sharded_out(moe.shared_expert, x, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
shared_shard = mx.sigmoid(moe.shared_expert_gate(x)) * shared_shard # pyright: ignore[reportAttributeAccessIssue]
return _all_gather_last(y_shard + shared_shard, self.sharding_group)
def _deepseek_style(self, x: mx.array) -> mx.array:
assert self.sharding_group is not None
x = sum_gradients(self.sharding_group)(x)
moe = self.original_layer
inds, scores = moe.gate(x) # pyright: ignore[reportAttributeAccessIssue]
y_shard = _switch_mlp_n_sharded_sharded_out(
moe.switch_mlp, x, inds, scores, self.sharding_group # pyright: ignore[reportAttributeAccessIssue]
)
if getattr(moe.config, "n_shared_experts", None) is not None: # pyright: ignore[reportAttributeAccessIssue]
y_shard = y_shard + _mlp_n_sharded_sharded_out(
moe.shared_experts, x, self.sharding_group # pyright: ignore[reportAttributeAccessIssue]
)
return _all_gather_last(y_shard, self.sharding_group)
def _minimax_style(self, x: mx.array) -> mx.array:
assert self.sharding_group is not None
x = sum_gradients(self.sharding_group)(x)
moe = self.original_layer
gates = moe.gate(x.astype(mx.float32)) # pyright: ignore[reportAttributeAccessIssue]
scores = mx.sigmoid(gates)
orig_scores = scores
scores = scores + moe.e_score_correction_bias # pyright: ignore[reportAttributeAccessIssue]
k = moe.num_experts_per_tok # pyright: ignore[reportAttributeAccessIssue]
inds = mx.argpartition(-scores, kth=k - 1, axis=-1)[..., :k]
scores = mx.take_along_axis(orig_scores, inds, axis=-1)
scores = scores / (mx.sum(scores, axis=-1, keepdims=True) + 1e-20)
scores = scores.astype(x.dtype)
y = _switch_mlp_n_sharded(moe.switch_mlp, x, inds, scores, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
return y
def _nemotronh_style(self, x: mx.array) -> mx.array:
"""Handles both NemotronH and Step35 (gate returns indices+weights directly)."""
assert self.sharding_group is not None
x = sum_gradients(self.sharding_group)(x)
moe = self.original_layer
gate_out = moe.gate(x) # pyright: ignore[reportAttributeAccessIssue]
if isinstance(gate_out, tuple):
inds, scores = gate_out
else:
inds, scores = self._route_softmax_topk(moe, x)
y = _switch_mlp_n_sharded(moe.switch_mlp, x, inds, scores, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
if getattr(moe, "share_expert", None) is not None:
y = y + _mlp_n_sharded(moe.share_expert, x, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
return y
def _nemotron_h_style(self, x: mx.array) -> mx.array:
assert self.sharding_group is not None
moe = self.original_layer
residuals = x
inds, scores = moe.gate(x) # pyright: ignore[reportAttributeAccessIssue]
if moe.moe_latent_size is not None: # pyright: ignore[reportAttributeAccessIssue]
x = moe.fc1_latent_proj(x) # pyright: ignore[reportAttributeAccessIssue]
y = _switch_fc_n_sharded(
moe.switch_mlp, # pyright: ignore[reportAttributeAccessIssue]
x,
inds,
self.sharding_group,
moe.switch_mlp.activation, # pyright: ignore[reportAttributeAccessIssue]
)
y = (y * scores[..., None]).sum(axis=-2).astype(y.dtype)
if moe.moe_latent_size is not None: # pyright: ignore[reportAttributeAccessIssue]
y = moe.fc2_latent_proj(y) # pyright: ignore[reportAttributeAccessIssue]
if moe.config.n_shared_experts is not None: # pyright: ignore[reportAttributeAccessIssue]
y = y + _mlp_n_sharded(moe.shared_experts, residuals, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
return y
def _gpt_oss_style(self, x: mx.array) -> mx.array:
from mlx_lm.models.gpt_oss import mlx_topk # pyright: ignore[reportUnknownVariableType]
assert self.sharding_group is not None
x = sum_gradients(self.sharding_group)(x)
moe = self.original_layer
g = moe.router(x) # pyright: ignore[reportAttributeAccessIssue]
expert_weights, indices = mlx_topk(g, k=moe.num_experts_per_tok, axis=-1) # pyright: ignore[reportAttributeAccessIssue]
expert_weights = mx.softmax(expert_weights, axis=-1, precise=True)
y = _switch_mlp_n_sharded(moe.experts, x, indices, expert_weights, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
return y
def _generic_switch_mlp_style(self, x: mx.array) -> mx.array:
assert self.sharding_group is not None
x = sum_gradients(self.sharding_group)(x)
moe = self.original_layer
gate_out = moe.gate(x) # pyright: ignore[reportAttributeAccessIssue]
if isinstance(gate_out, tuple):
inds, scores = gate_out
else:
inds, scores = self._route_softmax_topk(moe, x)
y = _switch_mlp_n_sharded(moe.switch_mlp, x, inds, scores, self.sharding_group) # pyright: ignore[reportAttributeAccessIssue]
return y
@@ -780,7 +1264,7 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_b_proj
)
layer.self_attn.o_proj = self.sharded_to_all_linear(layer.self_attn.o_proj)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.num_heads //= self.N
# Logic from upstream mlx
@@ -796,7 +1280,7 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
if isinstance(layer.mlp, Glm4MoeLiteMLP):
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
else:
@@ -804,15 +1288,15 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.gate_proj
)
self.sharded_to_all_linear_in_place(
layer.mlp.shared_experts.down_proj
)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.up_proj
)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.down_proj
)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.gate_proj)
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.up_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.down_proj)
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
layer.mlp.sharding_group = self.group # type: ignore
mx.eval(layer)
@@ -914,7 +1398,7 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_proj = self.all_to_sharded_linear(layer.self_attn.q_proj)
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)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.num_attention_heads //= self.N
layer.self_attn.num_key_value_heads //= self.N
@@ -925,12 +1409,12 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
self.all_to_sharded_linear_in_place(
layer.block_sparse_moe.switch_mlp.gate_proj
)
self.sharded_to_all_linear_in_place(
layer.block_sparse_moe.switch_mlp.down_proj
)
self.all_to_sharded_linear_in_place(
layer.block_sparse_moe.switch_mlp.up_proj
)
self.all_to_sharded_linear_in_place(
layer.block_sparse_moe.switch_mlp.down_proj
)
layer.block_sparse_moe = ShardedMoE(layer.block_sparse_moe) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
layer.block_sparse_moe.sharding_group = self.group # pyright: ignore[reportAttributeAccessIssue]
mx.eval(layer)
@@ -968,8 +1452,8 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
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
layer.self_attn.o_proj = NShardedLinear.from_linear(
layer.self_attn.o_proj, self.group
)
layer.self_attn.n_heads //= self.N
layer.self_attn.n_kv_heads //= self.N
@@ -1007,8 +1491,8 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
linear_attn.in_proj_a = self.all_to_sharded_linear(
linear_attn.in_proj_a
)
linear_attn.out_proj = self.sharded_to_all_linear(
linear_attn.out_proj
linear_attn.out_proj = NShardedLinear.from_linear(
linear_attn.out_proj, self.group
)
# Shard conv1d: depthwise conv with non-contiguous channel slicing.
@@ -1061,13 +1545,17 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
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
layer.self_attn.o_proj = NShardedLinear.from_linear(
layer.self_attn.o_proj, self.group
)
layer.self_attn.num_attention_heads //= self.N
layer.self_attn.num_key_value_heads //= self.N
# Shard the MoE.
# Shard the MoE. Down_proj is column-sharded (output dim) so each
# per-rank matmul runs the full K-reduction and its bf16 output is
# bit-exact per column to tp=1. ShardedMoE.__call__ inserts the
# required all_gather of the intermediate before down_proj and of
# the output after.
if isinstance(
layer.mlp,
(
@@ -1077,25 +1565,25 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
),
):
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.gate_proj)
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.up_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.down_proj)
if isinstance(
layer.mlp, (Qwen3NextSparseMoeBlock, Qwen3_5SparseMoeBlock)
):
self.all_to_sharded_linear_in_place(
layer.mlp.shared_expert.gate_proj
)
self.sharded_to_all_linear_in_place(
self.all_to_sharded_linear_in_place(layer.mlp.shared_expert.up_proj)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_expert.down_proj
)
self.all_to_sharded_linear_in_place(layer.mlp.shared_expert.up_proj)
layer.mlp = ShardedMoE(layer.mlp) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
layer.mlp.sharding_group = self.group
# Shard the MLP
else:
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
mx.eval(layer)
@@ -1118,30 +1606,30 @@ class Glm4MoeShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_proj = self.all_to_sharded_linear(layer.self_attn.q_proj)
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)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.n_heads //= self.N
layer.self_attn.n_kv_heads //= self.N
if isinstance(layer.mlp, MoE):
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.gate_proj)
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.up_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.down_proj)
if getattr(layer.mlp, "shared_experts", None) is not None:
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.gate_proj
)
self.sharded_to_all_linear_in_place(
layer.mlp.shared_experts.down_proj
)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.up_proj
)
self.all_to_sharded_linear_in_place(
layer.mlp.shared_experts.down_proj
)
layer.mlp = ShardedMoE(layer.mlp) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
layer.mlp.sharding_group = self.group
else:
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
mx.eval(layer)
@@ -1164,7 +1652,7 @@ class GptOssShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_proj = self.all_to_sharded_linear(layer.self_attn.q_proj)
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)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.num_attention_heads //= self.N
layer.self_attn.num_key_value_heads //= self.N
@@ -1180,8 +1668,8 @@ class GptOssShardingStrategy(TensorParallelShardingStrategy):
]
self.all_to_sharded_linear_in_place(layer.mlp.experts.gate_proj)
self.sharded_to_all_linear_in_place(layer.mlp.experts.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.experts.up_proj)
self.all_to_sharded_linear_in_place(layer.mlp.experts.down_proj)
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
layer.mlp.sharding_group = self.group # pyright: ignore[reportAttributeAccessIssue]
@@ -1205,7 +1693,7 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
layer.self_attn.q_proj = self.all_to_sharded_linear(layer.self_attn.q_proj)
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)
layer.self_attn.o_proj = NShardedLinear.from_linear(layer.self_attn.o_proj, self.group)
layer.self_attn.num_heads //= self.N
layer.self_attn.num_kv_heads //= self.N
@@ -1218,15 +1706,16 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
if isinstance(layer.mlp, Step35MLP):
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
else:
layer.mlp.sharding_group = self.group
self.all_to_sharded_linear_in_place(layer.mlp.share_expert.gate_proj)
self.all_to_sharded_linear_in_place(layer.mlp.share_expert.up_proj)
self.sharded_to_all_linear_in_place(layer.mlp.share_expert.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.share_expert.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.gate_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.up_proj)
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.down_proj)
layer.mlp = ShardedMoE(layer.mlp) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
layer.mlp.sharding_group = self.group
mx.eval(layer)
if on_layer_loaded is not None:
@@ -1252,7 +1741,7 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
mixer.q_proj = self.all_to_sharded_linear(mixer.q_proj)
mixer.k_proj = self.all_to_sharded_linear(mixer.k_proj)
mixer.v_proj = self.all_to_sharded_linear(mixer.v_proj)
mixer.o_proj = self.sharded_to_all_linear(mixer.o_proj)
mixer.o_proj = NShardedLinear.from_linear(mixer.o_proj, self.group)
mixer.num_heads //= self.N
mixer.num_key_value_heads //= self.N
@@ -1260,13 +1749,16 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
self._shard_mamba2_mixer(mixer, rank)
elif isinstance(mixer, NemotronHMoE):
# Shard routed experts (SwitchMLP uses fc1/fc2)
# N-shard both fc1 and fc2 so each per-rank matmul runs the
# full K reduction (bit-exact per column). ShardedMoE does the
# all_gather of the intermediate between fc1 and fc2 and of the
# fc2 output.
self.all_to_sharded_linear_in_place(mixer.switch_mlp.fc1)
self.sharded_to_all_linear_in_place(mixer.switch_mlp.fc2)
# Shard shared expert in-place (no all-reduce — ShardedMoE handles that)
self.all_to_sharded_linear_in_place(mixer.switch_mlp.fc2)
if hasattr(mixer, "shared_experts"):
self.all_to_sharded_linear_in_place(mixer.shared_experts.gate_proj)
self.all_to_sharded_linear_in_place(mixer.shared_experts.up_proj)
self.sharded_to_all_linear_in_place(mixer.shared_experts.down_proj)
self.all_to_sharded_linear_in_place(mixer.shared_experts.down_proj)
mixer = ShardedMoE(mixer) # pyright: ignore[reportArgumentType]
mixer.sharding_group = self.group
layer.mixer = mixer # pyright: ignore[reportAttributeAccessIssue]
@@ -1320,7 +1812,7 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
mixer.in_proj.weight = mixer.in_proj.weight[indices]
# === out_proj: input is intermediate_size (sharded) → hidden_size (reduce) ===
mixer.out_proj = self.sharded_to_all_linear(mixer.out_proj)
mixer.out_proj = NShardedLinear.from_linear(mixer.out_proj, self.group)
# === conv1d: depthwise conv on conv_dim channels ===
# conv_dim layout: [ssm_hidden:IS | B:NG*SS | C:NG*SS]
@@ -1368,12 +1860,13 @@ class WrappedGemma4Experts(CustomMlxLayer):
def __call__(
self, x: mx.array, top_k_indices: mx.array, top_k_weights: mx.array
) -> mx.array:
if self.sharding_group is not None:
x = sum_gradients(self.sharding_group)(x)
y: mx.array = self.original_layer(x, top_k_indices, top_k_weights)
if self.sharding_group is not None:
y = mx.distributed.all_sum(y, group=self.sharding_group)
return y
if self.sharding_group is None:
return cast(mx.array, self.original_layer(x, top_k_indices, top_k_weights))
x = sum_gradients(self.sharding_group)(x)
switch_glu = self.original_layer.switch_glu # pyright: ignore[reportAttributeAccessIssue]
return _switch_mlp_n_sharded(
switch_glu, x, top_k_indices, top_k_weights, self.sharding_group
)
class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
@@ -1393,18 +1886,18 @@ class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
attn.k_proj = self.all_to_sharded_linear(attn.k_proj)
if not attn.use_k_eq_v:
attn.v_proj = self.all_to_sharded_linear(attn.v_proj)
attn.o_proj = self.sharded_to_all_linear(attn.o_proj)
attn.o_proj = NShardedLinear.from_linear(attn.o_proj, self.group)
attn.n_heads //= self.N
attn.n_kv_heads //= self.N
layer.mlp.gate_proj = self.all_to_sharded_linear(layer.mlp.gate_proj)
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
layer.mlp.down_proj = NShardedLinear.from_linear(layer.mlp.down_proj, self.group)
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
if layer.enable_moe:
self.all_to_sharded_linear_in_place(layer.experts.switch_glu.gate_proj)
self.sharded_to_all_linear_in_place(layer.experts.switch_glu.down_proj)
self.all_to_sharded_linear_in_place(layer.experts.switch_glu.up_proj)
self.all_to_sharded_linear_in_place(layer.experts.switch_glu.down_proj)
layer.experts = WrappedGemma4Experts(layer.experts) # pyright: ignore[reportAttributeAccessIssue,reportArgumentType]
layer.experts.sharding_group = self.group
Generated
+7 -7
View File
@@ -485,7 +485,7 @@ dependencies = [
{ name = "hypercorn", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "loguru", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mflux", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260420+553a7adb", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#553a7adbb20ed1b71fe643f4075982a639aaef1b" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260421+9ed8c741", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#9ed8c7411f60b4f97e128031d19380255a67e7b0" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-lm", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-vlm", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "msgspec", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -1437,7 +1437,7 @@ dependencies = [
{ name = "hf-transfer", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "huggingface-hub", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "matplotlib", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260420+553a7adb", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#553a7adbb20ed1b71fe643f4075982a639aaef1b" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260421+9ed8c741", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#9ed8c7411f60b4f97e128031d19380255a67e7b0" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "numpy", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "opencv-python", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "piexif", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -1493,8 +1493,8 @@ wheels = [
[[package]]
name = "mlx"
version = "0.31.2.dev20260420+553a7adb"
source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#553a7adbb20ed1b71fe643f4075982a639aaef1b" }
version = "0.31.2.dev20260421+9ed8c741"
source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#9ed8c7411f60b4f97e128031d19380255a67e7b0" }
resolution-markers = [
"python_full_version >= '3.14' and sys_platform == 'darwin'",
"python_full_version < '3.14' and sys_platform == 'darwin'",
@@ -1542,10 +1542,10 @@ wheels = [
[[package]]
name = "mlx-lm"
version = "0.31.3"
source = { git = "https://github.com/rltakashige/mlx-lm?branch=leo%2Ffix-arrayscache-leak#f5b9c9f42ef7577c73c7f6eeeb15b35a6682ff57" }
source = { git = "https://github.com/rltakashige/mlx-lm?branch=leo%2Ffix-arrayscache-leak#a6acf2bf7c850b5ea0f1c78d55d1c4d70f112f94" }
dependencies = [
{ name = "jinja2", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260420+553a7adb", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#553a7adbb20ed1b71fe643f4075982a639aaef1b" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260421+9ed8c741", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#9ed8c7411f60b4f97e128031d19380255a67e7b0" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "numpy", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "protobuf", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pyyaml", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -1562,7 +1562,7 @@ dependencies = [
{ name = "fastapi", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "miniaudio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260420+553a7adb", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#553a7adbb20ed1b71fe643f4075982a639aaef1b" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.2.dev20260421+9ed8c741", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#9ed8c7411f60b4f97e128031d19380255a67e7b0" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-lm", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "numpy", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "opencv-python", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },