J-lens workspace reproduction + loop retrofit: lens, band looping, adapters, controls, multi-task evals

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>
This commit is contained in:
Nils
2026-07-14 00:54:12 +02:00
co-authored by Claude Fable 5
commit ef9c08966c
42 changed files with 5068 additions and 0 deletions
+62
View File
@@ -0,0 +1,62 @@
# J-lens reproduction on gemma-4-E2B-it
Reproduction of the core method and experiments of
["Verbalizable Representations Form a Global Workspace in Language Models"](https://transformer-circuits.pub/2026/workspace/index.html)
(Transformer Circuits, 2026) on `google/gemma-4-E2B-it`, running locally on a
DGX Spark (GB10, 128 GB unified memory).
The paper studies Claude Sonnet/Haiku/Opus; here we test whether its central
tool and headline results transfer to a small open-weights model.
## Method
**Jacobian lens.** For each layer , average the Jacobian from the residual
stream at layer , position t to the final-layer residual stream at position
t ≥ t, over positions and a pretraining-like corpus:
J̄_ = E[∂h_final,t / ∂h_,t] (d×d per layer)
Reading an activation: `lens(h) = softmax(W_U · finalnorm(J̄_ h))` → ranked
vocabulary tokens. The J-lens vector of token s at layer is row s of
`W_U J̄_`. Swap intervention: `h ← h + V(σ(c) c)` with `c = V⁺h`,
`V = [v_src; v_tgt]`, σ exchanging the two coordinates.
**Estimation.** Per prompt we compute the *exact* sum of Jacobians over all
(t, t′≥t) pairs with one forward pass and 1536 batched VJPs
(`is_grads_batched=True`): backward from Σ_t h_final,t with identity
cotangents; causal masking makes the per-position gradient equal the
row-sum over t ≥ t. ~20 s/prompt (bf16, eager attention, seq len 64).
Averaged over 300 fineweb-edu documents ⇒ ~620k (t,t) samples per layer.
Sanity check: at the last layer the pair-sum is exactly T·I (verified to 0 error).
## Layout
- `jlens/core.py` — model loading, residual capture, Jacobian estimation,
lens reading, J-lens vectors, swap hooks.
- `scripts/compute_jacobians.py` — corpus averaging (`results/jbar.pt`).
- `scripts/exp1_readouts.py` — two-hop spider, Chinese→English intermediates,
directed modulation (citrus), layer profile.
- `scripts/exp2_swaps.py` — spider↔ant answer flip, France→China broadcast,
country-swap grid.
- `scripts/exp4_regimes.py` — sensor/workspace/motor depth regimes + ignition.
- `scripts/plot_results.py` — heatmaps + layer profile plots.
## Run
```bash
uv venv && uv pip install --index-url https://download.pytorch.org/whl/cu130 \
--extra-index-url https://pypi.org/simple torch transformers accelerate datasets matplotlib
python scripts/compute_jacobians.py --n-prompts 300 --dtype bfloat16
python scripts/exp1_readouts.py results/jbar.pt
python scripts/exp2_swaps.py results/jbar.pt
python scripts/plot_results.py
```
Results: see `RESULTS.md` (E2B in `results/`, 12B in `results-12b/`; set
`JLENS_MODEL` and `JLENS_RESULTS` to switch). Headline: J-lens *readouts*
reproduce at both scales and strengthen with size (spider P 0.39→0.80).
Causal *swaps* need the right write basis, which migrates with scale: token
embeddings at 2B (6/30 capital grid), activation-derived concept vectors at
12B (30/30 grid, 4/4 broadcast incl. Seine→Yangtze) — the J-lens gate
transfers unchanged.