hallucination / mechanistic_interp /sequence_probe.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
7.18 kB
"""Attention-pooling sequence probes over raw residual-stream hidden states.
Self-contained port of multilayer-sae's ``experiment/training/sequence_probe.py``
(no ``experiment.*`` dependencies) for use by ``train_probe_latent.py``.
Replaces the linear + max-pool probe (``probe/probing.py:LinearProbe`` trained by
``train_binary_probe``) with a per-layer probe that reads the *sequence* of
per-token hidden states directly, with no fixed pooling:
* Operating on raw ``h_l`` (d_model) keeps the probe independent of any SAE
dictionary.
* A learned-query multi-head attention pool replaces max/mean pooling so no
information is discarded by a hard reduction, and the probe co-adapts to
*where* the concept lives in the token sequence (follows HyperSteer,
arXiv:2506.03292).
``forward_logits`` gives one per-sequence logit per layer (for training on a
per-caption label); ``forward_token_logits`` gives a per-token logit per layer
(for detecting the concept live during generation) — both reuse the same weights.
"""
from __future__ import annotations
import torch
import torch.nn as nn
class _AttnPoolProbe(nn.Module):
"""Single-layer attention-pooling probe: (B, T, d_model) -> (B,) logit."""
def __init__(
self,
d_model: int,
d_probe: int = 256,
n_heads: int = 4,
n_ctx_blocks: int = 1,
dropout: float = 0.0,
spectral_norm: bool = False,
):
super().__init__()
self.in_proj = nn.Linear(d_model, d_probe)
self.ctx_blocks = nn.ModuleList([
nn.TransformerEncoderLayer(
d_model=d_probe,
nhead=n_heads,
dim_feedforward=4 * d_probe,
dropout=dropout,
batch_first=True,
norm_first=True,
)
for _ in range(n_ctx_blocks)
])
# Learned query that attends over the token sequence (attention pooling).
self.query = nn.Parameter(torch.zeros(1, 1, d_probe))
nn.init.normal_(self.query, std=0.02)
self.pool_attn = nn.MultiheadAttention(
d_probe, n_heads, dropout=dropout, batch_first=True
)
self.norm = nn.LayerNorm(d_probe)
head = nn.Linear(d_probe, 1)
self.head = nn.utils.spectral_norm(head) if spectral_norm else head
def forward(
self, x: torch.Tensor, key_padding_mask: torch.Tensor | None
) -> torch.Tensor:
# x: (B, T, d_model); key_padding_mask: (B, T) bool, True = ignore (PAD).
x = self.in_proj(x) # (B, T, d_probe)
for blk in self.ctx_blocks:
x = blk(x, src_key_padding_mask=key_padding_mask)
B = x.shape[0]
q = self.query.expand(B, -1, -1) # (B, 1, d_probe)
pooled, _ = self.pool_attn(
q, x, x, key_padding_mask=key_padding_mask, need_weights=False
) # (B, 1, d_probe)
pooled = self.norm(pooled.squeeze(1)) # (B, d_probe)
return self.head(pooled).squeeze(-1) # (B,)
def token_logits(
self, x: torch.Tensor, key_padding_mask: torch.Tensor | None
) -> torch.Tensor:
"""Per-token logits (no attention pooling): (B, T, d_model) -> (B, T).
Reuses in_proj + context blocks + norm + head so a single probe module
supports both the per-sequence (attention-pooled) and per-token readouts.
Used to detect the concept live during generation.
"""
x = self.in_proj(x) # (B, T, d_probe)
for blk in self.ctx_blocks:
x = blk(x, src_key_padding_mask=key_padding_mask)
return self.head(self.norm(x)).squeeze(-1) # (B, T)
class SequenceLayerProbes(nn.Module):
"""One attention-pooling probe per monitored layer, run on raw hidden states."""
def __init__(
self,
layer_indices: list[int],
d_model: int,
d_probe: int = 256,
n_heads: int = 4,
n_ctx_blocks: int = 1,
dropout: float = 0.0,
spectral_norm: bool = False,
):
super().__init__()
self.layer_indices = list(layer_indices)
self.d_model = d_model
self.probes = nn.ModuleList([
_AttnPoolProbe(
d_model, d_probe=d_probe, n_heads=n_heads,
n_ctx_blocks=n_ctx_blocks, dropout=dropout,
spectral_norm=spectral_norm,
)
for _ in layer_indices
])
self._idx = {l: i for i, l in enumerate(self.layer_indices)}
def forward_logits(
self,
feats_seq: dict[int, torch.Tensor], # l -> (B, T, d_model)
key_padding_mask: torch.Tensor, # (B, T) bool, True = PAD/ignore
) -> list[torch.Tensor]:
# Probe runs in fp32 for numerical stability regardless of model dtype.
w_dtype = self.probes[0].in_proj.weight.dtype
return [
self.probes[self._idx[l]](
feats_seq[l].to(w_dtype), key_padding_mask
)
for l in self.layer_indices
]
def forward_token_logits(
self,
feats_seq: dict[int, torch.Tensor], # l -> (B, T, d_model)
key_padding_mask: torch.Tensor, # (B, T) bool, True = PAD/ignore
) -> list[torch.Tensor]:
"""Per-token logits per layer: list of (B, T). For live-generation detection."""
w_dtype = self.probes[0].in_proj.weight.dtype
return [
self.probes[self._idx[l]].token_logits(feats_seq[l].to(w_dtype), key_padding_mask)
for l in self.layer_indices
]
def forward(
self, feats_seq: dict[int, torch.Tensor], key_padding_mask: torch.Tensor
) -> list[torch.Tensor]:
return [torch.sigmoid(z) for z in self.forward_logits(feats_seq, key_padding_mask)]
def sequence_layer_probes_from_checkpoint(
path: str,
device: torch.device | str | None = None,
) -> SequenceLayerProbes:
"""Load a SequenceLayerProbes checkpoint saved by train_probe_latent.py.
The checkpoint stores both the state_dict and the construction meta
(layer_indices, d_model, probe_dim, ...) so the architecture is rebuilt
exactly. Handles a ``module.`` DDP prefix.
"""
# Checkpoint holds only tensors + a plain dict of ints/lists, so the safe
# weights_only loader suffices (no arbitrary-object unpickling).
ckpt = torch.load(path, map_location="cpu", weights_only=True)
meta = ckpt["meta"]
sd = ckpt["state_dict"]
if any(k.startswith("module.") for k in sd):
sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
probes = SequenceLayerProbes(
meta["layer_indices"], meta["d_model"],
d_probe=meta["d_probe"], n_heads=meta["n_heads"],
n_ctx_blocks=meta["n_ctx_blocks"], spectral_norm=meta["spectral_norm"],
)
probes.load_state_dict(sd, strict=True)
if device is not None:
dev = device if isinstance(device, torch.device) else torch.device(device)
probes = probes.to(dev)
return probes