Files
NilsandClaude Fable 5 ef9c08966c 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>
2026-07-14 00:54:12 +02:00

63 lines
3.0 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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.