Download modeling_gemma_4_e2b_it_hybrid.py from Cactus-Compute/gemma-4-e2b-it-hybrid: direct link, hf CLI and curl.
- Browser
- Download file 7.04 kB
-
https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid/resolve/main/modeling_gemma_4_e2b_it_hybrid.py
- Command line
-
hf download hf://Cactus-Compute/gemma-4-e2b-it-hybrid/modeling_gemma_4_e2b_it_hybrid.py
-
curl -L -o modeling_gemma_4_e2b_it_hybrid.py https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid/resolve/main/modeling_gemma_4_e2b_it_hybrid.py
7.04 kB
| """Cactus-Compute/gemma-4-e2b-it-hybrid — Gemma-4 causal LM with a handoff probe. | |
| ``Gemma4E2BItHybridForCausalLM`` is the stock ``Gemma4ForCausalLM`` plus a small | |
| "handoff probe" head (weight prefix ``handoff_probe.*``) that scores each | |
| generation with ``confidence = 1 - p_wrong``. Base weights keep identical keys, | |
| so the checkpoint is the stock checkpoint with eleven extra probe tensors. | |
| Probe contract (checkpoint layer 28, float32 math): | |
| - input: ``[T, 1536]`` — output of decoder layer index ``config.probe_layer`` | |
| at the position that predicts each generated token (row 0 = last prompt | |
| position at prefill, row t = position captured at generation step t). Only | |
| the first ``config.probe_max_tokens`` rows are used. | |
| - ``x = LayerNorm(x, eps=1e-5) * norm.weight + norm.bias`` | |
| - ``p = relu(x @ proj.weight.T + proj.bias)`` | |
| - ``s = p @ attn_query / sqrt(probe_proj_dim)``; ``w = softmax_T(s)`` | |
| - ``pooled = w @ p`` | |
| - ``h = relu(head.0 @ pooled); h = relu(head.2 @ h); logit = head.4 @ h`` | |
| - ``p_wrong = sigmoid(logit)``; ``confidence = 1 - p_wrong`` | |
| Layer capture uses a forward hook that keeps only the last position of each | |
| decode step (a ``[batch, hidden]`` row), never the full hidden-state stack. | |
| """ | |
| import math | |
| from contextlib import contextmanager | |
| import torch | |
| import torch.nn.functional as F | |
| from torch import nn | |
| from transformers.generation.utils import GenerationMixin | |
| from transformers.models.gemma4.modeling_gemma4 import Gemma4ForCausalLM | |
| try: | |
| from .configuration_gemma_4_e2b_it_hybrid import Gemma4E2BItHybridConfig | |
| except ImportError: # direct (non-package) execution | |
| from configuration_gemma_4_e2b_it_hybrid import Gemma4E2BItHybridConfig | |
| # 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. Weight keys match the probe checkpoint: | |
| ``norm.{weight,bias}``, ``proj.{weight,bias}``, ``attn_query``, | |
| ``head.{0,2,4}.{weight,bias}``. | |
| """ | |
| def __init__(self, feature_size: int, proj_dim: int = 32) -> None: | |
| 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 = nn.Parameter(torch.zeros(proj_dim)) | |
| self.head = nn.Sequential( | |
| nn.Linear(proj_dim, h1), | |
| nn.ReLU(), | |
| nn.Linear(h1, h2), | |
| nn.ReLU(), | |
| nn.Linear(h2, 1), | |
| ) | |
| def p_wrong(self, hidden_states: torch.Tensor, max_tokens: int = 1024) -> float: | |
| """Score ``[T, feature_size]`` generated-token hidden states. | |
| All math runs in float32 on the probe's own device, regardless of the | |
| dtype the module weights were loaded in; the input may arrive on any | |
| device (capture hooks store rows on CPU). | |
| """ | |
| if hidden_states.ndim != 2: | |
| raise ValueError(f"expected [tokens, features], got {tuple(hidden_states.shape)}") | |
| if hidden_states.shape[0] == 0: | |
| raise ValueError("cannot score an empty generation") | |
| x = hidden_states[:max_tokens].to(self.norm.weight.device, torch.float32) | |
| x = F.layer_norm( | |
| x, (x.shape[-1],), self.norm.weight.float(), self.norm.bias.float(), self.norm.eps | |
| ) | |
| projected = F.relu(F.linear(x, self.proj.weight.float(), self.proj.bias.float())) | |
| scores = projected @ self.attn_query.float() / math.sqrt(projected.shape[-1]) | |
| weights = torch.softmax(scores, dim=0) | |
| pooled = weights @ projected | |
| h = F.relu(F.linear(pooled, self.head[0].weight.float(), self.head[0].bias.float())) | |
| h = F.relu(F.linear(h, self.head[2].weight.float(), self.head[2].bias.float())) | |
| logit = F.linear(h, self.head[4].weight.float(), self.head[4].bias.float()) | |
| return float(torch.sigmoid(logit)[0].item()) | |
| class Gemma4E2BItHybridForCausalLM(Gemma4ForCausalLM): | |
| """Stock Gemma-4 causal LM plus the ``handoff_probe.*`` scoring head.""" | |
| config_class = Gemma4E2BItHybridConfig | |
| config: Gemma4E2BItHybridConfig | |
| def __init__(self, config: Gemma4E2BItHybridConfig) -> None: | |
| if config.probe_feature_size != config.hidden_size: | |
| raise ValueError( | |
| f"probe_feature_size={config.probe_feature_size} must equal " | |
| f"hidden_size={config.hidden_size}" | |
| ) | |
| super().__init__(config) | |
| self.handoff_probe = HandoffProbe( | |
| config.probe_feature_size, getattr(config, "probe_proj_dim", 32) | |
| ) | |
| #: Confidence of the most recent scored generation (``None`` before any). | |
| self.last_confidence: float | None = None | |
| def probe_capture(self): | |
| """Capture probe-layer rows during generation. | |
| Yields a list that fills with one ``[batch, hidden]`` float32 CPU tensor | |
| per decode step (row 0 comes from the prefill forward at the last prompt | |
| position). At most ``config.probe_max_tokens`` rows are kept. | |
| """ | |
| rows: list[torch.Tensor] = [] | |
| layer = self.model.layers[self.config.probe_layer] | |
| max_rows = self.config.probe_max_tokens | |
| def hook(module, args, output): | |
| if len(rows) < max_rows: | |
| hidden = output[0] if isinstance(output, tuple) else output | |
| rows.append(hidden[:, -1, :].detach().to(torch.float32).cpu()) | |
| handle = layer.register_forward_hook(hook) | |
| try: | |
| yield rows | |
| finally: | |
| handle.remove() | |
| def confidence_from_rows(self, rows: list[torch.Tensor]) -> float | None: | |
| """Turn captured probe rows into ``confidence = 1 - p_wrong``. | |
| Returns ``None`` when the capture is empty or not a single-sequence | |
| generation (batch size > 1, beam search, ...), for which the probe | |
| contract is undefined. | |
| """ | |
| if not rows or any(row.shape[0] != 1 for row in rows): | |
| return None | |
| states = torch.cat(rows, dim=0) # [T, hidden] | |
| p_wrong = self.handoff_probe.p_wrong(states, self.config.probe_max_tokens) | |
| self.last_confidence = 1.0 - p_wrong | |
| return self.last_confidence | |
| def generate_with_confidence(self, *args, **kwargs): | |
| """Run the stock ``generate`` and score it with the handoff probe. | |
| Returns ``(sequences, confidence)`` where ``sequences`` is exactly what | |
| the stock ``generate`` returns for the given arguments (no in-band | |
| trailer is appended) and ``confidence`` is a float in ``[0, 1]``, or | |
| ``None`` when the generation cannot be scored (batch > 1, beams, or | |
| assisted decoding). | |
| """ | |
| kwargs.pop("return_confidence", None) | |
| with self.probe_capture() as rows: | |
| sequences = GenerationMixin.generate(self, *args, **kwargs) | |
| confidence = self.confidence_from_rows(rows) | |
| return sequences, confidence | |
| __all__ = ["Gemma4E2BItHybridForCausalLM", "HandoffProbe"] | |