Reproduction of the 2026 workspace/J-lens paper on gemma-4 (E2B/12B/26B), plus the workspace-loop retrofit line: merge adapter, prompt-only latent planning (MBPP), carry variant, attribution controls (FF/pause/untrained), band-location ablation, Blocksworld harness, 12B replication scripts. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
262 lines
10 KiB
Python
262 lines
10 KiB
Python
"""Jacobian lens (J-lens) for gemma-4-E2B-it.
|
|
|
|
Reproduces the method of "Verbalizable Representations Form a Global
|
|
Workspace in Language Models" (transformer-circuits.pub/2026/workspace):
|
|
|
|
J_l = E_{prompt, t, t' >= t} [ d h_{final, t'} / d h_{l, t} ] (d x d per layer)
|
|
lens(h_l) = softmax(W_U . finalnorm(J_l h_l))
|
|
J-lens vector of token s at layer l: v_s = (W_U J_l)[s, :] = W_U[s] J_l
|
|
swap: h <- h + V (sigma(c) - c), c = pinv(V) h, V = [v_s; v_t]
|
|
|
|
h_l is the residual stream at the *output* of decoder layer l.
|
|
h_final is the output of the last decoder layer (pre final-RMSNorm).
|
|
"""
|
|
|
|
import os
|
|
|
|
import torch
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
|
|
MODEL_ID = os.environ.get("JLENS_MODEL", "google/gemma-4-E2B-it")
|
|
|
|
|
|
def load_model(dtype=torch.float32, device="cuda", model_id=None):
|
|
model_id = model_id or MODEL_ID
|
|
tok = AutoTokenizer.from_pretrained(model_id)
|
|
model = AutoModelForCausalLM.from_pretrained(
|
|
model_id, dtype=dtype, attn_implementation="eager"
|
|
).to(device).eval()
|
|
torch.backends.cuda.matmul.allow_tf32 = True
|
|
torch.backends.cudnn.allow_tf32 = True
|
|
return model, tok
|
|
|
|
|
|
def _text_model(model):
|
|
return model.model.language_model
|
|
|
|
|
|
class ResidualCapture:
|
|
"""Forward hooks that record each decoder layer's output residual stream."""
|
|
|
|
def __init__(self, model, layers=None, detach=True):
|
|
self.tm = _text_model(model)
|
|
self.layers = list(range(len(self.tm.layers))) if layers is None else layers
|
|
self.detach = detach
|
|
self.acts = {}
|
|
self.handles = []
|
|
|
|
def __enter__(self):
|
|
for i in self.layers:
|
|
def hook(mod, inp, out, i=i):
|
|
self.acts[i] = out.detach() if self.detach else out
|
|
self.handles.append(self.tm.layers[i].register_forward_hook(hook))
|
|
return self
|
|
|
|
def __exit__(self, *a):
|
|
for h in self.handles:
|
|
h.remove()
|
|
|
|
|
|
@torch.no_grad()
|
|
def collect_residuals(model, input_ids):
|
|
"""Return (n_layers, T, d) residual stream stack for a single prompt."""
|
|
with ResidualCapture(model) as cap:
|
|
model(input_ids=input_ids)
|
|
n = len(_text_model(model).layers)
|
|
return torch.stack([cap.acts[i][0] for i in range(n)])
|
|
|
|
|
|
@torch.no_grad()
|
|
def generate_with_residuals(model, tok, input_ids, max_new_tokens=40):
|
|
"""Greedy-generate while capturing the residual stream at every position
|
|
(prompt + generated). Returns (text, ids (T,), residuals (L, T, d))."""
|
|
tm = _text_model(model)
|
|
L = len(tm.layers)
|
|
steps = {i: [] for i in range(L)}
|
|
handles = []
|
|
for i in range(L):
|
|
def hook(mod, inp, out, i=i):
|
|
steps[i].append(out[0].detach())
|
|
handles.append(tm.layers[i].register_forward_hook(hook))
|
|
try:
|
|
out = model.generate(input_ids, max_new_tokens=max_new_tokens,
|
|
do_sample=False,
|
|
pad_token_id=tok.pad_token_id or tok.eos_token_id)
|
|
finally:
|
|
for h in handles:
|
|
h.remove()
|
|
hs = torch.stack([torch.cat(steps[i], dim=0) for i in range(L)])
|
|
text = tok.decode(out[0, input_ids.shape[1]:], skip_special_tokens=True)
|
|
return text, out[0], hs
|
|
|
|
|
|
def prompt_jacobian_pairsum(model, input_ids, chunk=128):
|
|
"""Sum over all pairs (t, t' >= t) of d h_{final,t'} / d h_{l,t} for one prompt.
|
|
|
|
One forward pass; backward from S = sum_{t'} h_{final,t'} with batched
|
|
identity cotangents. The gradient of <e_i, S> w.r.t. h_{l,t} is row i of
|
|
sum_{t' >= t} J^{(t',t)}_l (causality zeroes t' < t), so summing the
|
|
gradient over t gives row i of the pair-sum. Returns (fp32 (L, d, d), n_pairs).
|
|
"""
|
|
tm = _text_model(model)
|
|
L = len(tm.layers)
|
|
d = model.config.get_text_config().hidden_size
|
|
T = input_ids.shape[1]
|
|
dev = input_ids.device
|
|
|
|
with ResidualCapture(model, detach=False) as cap:
|
|
with torch.enable_grad():
|
|
model(input_ids=input_ids)
|
|
hs = [cap.acts[i] for i in range(L)] # each (1, T, d), in graph
|
|
S = hs[-1].sum(dim=1) # (1, d)
|
|
|
|
J = torch.zeros(L, d, d, device=dev, dtype=torch.float32)
|
|
eye = torch.eye(d, device=dev, dtype=S.dtype)
|
|
starts = list(range(0, d, chunk))
|
|
for ci, s in enumerate(starts):
|
|
e = s + min(chunk, d - s)
|
|
V = eye[s:e].unsqueeze(1) # (B, 1, d) cotangents
|
|
grads = torch.autograd.grad(
|
|
outputs=S, inputs=hs, grad_outputs=V,
|
|
retain_graph=(ci < len(starts) - 1), is_grads_batched=True,
|
|
)
|
|
for l in range(L):
|
|
# grads[l]: (B, 1, T, d); row i of pair-sum = sum over t
|
|
J[l, s:e] += grads[l].squeeze(1).sum(dim=1).float()
|
|
n_pairs = T * (T + 1) // 2
|
|
return J, n_pairs
|
|
|
|
|
|
class JLens:
|
|
"""Averaged-Jacobian lens: reading, J-lens vectors, swap interventions."""
|
|
|
|
def __init__(self, model, tok, Jbar):
|
|
"""Jbar: (L, d, d) fp32 averaged Jacobian per layer."""
|
|
self.model, self.tok = model, tok
|
|
self.Jbar = Jbar
|
|
self.tm = _text_model(model)
|
|
self.softcap = model.config.get_text_config().final_logit_softcapping
|
|
|
|
def _readout(self, x):
|
|
"""finalnorm + unembed (+softcap) of residual-space vectors x (..., d)."""
|
|
x = self.tm.norm(x.to(self.tm.norm.weight.dtype))
|
|
logits = self.model.lm_head(x)
|
|
if self.softcap:
|
|
logits = self.softcap * torch.tanh(logits / self.softcap)
|
|
return logits
|
|
|
|
@torch.no_grad()
|
|
def read(self, h, layer, topk=10):
|
|
"""J-lens reading of residual vector(s) h (..., d) at `layer`.
|
|
|
|
Returns (topk token ids, topk probs)."""
|
|
proj = h.float() @ self.Jbar[layer].T
|
|
logits = self._readout(proj)
|
|
probs = torch.softmax(logits.float(), dim=-1)
|
|
p, idx = probs.topk(topk, dim=-1)
|
|
return idx, p
|
|
|
|
@torch.no_grad()
|
|
def read_prompt(self, input_ids, layers, topk=8):
|
|
"""Full J-lens table for one prompt: {layer: (ids (T,k), probs (T,k))}."""
|
|
hs = collect_residuals(self.model, input_ids)
|
|
return {l: self.read(hs[l], l, topk=topk) for l in layers}
|
|
|
|
@torch.no_grad()
|
|
def concept_prob(self, input_ids, token_ids, layers):
|
|
"""P(token) under the lens for each (layer, position). (L, T, n_tokens)."""
|
|
hs = collect_residuals(self.model, input_ids)
|
|
out = []
|
|
for l in layers:
|
|
logits = self._readout(hs[l].float() @ self.Jbar[l].T)
|
|
probs = torch.softmax(logits.float(), dim=-1)
|
|
out.append(probs[:, token_ids])
|
|
return torch.stack(out)
|
|
|
|
def jlens_vector(self, token_id, layer):
|
|
"""v_s = W_U[s] J_l : residual-stream direction at `layer` for token s."""
|
|
wu = self.model.lm_head.weight[token_id].float()
|
|
return wu @ self.Jbar[layer]
|
|
|
|
def swap_hooks(self, pairs, layers=None, thr=0.005, alpha=1.0,
|
|
write="embed"):
|
|
"""Gated concept swap. `pairs` is a list of (src_str, tgt_str) tokens.
|
|
|
|
At each (layer, position) where the J-lens reads any source token with
|
|
probability > thr (the workspace "holds" the concept), transfer the
|
|
activation's content from source to target:
|
|
|
|
- write="embed" (default): move h's projection on the unit source
|
|
embedding onto the unit target embedding (tied-embedding write basis;
|
|
the J-lens supplies localization). This is the variant that works on
|
|
gemma-4-E2B.
|
|
- write="jlens": the paper's h <- h + V(sigma(c) - c), c = pinv(V) h,
|
|
with V rows = J-lens vectors of source+target tokens.
|
|
|
|
Returns (handles, fired) where fired collects (layer, n_positions)."""
|
|
layers = range(4, len(self.tm.layers) - 1) if layers is None else layers
|
|
enc = lambda s: self.tok.encode(s, add_special_tokens=False)[0]
|
|
s_ids = torch.tensor([enc(s) for s, _ in pairs], device="cuda")
|
|
t_ids = torch.tensor([enc(t) for _, t in pairs], device="cuda")
|
|
W = self.model.lm_head.weight
|
|
E_s = torch.nn.functional.normalize(W[s_ids].float(), dim=-1)
|
|
E_t = torch.nn.functional.normalize(W[t_ids].float(), dim=-1)
|
|
handles, fired = [], []
|
|
for l in layers:
|
|
V = pinv = None
|
|
m = len(pairs)
|
|
if write == "jlens":
|
|
V = torch.cat([
|
|
torch.stack([self.jlens_vector(int(s), l) for s in s_ids]),
|
|
torch.stack([self.jlens_vector(int(t), l) for t in t_ids])])
|
|
pinv = torch.linalg.pinv(V)
|
|
|
|
def hook(mod, inp, out, l=l, V=V, pinv=pinv, m=m):
|
|
h = out.float()
|
|
probs = torch.softmax(
|
|
self._readout(h @ self.Jbar[l].T).float(), -1)
|
|
gate = probs[..., s_ids].sum(-1) > thr # (B, T)
|
|
if not gate.any():
|
|
return None
|
|
if write == "embed":
|
|
proj = h @ E_s.T
|
|
delta = alpha * (proj @ E_t - proj @ E_s)
|
|
else:
|
|
c = h @ pinv
|
|
sig = torch.cat([c[..., m:], c[..., :m]], dim=-1)
|
|
delta = alpha * (sig - c) @ V
|
|
fired.append((l, int(gate.sum())))
|
|
return (h + delta * gate.unsqueeze(-1)).to(out.dtype)
|
|
|
|
handles.append(self.tm.layers[l].register_forward_hook(hook))
|
|
return handles, fired
|
|
|
|
@torch.no_grad()
|
|
def generate(self, input_ids, max_new_tokens=20):
|
|
out = self.model.generate(
|
|
input_ids, max_new_tokens=max_new_tokens, do_sample=False,
|
|
pad_token_id=self.tok.pad_token_id or self.tok.eos_token_id,
|
|
)
|
|
return self.tok.decode(out[0, input_ids.shape[1]:], skip_special_tokens=True)
|
|
|
|
@torch.no_grad()
|
|
def generate_swapped(self, input_ids, pairs, layers=None, thr=0.005,
|
|
alpha=1.0, write="embed", max_new_tokens=20):
|
|
handles, fired = self.swap_hooks(pairs, layers, thr, alpha, write)
|
|
try:
|
|
return self.generate(input_ids, max_new_tokens), fired
|
|
finally:
|
|
for h in handles:
|
|
h.remove()
|
|
|
|
|
|
def chat_ids(tok, user_msg, device="cuda", assistant_prefix=None):
|
|
enc = tok.apply_chat_template(
|
|
[{"role": "user", "content": user_msg}],
|
|
add_generation_prompt=True, return_tensors="pt", return_dict=True)
|
|
ids = enc["input_ids"]
|
|
if assistant_prefix:
|
|
pre = tok(assistant_prefix, add_special_tokens=False, return_tensors="pt")
|
|
ids = torch.cat([ids, pre["input_ids"]], dim=1)
|
|
return ids.to(device)
|