Files
jspace/WORKSPACE_LOOPING.md
T

328 lines
17 KiB
Markdown
Raw 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.
# 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`](RESULTS.md); the
schematic is [`results/jlens_diagram.png`](results/jlens_diagram.png); a
distilled worked/failed/pitfalls ledger is in [`LESSONS.md`](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](results-loop/loop_eval.png)
| 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](results-loop/loop_eval_code.png)
| 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
**Optimizer note (from McLeish et al., arXiv:2511.07384):** use **Muon**, not
AdamW, when unfreezing the band — they report AdamW loss-spikes to NaN when
training through recurrence; adapter-only training didn't hit this, band-LoRA
likely will. Also adopt their FLOPs convention for recurrent models
(FLOPs = (6N₁+2N₂)D, N₁ = params with grads, N₂ = forward-only) in any
compute-matched comparison. **Adaptive-α candidate (from Lys et al.,
arXiv:2602.14759):** their training-free softmax auto-alignment interpolation
is a cheap ablation against fixed α — directly relevant to the 12B α
miscalibration.
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`](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).
---
## 7. Status & related work
**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"](https://arxiv.org/abs/2502.05171)
(prelude→looped core→coda with input injection); [Mixture-of-Recursions](https://arxiv.org/abs/2507.10524)
(per-token learned recursion depth); [Relaxed Recursive Transformers](https://arxiv.org/abs/2410.20672)
(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.