Latent-CoT Graph Reachability — Experiment Log
Persistent working notes so progress/commands survive chat resets. Newest status at the top of "Current Status"; details and history below. Update this file whenever we decide/do something.
Repo: reasoning-by-superposition-main (NeurIPS 2025 "Reasoning by Superposition",
authors' official code). Env: conda env superposition (torch 2.5.1+cu121,
transformers 4.46.2). Node: slimgpu (8× RTX A5000).
Goal
Small from-scratch GPT-2 (symbol, 2-layer / 8-head / 768-dim) trained with Coconut continuous latent CoT on a 2-arm star-graph reachability task, where the token vocabulary is the set of graph node ids (node id == token id). This makes the LM head a "node dictionary": applying it (logit lens) to each latent thought reads out which graph node that thought points to, so we can interpret each latent as one BFS hop of reasoning depth. Then: scale depth (L6 -> L20) and study backtracking (retrain earlier latent stages when they regress).
Task detail: two disjoint star components, each with root + 2 arms of length L.
One component is reachable from the query root, the other is a distractor. Model
must answer which candidate leaf is reachable. bfs_variant: True => the "frontier"
at hop k is the two reachable-arm nodes at depth k. Key metric = per-hop
frontier accuracy (does latent k decode to a correct depth-k reachable node).
Current Status (2026-07-21)
- Depth-6 (L6): DONE, success. Per-hop
frontieraccuracy 0.98-1.0 across all 6 hops. Logit-lens probe (probe_latents.py) confirms each of the 6 latents decodes to a correct node one BFS hop deeper on the reachable arm, never the distractor. Checkpoints inckpts/star-coconut-L6-bfs/. Log:logs/star_coconut_L6_bfs.log. - Depth-20 (L20), epoch-scheduled version: FAILED to learn. Config
args/star_coconut_L20_bfs.yaml. The curriculum advanced the frontier on a fixed epoch timer (scheduled_stage = epoch // epochs_per_stage) and reached frontier 14, but per-hopfrontieracc stayed ~0.02-0.04 at EVERY hop (incl. hop 1). i.e. training compute got smeared across stages that were never mastered. This run was killed. - Depth-20, ACCURACY-GATED version: the current approach. Config
args/star_coconut_L20_bfs_accstage.yaml. Promotion is now driven by measured accuracy, not the epoch counter (see "Accuracy-gated curriculum" below).
2026-07-21 evening: root-caused the L20 failure = TRAINING DIVERGENCE
- The accuracy-gated run correctly HELD at stage 1 for 65 epochs (~1h) because hop-1 acc never cleared 0.9 -- but the real issue is that eval loss was monotonically INCREASING: ~4.62 -> 5.9+, while random loss for the 128-vocab is ln(128) ~= 4.85. So the model started ~chance and diverged past chance. hop-1 acc ~0.02 == uniform random over ~100 nodes. The old epoch-scheduled L20 run had the same near-random accuracy => this is an L20 training-instability problem, NOT a data or curriculum problem.
- Data verified well-formed (root/target/neighbor_k valid, node ids <= 99, 80 edges,
20 steps). Difference from working L6: ~330-token sequences (80 edges) vs ~6 edges,
wider vocab,
lr=1e-4, and NO gradient clipping. Classic long-seq divergence.
Fixes applied (this session)
- Stabilization in
run.py+ config: gradient clipping (grad_clip: 1.0), lowerlr1e-4 -> 3e-5, linear LRwarmup_steps: 200. - Gate off-by-one fix: metric "hop k" uses (k-1) latents to predict a depth-k node, so "stage s learned" == per-hop acc at hop (s+1). The promotion gate now requires min acc over hops 1..(cur_stage+1) >= threshold (previously 1..cur_stage, which never checked the current stage's own target).
- Limited backprop:
backprop_depth: 1(was 6) -- gradient only through the newest latent step + answer, matching the incremental curriculum. Simple, not full BPTT across stages.
- Status: relaunch
star_coconut_L20_bfs_accstage.yamlwith these fixes.
2026-07-21 late: first stabilization attempt still diverged -> FSDP clip bug
- With lr 3e-5 + grad_clip 1.0 + warmup, eval loss STILL climbed (4.62 -> 5.2 over 80 epochs) and stage 1 stayed at chance (HOLD 80 epochs). So clipping wasn't actually working.
- Root cause: the model is wrapped in FSDP (params sharded), but the clip used
torch.nn.utils.clip_grad_norm_(parallel_model.parameters(), ...), which computes the norm over only each rank's local shard -> under-counts -> effectively no clip. - Fix: use
parallel_model.clip_grad_norm_()when the model is FSDP (all-reduces the global norm). Also loweredlr3e-5 -> 1e-5. - If loss STILL rises after this, the remaining suspect is model capacity: a 2-layer GPT-2 doing multi-hop pointer chasing over ~80 shuffled edges (vs ~6 at L6). Next step would be a controlled L6-with-same-code sanity run, then consider more layers.
Accuracy-gated curriculum + backtracking (what we changed and why)
Motivation (K's guidance): stage promotion should be accuracy-dependent, not iteration-dependent. Only advance to stage k+1 once every stage 1..k has reached a target accuracy. And keep the sudoku-style backtracking: at each eval check all previous stages; if any falls below threshold, go back and rehearse/retrain it.
Implemented in run.py:
- New config flag
accuracy_staging: True. When on,scheduled_stage = cur_stage, a persistent frontier that only advances when earned. - After each per-hop eval: promote
cur_stage -> cur_stage+1only whenmin(frontier_acc[1..cur_stage]) >= promote_threshold(default 0.9). Because it checks all stages <= frontier, a regression anywhere blocks promotion. - Logs per-stage timing:
[acc-stage] PROMOTE stage k -> k+1 | solved in N epochs / Ts.HOLDlines mean the current frontier isn't solved yet. - Backtracking rehearsal preserved:
bt_r_current= earliest hop belowbacktrack_detect_threshold; the training-set builderget_graph_latent_cot_dataset_backtracksamples each example's stage from a rehearsal distribution (broad w.p.remember_rate, else targeted[r_current..frontier]). perhop_val_samplescaps val per-hop eval cost so frequent eval stays cheap.
Earlier enabling fixes (already in the code):
- Vectorized latent feedback in
coconut.py(clone+scatter instead of a Python double-loop of ~11k tiny GPU ops). Proven numerically identical. - O(batch^2 x depth) sync bug in
coconut.pylatent_listsconstruction (per-element.item()GPU->CPU syncs) replaced with one.tolist()+ pure-Python grouping. This was the real speed killer; latency went from 7357ms -> 561ms/it at stage 20 (~13x), roughly flat across depth. - Truncated BPTT
backprop_depth: Wincoconut.py: detaches fed-back thought- KV older than W latent steps => bounded backward memory/compute regardless of
depth. Default None = full BPTT (L6 baseline unchanged). L20 uses
backprop_depth: 6.
- KV older than W latent steps => bounded backward memory/compute regardless of
depth. Default None = full BPTT (L6 baseline unchanged). L20 uses
- Widened node vocab to 100 (
stokenizer.pyNUM_NODES=100), modelconfigs/symbol-2layer-8head-768dim-L20.jsonvocab_size=128, data regenerated at L=20 (data/star_2arm_L20_*_fo_bfs.json, 82 nodes/sample, 80 edges).
Curriculum semantics (what each stage learns)
From expand_data: at stage s the prompt has s latent tokens and the model
is trained to output the next hop's frontier node (depth s+1) -- NOT the final
leaf. Only the final stage (k = max_steps+1, all latents + [A]) predicts the
target leaf. So each intermediate stage teaches one more hop of the walk:
- stage 0: 0 latents -> predict a depth-1 node (root's neighbor)
- stage 1: 1 latent -> predict a depth-2 node
- stage s: s latents -> predict a depth-(s+1) node
- final: all latents + [A] -> predict the target leaf Per-hop metric "hop k" = (k-1) latents -> predict depth-k node, i.e. hop k tests training stage (k-1). (This is why the promotion gate checks hop cur_stage+1.)
Interpretability / probing (how latent -> node works)
- A latent "thought" is the previous position's last-layer hidden state h in R^768, fed back in place of a token embedding (never discretized).
- Logit lens: node_hat_k = argmax over node-id columns of (W_U h), where W_U is the model's own tied unembedding. Since node id == token id, this reads the node.
probe_latents.py: one teacher-forced validation forward pass over held-out graphs; reads the head at each latent-feeder position; scores set-membership in the hop-k reachable frontier vs the distractor arm.
Commands
Activate env:
source ~/miniforge3/etc/profile.d/conda.sh && conda activate superposition
Launch accuracy-gated L20 (single line; GPUs 2,3):
source ~/miniforge3/etc/profile.d/conda.sh && conda activate superposition && cd /egr/research-slim/ghoshavr/reasoning-by-superposition-main && CUDA_VISIBLE_DEVICES=2,3 WANDB_MODE=offline nohup torchrun --standalone --nnodes 1 --nproc_per_node 2 run.py args/star_coconut_L20_bfs_accstage.yaml > logs/star_coconut_L20_bfs_accstage.log 2>&1 & echo "PID=$!"
Kill a run:
pkill -9 -f 'star_coconut_L20_bfs_accstage.yaml'
Monitor:
grep -E 'acc-stage|scheduled_stage|eval per-hop|Accuracy on validation' logs/star_coconut_L20_bfs_accstage.log | tail -n 40
Probe L6 latents (example):
CUDA_VISIBLE_DEVICES=3 python probe_latents.py # see script for args
Known issues / gotchas
- GPU 2 has hung on CUDA init in this session (bad state). Probing all GPUs or
nvidia-smican hang the shell. Always pinCUDA_VISIBLE_DEVICESto a known-good device and wrap ad-hoc probes intimeout. (Note: the L20 training itself ran on 2,3 fine, so 2 may be intermittently OK.) - Cursor agent shell repeatedly wedges behind an unkillable CUDA-init; a window reload did not always free the agent's worker. Workaround: run launch/kill commands in a normal terminal; the agent can still read logs/files directly.
- Multi-line paste hazard: pasted multi-line blocks lost their newlines and ran
glued together (
superpositioncd ...). Use the single-line command forms above. - Resume-logic crash (fixed):
run.pytreated any non-emptyckpts/<name>/dir as a resume and crashed when nocheckpoint_*file existed ('NoneType'.split). Fixed to only resume when realcheckpoint_*files exist.
Next steps
- Confirm accuracy-gated L20 starts cleanly and learns hop 1 (stage 1 -> >=0.9). This is the real test: with all early compute focused on stage 1, does a single hop become learnable (unlike the smeared epoch-scheduled run)?
- Record stage-by-stage solve time (epochs + wall-clock) from
[acc-stage]logs. - Watch for backtracking
HOLDs (earlier stage regressed -> rehearsed). - If hop 1 still won't train, debug that specifically (LR, backprop_depth, data/vocab) rather than advancing the frontier. Consider an L10 bisection run.
- Re-run
probe_latents.pyon the best L20 checkpoint to see how deep the latents stay faithful.