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>
87 lines
2.8 KiB
Python
87 lines
2.8 KiB
Python
"""Experiment 3: J-space occupancy.
|
|
|
|
Decompose residual activations as sparse non-negative combinations of k<=25
|
|
J-lens vectors (matching pursuit + NNLS refit) and measure the fraction of
|
|
activation variance the J-space carries per layer (paper: ~10%, intermediate
|
|
layers only).
|
|
"""
|
|
|
|
import os, sys
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import torch
|
|
from scipy.optimize import nnls
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
|
from jlens.core import JLens, chat_ids, collect_residuals, load_model
|
|
|
|
RES = Path(os.environ.get("JLENS_RESULTS", "results"))
|
|
|
|
K = 25
|
|
|
|
PROMPTS = [
|
|
"The animal that spins webs has how many legs? Answer with just a number.",
|
|
"Write one sentence about the ocean.",
|
|
"What is the capital of France? Answer with just the city name.",
|
|
"Explain photosynthesis in one sentence.",
|
|
]
|
|
|
|
|
|
def pursuit_r2(D, h, k=K):
|
|
"""Non-negative matching pursuit of h onto rows of D. Returns R^2, support."""
|
|
r = h.clone()
|
|
sel = []
|
|
for _ in range(k):
|
|
scores = D @ r
|
|
if sel:
|
|
scores[torch.tensor(sel, device=D.device)] = -1e30
|
|
j = int(scores.argmax())
|
|
if scores[j] <= 0:
|
|
break
|
|
sel.append(j)
|
|
A = D[sel].T # (d, |sel|)
|
|
x, _ = nnls(A.detach().cpu().numpy().astype(np.float64),
|
|
h.detach().cpu().numpy().astype(np.float64))
|
|
approx = A @ torch.tensor(x, device=D.device, dtype=D.dtype)
|
|
r = h - approx
|
|
r2 = 1 - (r.norm() / h.norm()) ** 2
|
|
return float(r2), sel
|
|
|
|
|
|
def main():
|
|
model, tok = load_model(dtype=torch.bfloat16)
|
|
ck = torch.load(sys.argv[1] if len(sys.argv) > 1 else "results/jbar.pt",
|
|
map_location="cuda")
|
|
jl = JLens(model, tok, ck["Jbar"].cuda())
|
|
WU = model.lm_head.weight.detach().float() # (V, d)
|
|
|
|
L = len(model.model.language_model.layers)
|
|
layers = list(range(1, L, 3))
|
|
r2_by_layer = {l: [] for l in layers}
|
|
for p in PROMPTS:
|
|
ids = chat_ids(tok, p)
|
|
hs = collect_residuals(model, ids)
|
|
T = ids.shape[1]
|
|
positions = list(range(max(1, T - 12), T)) # skip BOS region
|
|
for l in layers:
|
|
D = WU @ jl.Jbar[l] # (V, d) J-lens dictionary at layer l
|
|
D = D / D.norm(dim=1, keepdim=True).clamp_min(1e-8)
|
|
for t in positions:
|
|
r2, _ = pursuit_r2(D, hs[l, t].float())
|
|
r2_by_layer[l].append(r2)
|
|
del D
|
|
torch.cuda.empty_cache()
|
|
print(f"done prompt: {p[:40]}...", flush=True)
|
|
|
|
print("\nlayer | mean R^2 of k<=25 non-negative J-lens pursuit")
|
|
means = {}
|
|
for l in layers:
|
|
means[l] = float(np.mean(r2_by_layer[l]))
|
|
print(f"L{l:>2} | {means[l]:.3f}")
|
|
torch.save(means, RES / "jspace_r2.pt")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|