← back to Exo
update deepseek sanitize to shard layers first before handle switch
a6bb8ddf41c02afd2f601401a397948a48d9e843 · 2024-07-28 12:58:09 +1000 · Anchen
Files touched
M exo/inference/mlx/models/deepseek_v2.py
Diff
commit a6bb8ddf41c02afd2f601401a397948a48d9e843
Author: Anchen <li.anchen.au@gmail.com>
Date: Sun Jul 28 12:58:09 2024 +1000
update deepseek sanitize to shard layers first before handle switch
---
exo/inference/mlx/models/deepseek_v2.py | 28 +++++++++++++++-------------
1 file changed, 15 insertions(+), 13 deletions(-)
diff --git a/exo/inference/mlx/models/deepseek_v2.py b/exo/inference/mlx/models/deepseek_v2.py
index 9cc8ea87..585798a6 100644
--- a/exo/inference/mlx/models/deepseek_v2.py
+++ b/exo/inference/mlx/models/deepseek_v2.py
@@ -89,19 +89,6 @@ class Model(nn.Module):
return out
def sanitize(self, weights):
- for l in range(self.args.num_hidden_layers):
- prefix = f"model.layers.{l}"
- for n, m in [("w1", "gate_proj"), ("w2", "down_proj"), ("w3", "up_proj")]:
- for k in ["weight", "scales", "biases"]:
- if f"{prefix}.mlp.experts.0.{m}.{k}" in weights:
- to_join = [
- weights.pop(f"{prefix}.mlp.experts.{e}.{m}.{k}") for e in range(self.args.n_routed_experts)
- ]
- weights[
- f"{prefix}.mlp.switch_mlp.{
- m}.{k}"
- ] = mx.stack(to_join)
-
shard_state_dict = {}
for key, value in weights.items():
@@ -113,6 +100,21 @@ class Model(nn.Module):
shard_state_dict[key] = value
elif self.args.shard.is_last_layer() and (key.startswith('model.norm') or key.startswith('lm_head')):
shard_state_dict[key] = value
+
+ for l in range(self.args.num_hidden_layers):
+ prefix = f"model.layers.{l}"
+ for n, m in [("w1", "gate_proj"), ("w2", "down_proj"), ("w3", "up_proj")]:
+ for k in ["weight", "scales", "biases"]:
+ if f"{prefix}.mlp.experts.0.{m}.{k}" in shard_state_dict:
+ to_join = [
+ shard_state_dict.pop(f"{prefix}.mlp.experts.{e}.{m}.{k}") for e in range(self.args.n_routed_experts)
+ ]
+ shard_state_dict[
+ f"{prefix}.mlp.switch_mlp.{
+ m}.{k}"
+ ] = mx.stack(to_join)
+
+
return shard_state_dict
@property
← cb217b7b format format.py
·
back to Exo
·
formatting 44413777 →