File size: 7,042 Bytes
94981e6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 | """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"]
|