"""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), ) @torch.no_grad() 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 @contextmanager 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"]