Fix MLX model_file: exclude probe from quantization, sanitize KV-shared layers
Browse files- 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:
|