Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 6 additions & 27 deletions mlx_lm/models/deepseek_v32.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention
from .cache import CacheList, KVCache
from .mla import MultiLinear
from .pipeline import PipelineMixin
from .rope_utils import initialize_rope
from .switch_layers import SwitchGLU

Expand Down Expand Up @@ -412,7 +413,7 @@ def __call__(
return h + r


class DeepseekV32Model(nn.Module):
class DeepseekV32Model(PipelineMixin, nn.Module):
def __init__(self, config: ModelArgs):
super().__init__()
self.vocab_size = config.vocab_size
Expand All @@ -421,28 +422,7 @@ def __init__(self, config: ModelArgs):
DeepseekV32DecoderLayer(config, idx)
for idx in range(config.num_hidden_layers)
]
self.start_idx = 0
self.end_idx = len(self.layers)
self.num_layers = self.end_idx

self.norm = nn.RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.pipeline_rank = 0
self.pipeline_size = 1

def pipeline(self, group):
# Split layers in reverse so rank=0 gets the last layers and
# rank=pipeline_size-1 gets the first
self.pipeline_rank = group.rank()
self.pipeline_size = group.size()
layers_per_rank = len(self.layers) // self.pipeline_size
extra = len(self.layers) - layers_per_rank * self.pipeline_size
if self.pipeline_rank < extra:
layers_per_rank += 1
self.start_idx = (self.pipeline_size - self.pipeline_rank - 1) * layers_per_rank
self.end_idx = self.start_idx + layers_per_rank
self.layers = self.layers[: self.end_idx]
self.layers[: self.start_idx] = [None] * self.start_idx
self.num_layers = len(self.layers) - self.start_idx

def __call__(
self,
Expand All @@ -455,18 +435,17 @@ def __call__(
pipeline_size = self.pipeline_size

if cache is None:
cache = [None] * self.num_layers
cache = [None] * len(self.pipeline_layers)
mask = create_attention_mask(
h, cache[0][0] if cache[0] else None, return_array=True
)

# Receive from the previous process in the pipeline

if pipeline_rank < pipeline_size - 1:
h = mx.distributed.recv_like(h, (pipeline_rank + 1))

for i in range(self.num_layers):
h = self.layers[self.start_idx + i](h, mask, cache[i])
for layer, c in zip(self.pipeline_layers, cache):
h = layer(h, mask, c)

# Send to the next process in the pipeline
if pipeline_rank != 0:
Expand Down Expand Up @@ -646,7 +625,7 @@ def shard_heads(w):

@property
def layers(self):
return self.model.layers[self.model.start_idx : self.model.end_idx]
return self.model.pipeline_layers

@property
def cast_predicate(self):
Expand Down
23 changes: 23 additions & 0 deletions tests/model_parallel_tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,29 @@ def test_pipeline(self):
"max_position_embeddings": 256,
"tie_word_embeddings": False,
},
{
"model_type": "deepseek_v32",
"vocab_size": 128,
"hidden_size": 64,
"index_head_dim": 16,
"index_n_heads": 2,
"index_topk": 4,
"intermediate_size": 128,
"moe_intermediate_size": 32,
"num_hidden_layers": 4,
"num_attention_heads": 4,
"num_key_value_heads": 4,
"n_shared_experts": 1,
"n_routed_experts": 4,
"kv_lora_rank": 8,
"q_lora_rank": 8,
"qk_rope_head_dim": 8,
"v_head_dim": 16,
"qk_nope_head_dim": 16,
"num_experts_per_tok": 2,
"first_k_dense_replace": 1,
"max_position_embeddings": 256,
},
]
mx.random.seed(0)
for config in test_configs:
Expand Down
Loading