jakmro commited on
Commit
723cf0b
·
verified ·
1 Parent(s): 94981e6

Fix MLX model_file: exclude probe from quantization, sanitize KV-shared layers

Browse files
Files changed (1) hide show
  1. gemma_4_e2b_it_hybrid.py +42 -0
gemma_4_e2b_it_hybrid.py CHANGED
@@ -186,6 +186,48 @@ class Model(gemma4_text.Model):
186
  per_layer_inputs=per_layer_inputs,
187
  )
188
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
189
  # --- scoring -----------------------------------------------------------
190
 
191
  def reset_probe(self) -> None:
 
186
  per_layer_inputs=per_layer_inputs,
187
  )
188
 
189
+ # --- weight handling ----------------------------------------------------
190
+
191
+ def sanitize(self, weights):
192
+ """Stock sanitize + drop KV tensors the module tree does not build.
193
+
194
+ The HF checkpoint ships k/v projections and norms for the KV-shared
195
+ layers; the mlx model creates no such modules there (``has_kv`` is
196
+ False), and ``load_weights`` is strict about extra keys.
197
+ """
198
+ weights = super().sanitize(weights)
199
+ shared_start = self.args.num_hidden_layers - self.args.num_kv_shared_layers
200
+ if self.args.num_kv_shared_layers <= 0:
201
+ return weights
202
+
203
+ def is_unused_shared_kv(key: str) -> bool:
204
+ parts = key.split(".")
205
+ return (
206
+ len(parts) >= 5
207
+ and parts[0] == "model"
208
+ and parts[1] == "layers"
209
+ and parts[2].isdigit()
210
+ and int(parts[2]) >= shared_start
211
+ and parts[3] == "self_attn"
212
+ and parts[4] in ("k_proj", "v_proj", "k_norm", "v_norm")
213
+ )
214
+
215
+ return {k: v for k, v in weights.items() if not is_unused_shared_kv(k)}
216
+
217
+ @property
218
+ def quant_predicate(self):
219
+ """Never quantize the probe: its math is float32 by contract, and its
220
+ tiny Linears would otherwise be 4-bit-quantized by ``mlx_lm.convert``
221
+ (their input dims are multiples of the group size)."""
222
+ base = gemma4_text.Model.quant_predicate.fget(self)
223
+
224
+ def predicate(path, module):
225
+ if path.startswith("handoff_probe"):
226
+ return False
227
+ return base(path, module)
228
+
229
+ return predicate
230
+
231
  # --- scoring -----------------------------------------------------------
232
 
233
  def reset_probe(self) -> None: