revert
This commit is contained in:
@@ -16,10 +16,7 @@ from mlx.nn.layers.distributed import (
|
||||
from mlx_lm.models.base import (
|
||||
scaled_dot_product_attention, # pyright: ignore[reportUnknownVariableType]
|
||||
)
|
||||
from mlx_lm.models.deepseek_v3 import (
|
||||
DeepseekV3MLP,
|
||||
group_expert_select, # pyright: ignore[reportAttributeAccessIssue,reportUnknownVariableType]
|
||||
)
|
||||
from mlx_lm.models.deepseek_v3 import DeepseekV3MLP
|
||||
from mlx_lm.models.deepseek_v3 import Model as DeepseekV3Model
|
||||
from mlx_lm.models.deepseek_v32 import DeepseekV32MLP
|
||||
from mlx_lm.models.deepseek_v32 import Model as DeepseekV32Model
|
||||
@@ -525,67 +522,6 @@ def _set_layers(model: nn.Module, layers: list[_LayerCallable]) -> None:
|
||||
raise ValueError("Model must have either a 'layers' or 'h' attribute")
|
||||
|
||||
|
||||
def _all_gather_last_axis(x: mx.array, group: mx.distributed.Group) -> mx.array:
|
||||
"""All-gather along the last axis by transposing, gathering on axis 0, transposing back."""
|
||||
batch_shape = x.shape[:-1]
|
||||
d = x.shape[-1]
|
||||
flat = x.reshape(-1, d).T # (D/N, batch_prod)
|
||||
gathered = mx.distributed.all_gather(flat, group=group) # (D, batch_prod)
|
||||
return gathered.T.reshape(*batch_shape, -1) # (..., D)
|
||||
|
||||
|
||||
class AllToAllLinear(nn.Module):
|
||||
"""Column-sharded linear that all_gathers output to reconstruct full result.
|
||||
|
||||
Wraps an AllToShardedLinear (created by shard_linear) and adds all_gather
|
||||
on the output so callers see the full unsharded output.
|
||||
"""
|
||||
|
||||
def __init__(self, sharded_linear: nn.Module, group: mx.distributed.Group):
|
||||
super().__init__()
|
||||
self.inner = sharded_linear
|
||||
self.group = group
|
||||
|
||||
def __call__(self, x: mx.array, *args: object, **kwargs: object) -> mx.array:
|
||||
out = cast(mx.array, self.inner(x, *args, **kwargs))
|
||||
return _all_gather_last_axis(out, self.group)
|
||||
|
||||
|
||||
class ShardedMoEGate(nn.Module):
|
||||
"""MoEGate with column-sharded weight and all_gather for routing scores."""
|
||||
|
||||
def __init__(self, gate: nn.Module, group: mx.distributed.Group):
|
||||
super().__init__()
|
||||
rank = group.rank()
|
||||
n = group.size()
|
||||
n_experts = int(gate.weight.shape[0]) # pyright: ignore[reportUnknownMemberType,reportUnknownArgumentType]
|
||||
shard_size = n_experts // n
|
||||
start = rank * shard_size
|
||||
end = start + shard_size
|
||||
|
||||
self.weight: mx.array = gate.weight[start:end] # pyright: ignore[reportUnknownMemberType]
|
||||
self.e_score_correction_bias: mx.array = gate.e_score_correction_bias # pyright: ignore[reportUnknownMemberType]
|
||||
self.top_k: int = gate.top_k # pyright: ignore[reportUnknownMemberType]
|
||||
self.norm_topk_prob: bool = gate.norm_topk_prob # pyright: ignore[reportUnknownMemberType]
|
||||
self.n_group: int = gate.n_group # pyright: ignore[reportUnknownMemberType]
|
||||
self.topk_group: int = gate.topk_group # pyright: ignore[reportUnknownMemberType]
|
||||
self.routed_scaling_factor: float = gate.routed_scaling_factor # pyright: ignore[reportUnknownMemberType]
|
||||
self.group = group
|
||||
|
||||
def __call__(self, x: mx.array) -> tuple[mx.array, mx.array]:
|
||||
scores_shard = x @ self.weight.T
|
||||
scores = _all_gather_last_axis(scores_shard, self.group)
|
||||
return group_expert_select( # pyright: ignore[reportUnknownVariableType]
|
||||
scores,
|
||||
self.e_score_correction_bias,
|
||||
self.top_k,
|
||||
self.n_group,
|
||||
self.topk_group,
|
||||
self.routed_scaling_factor,
|
||||
self.norm_topk_prob,
|
||||
)
|
||||
|
||||
|
||||
class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
@@ -594,16 +530,13 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
|
||||
on_timeout: TimeoutCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(DeepseekV3Model, model)
|
||||
|
||||
for layer in model.layers:
|
||||
eval_with_timeout(
|
||||
layer.parameters(), timeout_seconds / len(model.layers), on_timeout
|
||||
)
|
||||
|
||||
if layer.self_attn.q_lora_rank is not None:
|
||||
layer.self_attn.q_a_proj = AllToAllLinear( # pyright: ignore[reportAttributeAccessIssue]
|
||||
self.all_to_sharded_linear(layer.self_attn.q_a_proj), self.group
|
||||
)
|
||||
|
||||
# Shard attention heads
|
||||
if layer.self_attn.q_lora_rank is None:
|
||||
layer.self_attn.q_proj = self.all_to_sharded_linear(
|
||||
layer.self_attn.q_proj
|
||||
@@ -613,10 +546,11 @@ 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 = self.sharded_to_all_linear(
|
||||
layer.self_attn.o_proj
|
||||
)
|
||||
layer.self_attn.num_heads //= self.N
|
||||
|
||||
# Logic from upstream mlx
|
||||
num_heads = layer.self_attn.num_heads
|
||||
sh = self.group.rank() * num_heads
|
||||
eh = sh + num_heads
|
||||
@@ -633,14 +567,7 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
|
||||
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
|
||||
|
||||
# Shard the MoE. Shard in place since the MoE should be responsible
|
||||
# for aggregating the results.
|
||||
else:
|
||||
# Shard the gate weight for reduced compute
|
||||
n_experts = int(layer.mlp.gate.weight.shape[0])
|
||||
if n_experts % self.N == 0:
|
||||
layer.mlp.gate = ShardedMoEGate(layer.mlp.gate, self.group) # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
if getattr(layer.mlp, "shared_experts", None) is not None:
|
||||
self.all_to_sharded_linear_in_place(
|
||||
layer.mlp.shared_experts.gate_proj
|
||||
|
||||
Reference in New Issue
Block a user