gemma-4-e2b-it-hybrid / gemma_4_e2b_it_hybrid.py
jakmro's picture
Fix MLX model_file: exclude probe from quantization, sanitize KV-shared layers
723cf0b verified
Raw History Blame
9.94 kB
"""Cactus-Compute/gemma-4-e2b-it-hybrid — single-file mlx-lm model (``model_file`` mechanism).
Drop this file into the converted MLX repo and add ``"model_file":
"gemma_4_e2b_it_hybrid.py"`` to its ``config.json`` (mlx-lm >= 0.30.1, the mechanism
from mlx-lm PR #830). ``mlx_lm.load`` then builds ``Model``/``ModelArgs`` from
this file instead of the built-in ``gemma4_text`` classes.
The model is the stock ``mlx_lm.models.gemma4_text`` Gemma-4 text model plus
the handoff probe (weight prefix ``handoff_probe.*``). During generation the
output of decoder layer ``probe_layer`` is captured at the position that
predicts each generated token; after generation
``model.last_confidence`` -> float in [0, 1] or None
exposes ``confidence = 1 - p_wrong``. All probe math runs in float32.
Capture semantics (tuned to ``mlx_lm.generate_step``):
- multi-token (prefill-chunk) forwards reset the capture buffer and are never
kept; ``generate_step`` always feeds the final prompt token through a
1-token step, and that forward's last position is row 0 of the contract;
- ``generate_step`` pipelines one forward ahead, so the buffer ends with one
lookahead row for a token that is never emitted; ``last_confidence`` drops
that final row. Use ``model.confidence(num_tokens=N)`` to score exactly the
first N generated tokens instead.
Limitations:
- Server-side surfacing is Python-API only for MLX today: an mlx-lm ``Model``
cannot inject tokens into the stream, so there is no in-band
``[[hybrid:confidence=...]]`` trailer here (unlike the transformers repo).
- Speculative decoding is not supported (draft verification feeds multiple
tokens per forward, which resets the buffer).
- Prompts of exactly 2 tokens leave one extra leading row in the buffer (the
1-token prefill chunk is indistinguishable from a decode step). Chat-template
prompts are always far longer.
"""
import math
from dataclasses import dataclass
from typing import Any, List, Optional
import mlx.core as mx
import mlx.nn as nn
from mlx_lm.models import gemma4_text
@dataclass
class ModelArgs(gemma4_text.ModelArgs):
model_type: str = "gemma-4-e2b-it-hybrid"
probe_layer: int = 28
probe_feature_size: int = 1536
probe_max_tokens: int = 1024
probe_proj_dim: int = 32
def __post_init__(self):
super().__post_init__()
if not 0 <= self.probe_layer < self.num_hidden_layers:
raise ValueError(
f"probe_layer={self.probe_layer} must be in "
f"[0, num_hidden_layers={self.num_hidden_layers})"
)
# Hidden widths of the probe MLP head, fixed by the released checkpoint.
PROBE_HEAD_DIMS = (128, 64)
class HandoffProbe(nn.Module):
"""The released handoff probe; parameter tree matches the checkpoint keys
``norm.{weight,bias}``, ``proj.{weight,bias}``, ``attn_query``,
``head.{0,2,4}.{weight,bias}``."""
def __init__(self, feature_size: int, proj_dim: int = 32):
super().__init__()
h1, h2 = PROBE_HEAD_DIMS
self.norm = nn.LayerNorm(feature_size, eps=1e-5)
self.proj = nn.Linear(feature_size, proj_dim)
self.attn_query = mx.zeros((proj_dim,))
self.head = [
nn.Linear(proj_dim, h1),
nn.ReLU(),
nn.Linear(h1, h2),
nn.ReLU(),
nn.Linear(h2, 1),
]
def p_wrong(self, hidden_states: mx.array, max_tokens: int = 1024) -> mx.array:
"""Score ``[T, feature_size]`` rows; float32 math per the contract."""
x = hidden_states[:max_tokens].astype(mx.float32)
x = mx.fast.layer_norm(
x,
self.norm.weight.astype(mx.float32),
self.norm.bias.astype(mx.float32),
1e-5,
)
projected = mx.maximum(
x @ self.proj.weight.astype(mx.float32).T + self.proj.bias.astype(mx.float32),
0.0,
)
scores = projected @ self.attn_query.astype(mx.float32)
scores = scores / math.sqrt(projected.shape[-1])
weights = mx.softmax(scores - scores.max(), axis=-1)
pooled = weights @ projected
h0, h2_, h4 = self.head[0], self.head[2], self.head[4]
h = mx.maximum(
pooled @ h0.weight.astype(mx.float32).T + h0.bias.astype(mx.float32), 0.0
)
h = mx.maximum(
h @ h2_.weight.astype(mx.float32).T + h2_.bias.astype(mx.float32), 0.0
)
logit = h @ h4.weight.astype(mx.float32).T + h4.bias.astype(mx.float32)
return mx.sigmoid(logit)[0]
class _ProbeState:
"""Plain-object row buffer, invisible to the MLX module tree."""
def __init__(self):
self.rows: List[mx.array] = []
def reset(self):
self.rows.clear()
class _CaptureDecoderLayer(gemma4_text.DecoderLayer):
"""Stock decoder layer that reports its output to a capture sink.
Same submodule tree as ``DecoderLayer``, so weight keys are unchanged.
"""
def __call__(self, x, *args, **kwargs):
out = super().__call__(x, *args, **kwargs)
sink = getattr(self, "capture_sink", None)
if sink is not None:
sink(out[0])
return out
class Model(gemma4_text.Model):
def __init__(self, args: ModelArgs):
super().__init__(args)
self.args = args
self.handoff_probe = HandoffProbe(args.probe_feature_size, args.probe_proj_dim)
self._probe_state = _ProbeState()
capture = _CaptureDecoderLayer(args, layer_idx=args.probe_layer)
capture.capture_sink = self._capture_row
self.model.layers[args.probe_layer] = capture
# --- capture -----------------------------------------------------------
def _capture_row(self, hidden: mx.array) -> None:
"""Keep the last position of each 1-token probe-layer forward.
``generate_step``'s prefill loop always leaves the final prompt token
to a 1-token ``_step`` call, so multi-token (prefill-chunk) forwards
never produce a contract row — they just reset the buffer. Row 0 is
the step that consumes the last prompt token (= last prompt position).
"""
rows = self._probe_state.rows
if hidden.shape[1] > 1:
rows.clear()
return
# Keep one row beyond the window so dropping the lookahead row cannot
# lose a real one.
if len(rows) <= self.args.probe_max_tokens:
rows.append(hidden[:, -1, :])
def __call__(
self,
inputs: mx.array,
cache=None,
input_embeddings: Optional[mx.array] = None,
per_layer_inputs: Optional[mx.array] = None,
):
# A fresh cache (or none) means a new generation: reset the buffer.
if cache is None or (len(cache) > 0 and getattr(cache[0], "offset", 0) == 0):
self._probe_state.reset()
return super().__call__(
inputs,
cache=cache,
input_embeddings=input_embeddings,
per_layer_inputs=per_layer_inputs,
)
# --- weight handling ----------------------------------------------------
def sanitize(self, weights):
"""Stock sanitize + drop KV tensors the module tree does not build.
The HF checkpoint ships k/v projections and norms for the KV-shared
layers; the mlx model creates no such modules there (``has_kv`` is
False), and ``load_weights`` is strict about extra keys.
"""
weights = super().sanitize(weights)
shared_start = self.args.num_hidden_layers - self.args.num_kv_shared_layers
if self.args.num_kv_shared_layers <= 0:
return weights
def is_unused_shared_kv(key: str) -> bool:
parts = key.split(".")
return (
len(parts) >= 5
and parts[0] == "model"
and parts[1] == "layers"
and parts[2].isdigit()
and int(parts[2]) >= shared_start
and parts[3] == "self_attn"
and parts[4] in ("k_proj", "v_proj", "k_norm", "v_norm")
)
return {k: v for k, v in weights.items() if not is_unused_shared_kv(k)}
@property
def quant_predicate(self):
"""Never quantize the probe: its math is float32 by contract, and its
tiny Linears would otherwise be 4-bit-quantized by ``mlx_lm.convert``
(their input dims are multiples of the group size)."""
base = gemma4_text.Model.quant_predicate.fget(self)
def predicate(path, module):
if path.startswith("handoff_probe"):
return False
return base(path, module)
return predicate
# --- scoring -----------------------------------------------------------
def reset_probe(self) -> None:
"""Clear the capture buffer (e.g. between manual forward calls)."""
self._probe_state.reset()
def confidence(self, num_tokens: Optional[int] = None) -> Optional[float]:
"""``1 - p_wrong`` for the captured generation.
``num_tokens`` scores exactly the first N captured rows. Without it,
the final row is dropped to discard ``generate_step``'s one-step
lookahead forward. Returns ``None`` when nothing (scoreable) was
captured.
"""
rows = list(self._probe_state.rows)
if not rows or rows[0].shape[0] != 1:
return None
if num_tokens is not None:
rows = rows[:num_tokens]
elif len(rows) >= 2:
rows = rows[:-1]
if not rows:
return None
stacked = mx.concatenate(rows, axis=0) # [T, features]
p_wrong = self.handoff_probe.p_wrong(stacked, self.args.probe_max_tokens)
return float(1.0 - p_wrong.item())
@property
def last_confidence(self) -> Optional[float]:
"""Confidence of the most recent generation (None before any)."""
return self.confidence()