Files
jspace/WORKSPACE_LOOPING.md
T
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

16 KiB
Raw Blame History

Latent workspace recursion via a trained merge layer

A design note on one idea that came out of the J-lens work: can we loop the workspace band of a pretrained model to spend more compute per token — and if so, how do we make it actually improve reasoning rather than just idle?

Reproduction results this builds on are in RESULTS.md; the schematic is results/jlens_diagram.png; a distilled worked/failed/pitfalls ledger is in LESSONS.md. All experiments below are on google/gemma-4-E2B-it (35 layers, d=1536) on a DGX Spark.


1. Why the workspace is a candidate for looping

The J-lens analysis gives every layer a readout of "what verbalizable concept is here." Across models that partitions depth into four regimes:

regime E2B layers what the J-lens reads
transduction 05 noise / surface lexical form
sensor 613 the current input token (echo ↑26%)
workspace 1430 unspoken abstract content ('spider', 'big'); held & broadcast
motor 3134 the next output token (→68%)

Two measured properties make the workspace look like a recurrent variable that is currently unrolled across distinct layers rather than iterated:

  • It is slow / persistent. Top-10 J-lens content overlaps ~0.25 (Jaccard) between adjacent token positions — the workspace barely changes token to token.
  • It holds unspoken intermediates. On "the animal that spins webs has how many legs?" the lens surfaces spider mid-network before the model emits 8; on multi-step arithmetic it stages 21 → 42 → 49 across layers.

If that band is effectively approximating "iterate toward a settled abstract representation," then looping it should be a way to add sequential compute without adding parameters — the premise behind Universal Transformers, Deep Equilibrium Models, and recent recurrent-depth LLMs (Huginn, Ouro, Mixture-of-Recursions).

Workspace-band parameter mass (decoder only), for reference:

model band params in band active/token
E2B L1430 1.08 B (58%) 1.08 B
12B L3645 2.26 B (21%) 2.26 B
26B-MoE L2027 6.53 B (27%) ≈0.82 B

2. What we tried, and what actually happens

Probe: inject a residual at the band entrance, run the band, feed the result back, and track (a) convergence |Δ|/|·|, (b) the J-lens 'spider' concept, (c) the emitted answer. Baseline answer is 8.

2a. Naive band loop — fails (ill-posed)

Feed L30's output straight back into L14. Collapses in one iteration (8 → )). Reason: the band is not a self-map — its output lives 17 layers of accumulation downstream of its input (larger norm, "settled" content), so re-injecting it double-counts the band's contribution and leaves the input distribution L14 expects.

2b. Anchored H ← H0 + bandΔ(H) — also fails

Merging the fixed pre-band residual H0 as an additive anchor does not help on its own: bandΔ is a large forward displacement, so H0 + bandΔ lands back in post-band space and diverges (|ΔH|/|H| ≈ 2 per step).

2c. Single tied layer h ← h + β·Δ_L(h) — well-posed but not useful

One layer is a genuine self-map (same residual point in and out). Result: stable, converges (|ΔH| decreases monotonically), and low-pass damping extends it — β=0.5 at L26 keeps 8 through ~6 loops vs. immediate decay at β=1.0. But a frozen layer has no useful fixed point: pushed far enough it relaxes to a degenerate attractor (Set, number, …) and the J-space concept never sharpens. Confirms the Relaxed-Recursive-Transformer lesson: layers not trained to be looped only tolerate it, they don't benefit.

2d. Merge layer between L13 and L14 — the fix (structurally)

Insert a merge that reconciles the looped-back post-band state s with the fixed pre-band anchor e = L13 output:

L14_in = (1−α)·e + α·(s renormalized to |e|)

Untrained, anchor-dominant (α=0.3):

| k | |Δs|/|s| | |s|/|e| | answer | P(ans) | |---|---|---|---|---| | 0 | 0.00 | 2.23 | 8 | 1.000 | | 1 | 0.29 | 2.11 | 8 | 1.000 | | 5 | 0.12 | 2.09 | 8 | 1.000 | | 10 | 0.10 | 2.09 | 8 | 1.000 |

The band is now a well-posed, stable, answer-preserving recurrence: converges to a fixed point (|Δs| 0.29→0.10), norm controlled, 8 intact through 11+ loops — versus collapse-in-one-step for the naive loop.

Caveats: it only stabilizes when the anchor dominates (α≥0.6 drifts to garbage), and it holds rather than sharpens — the 'spider' J-space concept fluctuates (~0.010.04) rather than growing. Untrained, the merge buys a stable substrate, not extra reasoning.

This is the DEQ / Huginn "input-injection" structure: a fixed input (anchor) merged with an iterated state makes the core a self-map. Our merge is that injection, done by hand.


3. The proposal: train only the merge layer

The merge layer converts the ill-posed loop into a stable, answer-preserving substrate, and it is the single small component to train — a little MLP / LoRA at L13→L14 — to make iterations improve the answer, with L1430 and the rest of the model frozen. This is the cheapest possible "train the loop in."

        e (L13 out, fixed anchor)
                │
  s_k ──▶ [ merge adapter θ ] ──▶ L14 … L30 ──▶ s_{k+1}   (× k loops)
                                        │
                                    L31 … L34 ─▶ logits (answer)

Objective (self-contained, no bigger teacher)

Latent chain-of-thought distillation (Coconut / STaR-style):

  1. Let the frozen base model produce CoT answers; keep only traces whose final answer is correct (rejection sampling).
  2. Train the adapter so the latent looped forward — k loops, no CoT emitted — reaches that correct answer. Loss = cross-entropy on answer tokens only; adapter is the sole trainable tensor; band unrolled k times.

The supervision is answers the model can reach with CoT but not without — so loop depth substitutes for reasoning tokens, and everything stays in-distribution for the frozen base.

Training data

  • GSM8K (already cached: openai/gsm8k) — multi-step arithmetic; single verifiable numeric answer; exactly the 21→42→49 staging we observed.
  • Synthetic multi-hop templates (the spider family) — controllable hop count for a difficulty curriculum, and known intermediates so the J-lens can check whether looping sharpens them. Optionally real 2-hop QA (2WikiMultiHop / StrategyQA) for naturalness.
  • Avoid: raw pretraining text (no reasoning signal → re-teaches preservation), pure 1-pass self-distillation (target = no-loop answer, loop can't beat it), bigger-model distillation (conflates "loop helps" with "bigger helps").

The design point that makes or breaks it

Force the loop to be used: train across variable k with a difficulty→depth curriculum (easy items solvable at k=1, hard items only at k>1). Otherwise the adapter learns to solve everything at k=1 and ignores the recurrence — the "limited evidence of latent CoT" failure the recurrent-depth interpretability papers report.

Evaluation (closes the loop with our tooling)

  • Accuracy vs. k on held-out two-hop templates + GSM8K test. Success = the curve rises with k — the thing that flatly did not happen untrained.
  • J-lens sharpening: does the intermediate concept (spider, arithmetic partial) grow across loops after training, where it only fluctuated before?

4. Results of the trained run

The §3 experiment was run (July 2026, ~25 min on the GB10): 1,024 GSM8K train items labeled by the frozen model (direct pass vs CoT pass → 165 easy / 416 hard / 443 dropped as unreachable; baselines 16% direct, 53% CoT), then 800 steps of adapter-only training (1.6M params, batch 8, CE on answer tokens) with the curriculum k=1 easy / k=2 mixed / k=4 hard-only. Eval: greedy decode through the looped forward on 256 held-out test items (frozen-model baselines there: 11.3% direct, 55.9% CoT), trained adapter vs the untrained α-merge.

accuracy vs k and J-lens sharpening

k untrained all trained all untrained hard trained hard
0 .117 .117 .008 .008
1 .082 .039 .016 .016
2 .094 .090 .039 .063
4 .090 .082 .047 .063
8 .086 .082 .039 .055

What worked (the pre-registered J-lens criterion). After training, loops sharpen the latent concept instead of merely holding it: P('spider') under the lens at L30 rises 0.015 → 0.107 → 0.130 across k=0→4, ~8× the untrained control (which stays at 0.0130.033, as in probe 2d), with the answer intact (P(8)≈1.0) and the recurrence stable through k=8. Training the merge changed the loop from a preserver into an amplifier.

What partially worked. On the hard (CoT-only) bucket, accuracy rises with depth: 0.8% at k=0 → 6.3% at k=24 trained, vs 3.94.7% untrained. Depth buys a small number of answers the model cannot produce in one pass — but it is 8/127 items, not a CoT replacement.

What failed. Overall accuracy does not rise with k: pushing free-running generation through the loop damages easy items (97% → 2135% trained, → 4152% even untrained), which swamps the hard-bucket gain; k=0 remains best overall. The teacher-forced val looked much better than free-run decode — an exposure-bias gap: training only ever saw gold answer tokens, generation feeds back its own. Note the damage is mostly there in the untrained loop path too, so it is a property of the substrate, not something training introduced.

Obvious next steps. (1) Free-run/scheduled-sampling training to close the exposure-bias gap; (2) a k=0-preservation loss on easy items so looping never costs known answers; (3) per-token adaptive depth (MoR-style) instead of uniform k; (4) more data at grade-school difficulty (Orca-Math) and a contamination-free eval with a difficulty dial (GSM-Symbolic variants mapped to target k); (5) bigger adapter / more steps — 800 steps on 517 items is a first probe, not a training run.


5. Latent planning on code (MBPP) — first overall win

Transposing to coding with one structural change: the loop applies only to the prompt span (loop_mask) — the model "plans silently" while reading the problem, then writes code through the plain path (generated tokens never loop, so the §4 exposure-bias failure is absent by construction). Since causality makes the looped prompt states constant across token steps, they are computed once and frozen; generation then runs at native speed with a KV cache (generate_frozen_prompt, verified bit-identical to the recomputing path).

Data: MBPP (974 basic Python problems, 3 asserts each). Frozen-model labeling: direct pass@1 52-57%; plan-first 63% reachable (caution: with a 380-token budget the plan pass scored 17% — "planning hurts" was purely truncation; at 700 tokens planning adds ~11 points). Labels: ~260 easy / 66 hard (plan-only) / rest drop. Training: same recipe, supervision = the model's own verified passing code, CE on code tokens, curriculum k=1 easy / 2 mixed / 4 hard; the 38-item hard pool overfits past ~step 300 (evaluated the step-399 snapshot).

MBPP pass@1 vs k

k untrained all trained all untrained hard trained hard trained easy
0 .488 .488 .036 .036 .984
1 .472 .488 .179 .429 .852
2 .476 .520 .179 .357 .926
4 .472 .512 .143 .464 .885
  • Overall pass@1 beats the no-loop baseline (52.0% vs 48.8% at k=2) — the criterion the GSM8K run failed. The untrained loop never does (≤47.6%).
  • Planning-dependent problems: 3.6% → 46.4% at k=4 (26/56 items; McNemar p<1e-5). Silent loops recover ~46% of what explicit written planning achieves. Training matters: the untrained substrate reaches only 14-18%.
  • The easy-item dip (98→89%) is substrate-inherent (untrained dips equally); training partially repairs it (92.6% vs 88.5% at k=2). Fix on the board: a per-prompt gate (probe on the k=0 workspace state predicting "will looping help?" — the STaR labels are its free supervision; a perfect gate scores ~86% on this mix).

Rung-2 design note: depth-graded band unfreezing

When the band itself is unfrozen (per-iteration LoRA), unfreeze entrance-faded: full trainability at L14 decaying to frozen by ~L22 (e.g. LoRA α × max(0, 1(ℓ−14)/8)). Rationale: the only novel inputs in the system are what the first band layers see (merged fed-back states); the band's exit distribution must stay on-manifold for the frozen motor/suffix consumers, and the k-fold application of the band amplifies any trained change — a depth-decaying profile is a trust region. Risk: if sharpening is implemented by late-band broadcasts, capping their plasticity caps gains. Ablation triplet (shares everything but the LoRA mask): uniform / entrance-faded / exit-faded (falsification control — parity would mean placement is irrelevant and only capacity matters). Composes with, but is run before, per-iteration strength fading (Relaxed-Recursive axis).

6. Loop dynamics & cross-token carry probes

Dynamics (results-loop/loop_dynamics.png): the trained within-token loop is a true fixed-point iteration — a large first step (cos(s₁,s₀)=0.926 vs 0.977 untrained), then bit-exact convergence by k≈3-4 (cos=1.000), norm settled ~1.7% above the untrained level. Convergence depth coincides with where accuracy and J-lens sharpening plateau; k>4 is a literal no-op.

Cross-token carry (probe_carry.py): instead of k loops per token, carry the band output across token steps — one band pass per token seeded with the previous step's state (the amortized loop; cf. Feedback Transformer). Untrained: stable and coherent over 60 free-running tokens (all three seeding modes — last-position / mean / recency-weighted — with mean/EMA marginally cleaner), only a mild norm drift (+12%/60 tokens). With the §4-trained adapter: the 8× sharpening transfers to the carried regime on short answers (P('spider') 0.015→0.12 in ~2 token steps) but long free-running text collapses — the adapter never saw its own continuations. Conclusion: the carry loop must be trained as the carry loop (sequential unroll on its own trajectory; planned objective: fixed-point distillation — train one carried band pass to land where the within-token k-loop converges — plus STaR rollouts).


Status. Probes 2a2d, the trained-merge experiments (§35), the dynamics analysis and carry probes (§6) are run and reproducible on E2B: scripts/loop_common.py (band re-run machinery, bit-exact; loop_mask; generate_frozen_prompt fast path), prep_star_data.py / prep_mbpp.py (+ prep_mbpp_fix.py), train_merge.py / train_merge_code.py, eval_loop.py / eval_loop_code.py, probe_carry.py, plotting scripts; data, adapters and eval JSONs in results-loop/. Toward a paper: needs seeds/full test sets + CIs, the pause-token (equal-FLOPs) baseline, a second model scale (12B), a contamination-free eval (GSM-Symbolic), and the per-prompt gate; the cross-token carry (§6) is the successor experiment.

Related work. Universal Transformers (adaptive-depth recurrence); Deep Equilibrium Models (solve for the fixed point directly); Geiping et al. recurrent-depth latent reasoning / "Huginn" (prelude→looped core→coda with input injection); Mixture-of-Recursions (per-token learned recursion depth); Relaxed Recursive Transformers (layer-tying + per-loop LoRA); Coconut (latent chain-of-thought). The novel piece here is using an interpretability signal (the J-lens) both to choose the band to loop and to measure whether looping deepens computation, and training only a merge adapter on top of a frozen pretrained model.