gemma-4-e2b-it-hybrid / modeling_gemma_4_e2b_it_hybrid.py
jakmro's picture
Add files using upload-large-folder tool
94981e6 verified
Raw History Blame Contribute Delete
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),
)
@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"]