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"]