From ee044da0a8b906fbb903c2eba7700fb7d50ef9f3 Mon Sep 17 00:00:00 2001 From: Awni Hannun Date: Tue, 18 Mar 2025 08:57:53 -0700 Subject: [PATCH] dequantize dsv3 (#32) --- mlx_lm/models/deepseek_v3.py | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/mlx_lm/models/deepseek_v3.py b/mlx_lm/models/deepseek_v3.py index 5cd40a0..348e2a8 100644 --- a/mlx_lm/models/deepseek_v3.py +++ b/mlx_lm/models/deepseek_v3.py @@ -483,6 +483,35 @@ class Model(nn.Module): return self.lm_head(out) def sanitize(self, weights): + def dequant(weight, scale_inv): + bs = 128 # block size + m, n = weight.shape + pad_bottom = (-m) % bs + pad_side = (-n) % bs + weight = mx.pad(weight, ((0, pad_bottom), (0, pad_side))) + weight = weight.reshape( + ((m + pad_bottom) // bs, bs, (n + pad_side) // bs, bs) + ) + scale_inv = scale_inv.astype(weight.dtype) + weight = (weight * scale_inv[:, None, :, None]).reshape( + m + pad_bottom, n + pad_side + ) + return weight[:m, :n] + + # Dequantize + new_weights = {} + for k, v in weights.items(): + if "weight_scale_inv" in k: + scale_inv = v + wk = k.replace("_scale_inv", "") + weight = weights[wk] + weight = dequant(weight, scale_inv) + new_weights[wk] = weight + elif k not in new_weights: + new_weights[k] = v + weights = new_weights + + # Stack experts 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")]: