Download gemma_4_e2b_it_hybrid.py from Cactus-Compute/gemma-4-e2b-it-hybrid: direct link, hf CLI and curl.
- Browser
- Download file 9.94 kB
-
https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid/resolve/3504e1951251bc129bbc8a346d58141a41f15c9d/gemma_4_e2b_it_hybrid.py
- Command line
-
hf download hf://Cactus-Compute/gemma-4-e2b-it-hybrid@3504e1951251bc129bbc8a346d58141a41f15c9d/gemma_4_e2b_it_hybrid.py
-
curl -L -o gemma_4_e2b_it_hybrid.py https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid/resolve/3504e1951251bc129bbc8a346d58141a41f15c9d/gemma_4_e2b_it_hybrid.py
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 | |
| 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)} | |
| 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()) | |
| def last_confidence(self) -> Optional[float]: | |
| """Confidence of the most recent generation (None before any).""" | |
| return self.confidence() | |