Avra98 commited on
Commit
8f46582
·
verified ·
1 Parent(s): 34a8c0a

Add training code (same as GitHub reasoning-by-superposition-latent)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +5 -0
  2. .gitignore +38 -0
  3. EXPERIMENT_LOG.md +206 -0
  4. LICENSE +21 -0
  5. NOTES_curriculum_principles.md +79 -0
  6. ONBOARDING.md +294 -0
  7. README.md +152 -0
  8. args/L10_w1_prom098_bt098.yaml +58 -0
  9. args/L10_w1_prom099_bt099.yaml +56 -0
  10. args/L10_w2_prom095_bt095.yaml +56 -0
  11. args/L10_w2_prom098_bt098.yaml +56 -0
  12. args/L15_push_2L_ce95_100k.yaml +56 -0
  13. args/L15_push_2L_ce95_100k_w1.yaml +59 -0
  14. args/L15_push_2L_ce95_100k_w1_prom098.yaml +59 -0
  15. args/L15_push_2L_ce95_50k.yaml +56 -0
  16. args/L15_push_4L_ce90_50k.yaml +56 -0
  17. args/L15_push_4L_ce95_50k.yaml +56 -0
  18. args/L15_push_s0_4L.yaml +56 -0
  19. args/L15_w1_prom098_bt098.yaml +56 -0
  20. args/L15_w1_prom099_bt099.yaml +56 -0
  21. args/L15_w2_prom095_bt095.yaml +57 -0
  22. args/L15_w2_s1_prom095_bt095.yaml +57 -0
  23. args/L15_w5_prom095_bt095.yaml +57 -0
  24. args/L15_w5_prom098_bt098.yaml +56 -0
  25. args/L15_w5_s1_prom095_bt095.yaml +57 -0
  26. args/L20_cso_prom090.yaml +59 -0
  27. args/L20_w20_s1_prom090_bt090.yaml +58 -0
  28. args/L20_w2_s1_prom090_bt090.yaml +58 -0
  29. args/L20_w5_s1_prom090_bt090.yaml +57 -0
  30. args/diag_L15_ce_thr050.yaml +55 -0
  31. args/diag_L15_ce_thr090.yaml +53 -0
  32. args/diag_L15_ce_thr095.yaml +53 -0
  33. args/diag_L15_frontier_nobt_095.yaml +54 -0
  34. args/diag_L15_frontier_thr085.yaml +53 -0
  35. args/diag_L15_frontier_thr090.yaml +53 -0
  36. args/diag_L15_frontier_thr095.yaml +55 -0
  37. args/diag_L15_frontier_thr099.yaml +55 -0
  38. args/diag_L15_promCE090_btCE090.yaml +58 -0
  39. args/diag_L15_promF095_btCE050.yaml +58 -0
  40. args/diag_L15_promF095_btCE090.yaml +58 -0
  41. args/diag_L15_promF095_btCE095.yaml +58 -0
  42. args/diag_L15_promF095_btF095.yaml +58 -0
  43. args/diag_L15_promF095_btNONE.yaml +58 -0
  44. args/diag_L15_promF099_btCE050.yaml +58 -0
  45. args/diag_L15_promF099_btCE090.yaml +58 -0
  46. args/diag_ctrl.yaml +36 -0
  47. args/diag_lr1e3.yaml +36 -0
  48. args/diag_lr3e4.yaml +36 -0
  49. args/diag_small.yaml +36 -0
  50. args/diag_smallL.yaml +36 -0
.gitattributes CHANGED
@@ -33,3 +33,8 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ figs/interventions/L20_all_methods.png filter=lfs diff=lfs merge=lfs -text
37
+ overleaf_latent_backtracking/figs/confusion_heatmaps.png filter=lfs diff=lfs merge=lfs -text
38
+ overleaf_latent_backtracking/figs/curriculum_analysis.png filter=lfs diff=lfs merge=lfs -text
39
+ overleaf_latent_backtracking/figs/panelA_stagewise_learning.png filter=lfs diff=lfs merge=lfs -text
40
+ overleaf_latent_backtracking/figs/panelE_superposition.png filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ __pycache__/
2
+ *.py[cod]
3
+ *.egg-info/
4
+ .venv/
5
+ wandb/
6
+ # Ignore local checkpoints, but ship known-good L10 stage-0 init
7
+ ckpts/**
8
+ !ckpts/
9
+ !ckpts/star-coconut-L10-bfs-stage0/
10
+ ckpts/star-coconut-L10-bfs-stage0/**
11
+ !ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
12
+ **/ckpts/**
13
+ !**/ckpts/
14
+ !**/ckpts/star-coconut-L10-bfs-stage0/
15
+ **/ckpts/star-coconut-L10-bfs-stage0/**
16
+ !**/ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
17
+ logs/
18
+ **/logs/
19
+ data/
20
+ *.out
21
+ *.err
22
+ *.pt
23
+ *.bin
24
+ *.safetensors
25
+ .DS_Store
26
+ .ipynb_checkpoints/
27
+ # large local artifacts
28
+ *.pdf
29
+ Curriculum_Learning.pdf
30
+ latent_backtracking_overleaf.zip
31
+ run.py.bak
32
+ # keep code figs small optional — ignore bulky plots by default
33
+ figs/**/*.png
34
+ # exception: intervention figures linked from the README
35
+ !figs/interventions/*.png
36
+ stage_q_candidates/ckpts/
37
+ stage_q_candidates/logs/
38
+
EXPERIMENT_LOG.md ADDED
@@ -0,0 +1,206 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Latent-CoT Graph Reachability — Experiment Log
2
+
3
+ Persistent working notes so progress/commands survive chat resets. Newest status
4
+ at the top of "Current Status"; details and history below. Update this file
5
+ whenever we decide/do something.
6
+
7
+ Repo: `reasoning-by-superposition-main` (NeurIPS 2025 "Reasoning by Superposition",
8
+ authors' official code). Env: conda env `superposition` (torch 2.5.1+cu121,
9
+ transformers 4.46.2). Node: `slimgpu` (8× RTX A5000).
10
+
11
+ ---
12
+
13
+ ## Goal
14
+
15
+ Small from-scratch GPT-2 (symbol, 2-layer / 8-head / 768-dim) trained with Coconut
16
+ continuous latent CoT on a 2-arm **star-graph reachability** task, where the token
17
+ vocabulary **is the set of graph node ids** (node id == token id). This makes the
18
+ LM head a "node dictionary": applying it (logit lens) to each latent thought reads
19
+ out *which graph node that thought points to*, so we can interpret each latent as
20
+ one BFS hop of reasoning depth. Then: scale depth (L6 -> L20) and study
21
+ backtracking (retrain earlier latent stages when they regress).
22
+
23
+ Task detail: two disjoint star components, each with root + 2 arms of length L.
24
+ One component is reachable from the query root, the other is a distractor. Model
25
+ must answer which candidate leaf is reachable. `bfs_variant: True` => the "frontier"
26
+ at hop k is the two reachable-arm nodes at depth k. Key metric = per-hop
27
+ `frontier` accuracy (does latent k decode to a correct depth-k reachable node).
28
+
29
+ ---
30
+
31
+ ## Current Status (2026-07-21)
32
+
33
+ - **Depth-6 (L6): DONE, success.** Per-hop `frontier` accuracy 0.98-1.0 across all
34
+ 6 hops. Logit-lens probe (`probe_latents.py`) confirms each of the 6 latents
35
+ decodes to a correct node one BFS hop deeper on the reachable arm, never the
36
+ distractor. Checkpoints in `ckpts/star-coconut-L6-bfs/`. Log:
37
+ `logs/star_coconut_L6_bfs.log`.
38
+ - **Depth-20 (L20), epoch-scheduled version: FAILED to learn.** Config
39
+ `args/star_coconut_L20_bfs.yaml`. The curriculum advanced the frontier on a fixed
40
+ epoch timer (`scheduled_stage = epoch // epochs_per_stage`) and reached frontier
41
+ 14, but per-hop `frontier` acc stayed ~0.02-0.04 at EVERY hop (incl. hop 1).
42
+ i.e. training compute got smeared across stages that were never mastered.
43
+ **This run was killed.**
44
+ - **Depth-20, ACCURACY-GATED version: the current approach.** Config
45
+ `args/star_coconut_L20_bfs_accstage.yaml`. Promotion is now driven by measured
46
+ accuracy, not the epoch counter (see "Accuracy-gated curriculum" below).
47
+
48
+ ### 2026-07-21 evening: root-caused the L20 failure = TRAINING DIVERGENCE
49
+ - The accuracy-gated run correctly HELD at stage 1 for 65 epochs (~1h) because
50
+ hop-1 acc never cleared 0.9 -- but the real issue is that **eval loss was
51
+ monotonically INCREASING**: ~4.62 -> 5.9+, while random loss for the 128-vocab is
52
+ ln(128) ~= 4.85. So the model started ~chance and diverged past chance. hop-1 acc
53
+ ~0.02 == uniform random over ~100 nodes. The old epoch-scheduled L20 run had the
54
+ same near-random accuracy => this is an L20 training-instability problem, NOT a
55
+ data or curriculum problem.
56
+ - Data verified well-formed (root/target/neighbor_k valid, node ids <= 99, 80 edges,
57
+ 20 steps). Difference from working L6: ~330-token sequences (80 edges) vs ~6 edges,
58
+ wider vocab, `lr=1e-4`, and NO gradient clipping. Classic long-seq divergence.
59
+
60
+ ### Fixes applied (this session)
61
+ 1. **Stabilization** in `run.py` + config: gradient clipping (`grad_clip: 1.0`),
62
+ lower `lr` 1e-4 -> 3e-5, linear LR `warmup_steps: 200`.
63
+ 2. **Gate off-by-one fix**: metric "hop k" uses (k-1) latents to predict a depth-k
64
+ node, so "stage s learned" == per-hop acc at hop (s+1). The promotion gate now
65
+ requires min acc over hops 1..(cur_stage+1) >= threshold (previously 1..cur_stage,
66
+ which never checked the current stage's own target).
67
+ 3. **Limited backprop**: `backprop_depth: 1` (was 6) -- gradient only through the
68
+ newest latent step + answer, matching the incremental curriculum. Simple, not
69
+ full BPTT across stages.
70
+ - Status: relaunch `star_coconut_L20_bfs_accstage.yaml` with these fixes.
71
+
72
+ ### 2026-07-21 late: first stabilization attempt still diverged -> FSDP clip bug
73
+ - With lr 3e-5 + grad_clip 1.0 + warmup, eval loss STILL climbed (4.62 -> 5.2 over
74
+ 80 epochs) and stage 1 stayed at chance (HOLD 80 epochs). So clipping wasn't
75
+ actually working.
76
+ - Root cause: the model is wrapped in **FSDP** (params sharded), but the clip used
77
+ `torch.nn.utils.clip_grad_norm_(parallel_model.parameters(), ...)`, which computes
78
+ the norm over only each rank's local shard -> under-counts -> effectively no clip.
79
+ - Fix: use `parallel_model.clip_grad_norm_()` when the model is FSDP (all-reduces the
80
+ global norm). Also lowered `lr` 3e-5 -> 1e-5.
81
+ - If loss STILL rises after this, the remaining suspect is model capacity: a 2-layer
82
+ GPT-2 doing multi-hop pointer chasing over ~80 shuffled edges (vs ~6 at L6). Next
83
+ step would be a controlled L6-with-same-code sanity run, then consider more layers.
84
+
85
+ ---
86
+
87
+ ## Accuracy-gated curriculum + backtracking (what we changed and why)
88
+
89
+ Motivation (K's guidance): stage promotion should be **accuracy-dependent**, not
90
+ iteration-dependent. Only advance to stage k+1 once every stage 1..k has reached a
91
+ target accuracy. And keep the sudoku-style backtracking: at each eval check all
92
+ previous stages; if any falls below threshold, go back and rehearse/retrain it.
93
+
94
+ Implemented in `run.py`:
95
+ - New config flag `accuracy_staging: True`. When on, `scheduled_stage = cur_stage`,
96
+ a persistent frontier that only advances when earned.
97
+ - After each per-hop eval: promote `cur_stage -> cur_stage+1` only when
98
+ `min(frontier_acc[1..cur_stage]) >= promote_threshold` (default 0.9). Because it
99
+ checks *all* stages <= frontier, a regression anywhere blocks promotion.
100
+ - Logs per-stage timing: `[acc-stage] PROMOTE stage k -> k+1 | solved in N epochs / Ts`.
101
+ `HOLD` lines mean the current frontier isn't solved yet.
102
+ - Backtracking rehearsal preserved: `bt_r_current` = earliest hop below
103
+ `backtrack_detect_threshold`; the training-set builder
104
+ `get_graph_latent_cot_dataset_backtrack` samples each example's stage from a
105
+ rehearsal distribution (broad w.p. `remember_rate`, else targeted
106
+ `[r_current..frontier]`).
107
+ - `perhop_val_samples` caps val per-hop eval cost so frequent eval stays cheap.
108
+
109
+ Earlier enabling fixes (already in the code):
110
+ - **Vectorized latent feedback** in `coconut.py` (clone+scatter instead of a
111
+ Python double-loop of ~11k tiny GPU ops). Proven numerically identical.
112
+ - **O(batch^2 x depth) sync bug** in `coconut.py` `latent_lists` construction
113
+ (per-element `.item()` GPU->CPU syncs) replaced with one `.tolist()` + pure-Python
114
+ grouping. This was the real speed killer; latency went from 7357ms -> 561ms/it at
115
+ stage 20 (~13x), roughly flat across depth.
116
+ - **Truncated BPTT** `backprop_depth: W` in `coconut.py`: detaches fed-back thought
117
+ + KV older than W latent steps => bounded backward memory/compute regardless of
118
+ depth. Default None = full BPTT (L6 baseline unchanged). L20 uses `backprop_depth: 6`.
119
+ - Widened node vocab to 100 (`stokenizer.py` `NUM_NODES=100`), model
120
+ `configs/symbol-2layer-8head-768dim-L20.json` `vocab_size=128`, data regenerated
121
+ at L=20 (`data/star_2arm_L20_*_fo_bfs.json`, 82 nodes/sample, 80 edges).
122
+
123
+ ---
124
+
125
+ ## Curriculum semantics (what each stage learns)
126
+
127
+ From `expand_data`: at **stage s** the prompt has **s latent tokens** and the model
128
+ is trained to output the **next hop's frontier node (depth s+1)** -- NOT the final
129
+ leaf. Only the final stage (`k = max_steps+1`, all latents + `[A]`) predicts the
130
+ target leaf. So each intermediate stage teaches one more hop of the walk:
131
+ - stage 0: 0 latents -> predict a depth-1 node (root's neighbor)
132
+ - stage 1: 1 latent -> predict a depth-2 node
133
+ - stage s: s latents -> predict a depth-(s+1) node
134
+ - final: all latents + [A] -> predict the target leaf
135
+ Per-hop metric "hop k" = (k-1) latents -> predict depth-k node, i.e. hop k tests
136
+ training stage (k-1). (This is why the promotion gate checks hop cur_stage+1.)
137
+
138
+ ## Interpretability / probing (how latent -> node works)
139
+
140
+ - A latent "thought" is the previous position's last-layer hidden state
141
+ h in R^768, fed back in place of a token embedding (never discretized).
142
+ - Logit lens: node_hat_k = argmax over node-id columns of (W_U h), where W_U is the
143
+ model's own tied unembedding. Since node id == token id, this reads the node.
144
+ - `probe_latents.py`: one teacher-forced validation forward pass over held-out
145
+ graphs; reads the head at each latent-feeder position; scores set-membership in
146
+ the hop-k reachable frontier vs the distractor arm.
147
+
148
+ ---
149
+
150
+ ## Commands
151
+
152
+ Activate env:
153
+ ```
154
+ source ~/miniforge3/etc/profile.d/conda.sh && conda activate superposition
155
+ ```
156
+
157
+ Launch accuracy-gated L20 (single line; GPUs 2,3):
158
+ ```
159
+ 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=$!"
160
+ ```
161
+
162
+ Kill a run:
163
+ ```
164
+ pkill -9 -f 'star_coconut_L20_bfs_accstage.yaml'
165
+ ```
166
+
167
+ Monitor:
168
+ ```
169
+ grep -E 'acc-stage|scheduled_stage|eval per-hop|Accuracy on validation' logs/star_coconut_L20_bfs_accstage.log | tail -n 40
170
+ ```
171
+
172
+ Probe L6 latents (example):
173
+ ```
174
+ CUDA_VISIBLE_DEVICES=3 python probe_latents.py # see script for args
175
+ ```
176
+
177
+ ---
178
+
179
+ ## Known issues / gotchas
180
+
181
+ - **GPU 2 has hung on CUDA init** in this session (bad state). Probing all GPUs or
182
+ `nvidia-smi` can hang the shell. Always pin `CUDA_VISIBLE_DEVICES` to a known-good
183
+ device and wrap ad-hoc probes in `timeout`. (Note: the L20 training itself ran on
184
+ 2,3 fine, so 2 may be intermittently OK.)
185
+ - **Cursor agent shell repeatedly wedges** behind an unkillable CUDA-init; a window
186
+ reload did not always free the agent's worker. Workaround: run launch/kill
187
+ commands in a normal terminal; the agent can still read logs/files directly.
188
+ - **Multi-line paste hazard:** pasted multi-line blocks lost their newlines and ran
189
+ glued together (`superpositioncd ...`). Use the single-line command forms above.
190
+ - **Resume-logic crash (fixed):** `run.py` treated any non-empty `ckpts/<name>/`
191
+ dir as a resume and crashed when no `checkpoint_*` file existed
192
+ (`'NoneType'.split`). Fixed to only resume when real `checkpoint_*` files exist.
193
+
194
+ ---
195
+
196
+ ## Next steps
197
+
198
+ 1. Confirm accuracy-gated L20 starts cleanly and learns hop 1 (stage 1 -> >=0.9).
199
+ This is the real test: with all early compute focused on stage 1, does a single
200
+ hop become learnable (unlike the smeared epoch-scheduled run)?
201
+ 2. Record stage-by-stage solve time (epochs + wall-clock) from `[acc-stage]` logs.
202
+ 3. Watch for backtracking `HOLD`s (earlier stage regressed -> rehearsed).
203
+ 4. If hop 1 still won't train, debug that specifically (LR, backprop_depth,
204
+ data/vocab) rather than advancing the frontier. Consider an L10 bisection run.
205
+ 5. Re-run `probe_latents.py` on the best L20 checkpoint to see how deep the latents
206
+ stay faithful.
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) Meta Platforms, Inc. and affiliates.
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
NOTES_curriculum_principles.md ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Operating principles for the latent-curriculum runs
2
+
3
+ Working rules for L15/L20 and any deeper/wider graph runs. Follow these before
4
+ reaching for a looser gate or killing a run.
5
+
6
+ ## 1. Tighter gates early pay for themselves later
7
+
8
+ Promote only when the balance score is genuinely met. A tight early gate buys
9
+ cheaper later stages and fewer backtracks; a loose one lets balance rot behind
10
+ the frontier and the debt comes due at depth.
11
+
12
+ Do not loosen a threshold to make a run "progress". Progress bought that way is
13
+ the failure mode we already diagnosed: frontier@0.95 promotion kept advancing
14
+ L15 while ce_score decayed 0.96 -> 0.44 across hops.
15
+
16
+ ## 2. Early stages are slow; later stages get smooth
17
+
18
+ Measured on the L10 ce_score@0.95 arm, the only run that completed all stages:
19
+
20
+ | stage | 1->2 | 2->3 | 3->4 | 4->5 | 5->6 | 6->7 | 7->8 | 8->9 | 9->10 |
21
+ |-------|------|------|------|------|------|------|------|------|-------|
22
+ | epochs| 140 | 45 | 15 | 35 | 10 | 20 | 45 | 20 | 30 |
23
+
24
+ Stage 1 cost 140 epochs, ~3x any later stage. Everything after cleared in 10-45.
25
+
26
+ Consequence: a run sitting at stage 1 or 2 for ~150 epochs is on trajectory, not
27
+ stalled. Judge a run against this curve before concluding anything. The number
28
+ of backtracks should also fall as stages advance -- track it as a health signal,
29
+ since a rising count is the real warning sign.
30
+
31
+ ## 3. Stalling usually means more training, not a smaller threshold
32
+
33
+ Check in this order before touching the gate: how many epochs has this stage
34
+ actually had, relative to the 140-epoch stage-1 reference; is the metric flat or
35
+ still creeping; is backtracking firing at all (a gate that never triggers cannot
36
+ repair anything).
37
+
38
+ ## 4. Stage 0 is not the bottleneck
39
+
40
+ L15 hop-1 saturates at ce_score ~0.96 by epoch 45 and holds 0.94-0.98 for the
41
+ rest of the curriculum. Extra stage-0 epochs buy ~0.016. Do not spend GPU time
42
+ there.
43
+
44
+ ## 5. The point of tightening is depth and branching
45
+
46
+ The payoff is deeper graphs with more branches. A recipe that only works by
47
+ loosening the gate will not extend, so prefer the tight-gate result even when it
48
+ is slower to obtain.
49
+
50
+ ## 6. Once early latents are robust, W=1 BPTT should be enough
51
+
52
+ If stages 0..k-1 are already solid (high ce_score, rarely backtracked), the
53
+ forward chain of latents is carrying the right frontier state, so gradients for
54
+ stage k only need to flow through the *last* latent recurrence step
55
+ (`backprop_depth: 1`). Earlier latents should not need re-learning via full
56
+ BPTT -- that is exactly what the long/robust early training bought.
57
+
58
+ Implication for experiments:
59
+ - Do NOT test W=1 from a cold or weakly-gated start. Our earlier W=1 L15/L20
60
+ runs failed for this reason: early hops were not solid, so truncating BPTT
61
+ cut the only path that could still fix them.
62
+ - DO test W=1 by warm-starting from a checkpoint where hops 1..k already meet
63
+ the tight ce_score gate, then training stage k+1..L with `backprop_depth: 1`.
64
+ - The success criterion is: same promote/backtrack behaviour as full BPTT, with
65
+ fewer backtracks and similar stage-clearing times. If W=1 needs constant
66
+ repair of early hops, early training was not robust enough yet.
67
+
68
+ ## Fixed setup details worth not rediscovering
69
+
70
+ - Mirror the L10 winner: `promote_metric` and `backtrack_metric` both
71
+ `ce_score`, both at 0.95, `init_stage: 1`, warm start from the pinned stage-0
72
+ checkpoint.
73
+ - Run fp32. `bf16: True` adds mantissa noise to a metric that compares
74
+ near-equal log-probabilities between the two arms.
75
+ - Default to full BPTT while building the early stages. Switch to W=1 only as
76
+ a *transfer* test from a mastered early-stage checkpoint (principle 6).
77
+ - Point warm starts at the *latest* stage-0 checkpoint. The diag sweep silently
78
+ used a stale `checkpoint_99` (ce_score 0.947) when `checkpoint_150` (0.961)
79
+ existed.
ONBOARDING.md ADDED
@@ -0,0 +1,294 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Onboarding / Collaboration Guide
2
+
3
+ This repo (`reasoning-by-superposition`) is where the active **continuous-latent reasoning**
4
+ experiments live. This document explains the two codebases in play, exactly what we changed here
5
+ vs. the upstream paper code, how to set things up and run experiments, and the git workflow for
6
+ contributing.
7
+
8
+ If you just want to run something, jump to **[Quick Start](#quick-start)**. If you want to
9
+ understand the experiment design, read **[Concepts](#concepts)** first.
10
+
11
+ ---
12
+
13
+ ## 1. The two codebases
14
+
15
+ There are two related but separate repos. They share conceptual DNA (Coconut-style continuous
16
+ latent chain-of-thought on **directed graph reachability**), but differ in scale, data format, and
17
+ which questions they answer.
18
+
19
+ | Repo | Where | Base model | Task/data | Role |
20
+ |---|---|---|---|---|
21
+ | **`reasoning-by-superposition`** (this repo, "rbs") | Berkeley SCF `/scratch/users/gatmiry/reasoning-by-superposition`; backup at `github.com/seyedparsa/reasoning-by-superposition` (private) | 2-layer, 8-head, 768-dim GPT-2, **randomly initialized** (~15M params) | Symbolic graph reachability, integer node tokens 0–30 (`stokenizer.py`) | **Active work.** Small, fast, from-scratch. Where the 2-arm-star + final-only experiments run. |
22
+ | **`thoughtformer`** | `github.com/seyedparsa/thoughtformer` | Qwen3-0.6B-Base / GPT-2 (pretrained) | Bridge / Tree-n100 / ProsQA, natural-language premises | The parent project. Larger-scale Coconut fork with a full curriculum/staging framework. See its `CLAUDE.md`. |
23
+
24
+ **Lineage.** `thoughtformer` is a fork of the original [Coconut](https://arxiv.org/abs/2412.06769).
25
+ `reasoning-by-superposition` is the *authors'* code for the NeurIPS 2025 paper
26
+ [Reasoning by Superposition](https://arxiv.org/abs/2505.12514) (also a Coconut descendant); we
27
+ cloned it and built our experiments on top. **Our `origin` is the upstream author repo (`Ber666`,
28
+ read-only for us); all our work lives on the `mine` remote** — see [Git workflow](#5-git-workflow).
29
+
30
+ For the rest of this doc, "we/our" = the changes made on top of the upstream `Ber666` code.
31
+
32
+ ---
33
+
34
+ ## 2. What we changed in this repo (rbs)
35
+
36
+ Upstream trains one thing: standard per-hop Coconut on ProsQA. We added a new **task** (2-arm star),
37
+ a new **training variant** (final-only), a new **staging algorithm** (revert), per-depth **metrics**,
38
+ and a pile of configs/scripts. All of it is additive and flag-gated — the upstream ProsQA path is
39
+ byte-identical when the new flags are off.
40
+
41
+ ### Commit history (our 10 commits on top of `origin/main`)
42
+
43
+ ```
44
+ fc56a19 Backup: configs, gen/probe/verify/wb scripts, launchers, stokenizer tweak; gitignore data/
45
+ 12706bd final-only d1 troubleshoot: diagnostic matrix (lr sweep, small model, high wd, base-coconut control)
46
+ ab729fa depth-1 final-only scaling sweep: L=1 data + configs (14k-960k)
47
+ ff52afe add revert-staged final-only variant (gates depth on per-depth acc) + configs
48
+ 282a0f8 final-only: log reachable/frontier/optimal per depth alongside acc (classify [A] output)
49
+ 8c5fbcd add general final-only reachability variant (per-depth [A] training) + 2-arm-star fo data/configs
50
+ 64c9172 add 2-arm star dataset generator + standard/revert configs (L=6)
51
+ 7d06b5f Revert "add hysteresis staging variant ..." <- dead end, reverted
52
+ 9f6a050 add hysteresis staging variant ... <- dead end
53
+ 4f2efd3 wandb: resume same run id across preemptions; key eval metrics on train/epoch
54
+ ```
55
+
56
+ ### File-by-file
57
+
58
+ - **`generate_2arm_star.py`** — *new.* Generates the 2-arm-star dataset. Emits `neg_root` /
59
+ `neg_neighbor_k` (the unreachable component's root and per-depth frontier) alongside the usual
60
+ `root` / `neighbor_k`, in **two flavors**: `_coconut` (standard = shortest-path frontier, one node
61
+ per depth) and `_bfs` (full BFS frontier). Key funcs: `make_star`, `to_sample(star, flavor)`,
62
+ `gen(L, n, rng, seen)`, `main`.
63
+ - **`graph_metrics.py`** — *new.* Per-depth evaluation metrics. `perhop_categorize` classifies an
64
+ intermediate-node prediction into **reachable / frontier / optimal** (via BFS distances from
65
+ root and target). `finalonly_categorize` scores the final `[A]` answer per depth. `_distances`,
66
+ `_bfs`, `_classify`, `category_log_dict` are helpers.
67
+ - **`dataset.py`** — added **`get_graph_finalonly_dataset`** (the final-only builder). The base
68
+ builders (`get_graph_latent_cot_dataset` = standard per-hop, `get_graph_latent_question_dataset`,
69
+ `expand_data`, `MyCollator`, etc.) are upstream and untouched.
70
+ - **`run.py`** — flag-gated dispatch inserts (all `elif getattr(configs, "final_only", False)` /
71
+ `getattr(configs, "revert_staging", False)`):
72
+ - **wandb**: resume the same run id across SLURM preemptions; key eval metrics on `train/epoch`.
73
+ - **`scheduled_stage`** selection: `revert_staging` branch sets it from `revert_next_stage`
74
+ (else the upstream `epoch // epochs_per_stage`).
75
+ - **train + val-loss datasets**: route to `get_graph_finalonly_dataset` when `final_only`.
76
+ - **metrics block**: `final_only` -> `finalonly_categorize` (logs `eval/acc_hop{k}` etc.);
77
+ else upstream `perhop_categorize`. Revert logic: `revert_next_stage = first depth k with
78
+ metric < revert_threshold, minus 1` (metric = `acc` for final-only, `frontier` for per-hop).
79
+ - **`stokenizer.py`** — minor tweak (symbolic tokenizer; nodes 0–30 + special tokens, vocab 40).
80
+ - **`configs/symbol-2layer-8head-768dim.json`** — the from-scratch model definition (GPT-2, 2L/8H/768d).
81
+ (`...128dim.json` is a smaller diagnostic variant.)
82
+ - **`args/*.yaml`** — all the run configs (see [Configs](#4-configs)).
83
+ - **Helper scripts** (not part of the pipeline, but committed for reproducibility): `gen_*.py`
84
+ (data generation for specific sweeps), `probe_*.py` / `verify_*.py` / `show_examples.py` (sanity
85
+ checks on data + model outputs), `wb_*.py` (wandb metric pulls), `run_repro*.sbatch` (SLURM launchers).
86
+
87
+ ### What we changed in `thoughtformer` (the parent repo)
88
+
89
+ `thoughtformer` already had the full framework we're conceptually mirroring here — methods
90
+ (`no_cot` / `cot` / `coconut` / `thoughtformer`), two curriculum axes (`data_curriculum`,
91
+ `thought_curriculum`), and three staging modes (`epoch` / `adaptive` / `two_pointers`). Read its
92
+ **`CLAUDE.md`** for the complete map. Our recent changes there were limited to plotting/analysis
93
+ scripts and result artifacts, not the training core.
94
+
95
+ ---
96
+
97
+ ## 3. Concepts
98
+
99
+ ### The task: 2-arm star (graph reachability)
100
+
101
+ Two **disconnected, mirror-image** star graphs:
102
+ - **Reachable component**: root `R` with a *target arm* (length `L`) and a *decoy arm* (length `L`).
103
+ - **Unreachable component**: neg-root `R2` with `neg_target` (plays target's role) and a decoy.
104
+
105
+ Nodes are integer tokens; a star with arm length `L` has `2 + 4L` nodes. The model reads shuffled
106
+ edge premises + two candidate targets + the start root, and must pick the reachable candidate.
107
+ `difficulty` / depth = hops from root to target.
108
+
109
+ Sequence format (single example):
110
+ ```
111
+ [Q] cand1 cand2 [R] root <|latent|> ... <|latent|> [A] answer
112
+ ```
113
+ with `d` latent tokens at depth `d`.
114
+
115
+ ### Two training variants
116
+
117
+ 1. **Standard per-hop Coconut** (`get_graph_latent_cot_dataset`, upstream). At stage `s`, supervise
118
+ an **intermediate path node** (`neighbor_k[s+1]`) with `s` latents and *no* `[A]`; only the final
119
+ stage trains `[A] -> target`. Two flavors: **standard** (`neighbor_k` = shortest-path frontier,
120
+ unique node) and **BFS** (`bfs_variant: True`, full frontier). This is the setting the paper
121
+ studies; the model learns to traverse the graph in latent space.
122
+
123
+ 2. **Final-only** (`get_graph_finalonly_dataset`, `final_only: True`, *ours*). At depth `d`, form two
124
+ depth-`d` candidates — `reach = choice(neighbor_k[d])`, `neg = choice(neg_neighbor_k[d])` — give
125
+ `d` latents + `[A]`, and train only to output `reach`. This **front-loads the reachability
126
+ decision to every depth** instead of supervising intermediate nodes. Same standard/BFS flavor
127
+ distinction (inherited from the data file). Motivation: isolate whether the model can *decide
128
+ reachability* (a 1-bit selection) vs. *traverse* — we found binary final-only supervision tends
129
+ to **memorize** and is far more data-hungry than per-hop chain supervision, which generalizes.
130
+
131
+ ### Staging (when to add latents / advance depth)
132
+
133
+ - **Fixed** (upstream): `scheduled_stage = epoch // epochs_per_stage`.
134
+ - **Revert** (`revert_staging: True`, *ours*): each eval, `scheduled_stage = (first depth k whose
135
+ metric < revert_threshold) - 1`. This **gates advancement on generalization** — depth only grows
136
+ once the model actually solves the current depth on held-out data. `revert_metric` = `acc` for
137
+ final-only, `frontier` for per-hop.
138
+
139
+ ### Metrics (logged to wandb per depth `k`)
140
+
141
+ - Per-hop runs: `{reachable,frontier,optimal}_hop{k}` — is the predicted intermediate node reachable
142
+ from root / on *a* shortest-path frontier / on the *target's* shortest path, respectively.
143
+ - Final-only runs: `eval/acc_hop{k}` and `train/acc_hop{k}` (per-depth answer accuracy) + `eval/acc`
144
+ (deepest-depth accuracy). `revert/stage` tracks the current scheduled stage.
145
+
146
+ ---
147
+
148
+ ## 4. Configs
149
+
150
+ Configs live in `args/*.yaml`. CLI overrides work: `run.py args/foo.yaml lr=3e-4 num_epochs=200`.
151
+
152
+ Naming: `star_` = 2-arm star; `prosqa_` = upstream ProsQA task; `diag_` = depth-1 diagnostics.
153
+ `_fo` / `finalonly` = final-only variant; `_bfs` = BFS flavor; `_rev` / `revert` = revert staging;
154
+ `d1` / `L2` / `L6` = arm length; trailing number = train-set size.
155
+
156
+ Key fields:
157
+
158
+ | field | meaning |
159
+ |---|---|
160
+ | `coconut: True` | use the Coconut latent wrapper |
161
+ | `c_thought` | latent tokens per reasoning step (`0` = no-cot) |
162
+ | `max_latent_stage` | max depth / latent count (arm length `L`) |
163
+ | `pad_latent_to_max` | cap latent count at `max_latent_stage` |
164
+ | `final_only` | use the final-only training variant (ours) |
165
+ | `bfs_variant` | use full BFS frontier instead of shortest-path |
166
+ | `revert_staging` | gate depth advancement on generalization (ours) |
167
+ | `revert_metric` / `revert_threshold` | metric + bar for revert staging (e.g. `acc` / `0.8`) |
168
+ | `epochs_per_stage` | (fixed staging) epochs before advancing a stage |
169
+ | `uniform_prob` | prob. of sampling a *shallower* depth than scheduled (curriculum mix) |
170
+ | `model_id` | model def, e.g. `configs/symbol-2layer-8head-768dim.json` |
171
+ | `train_path` / `val_path` | data files (regenerate with the gen scripts — see below) |
172
+ | `lr`, `weight_decay`, `batch_size_training`, `num_epochs`, `seed` | standard knobs |
173
+ | `bf16: False` | keep fp32 (bf16 NaNs with ≥3 latents) |
174
+ | `save_only_improve` | only checkpoint on val improvement |
175
+
176
+ Representative starting points: `star_coconut_full.yaml` (standard per-hop, L=6),
177
+ `star_fo_L6_100k_rev.yaml` (final-only + revert, L=6), `star_finalonly_bfs.yaml` (BFS flavor).
178
+
179
+ ---
180
+
181
+ ## Quick Start
182
+
183
+ ### Environment
184
+
185
+ Python 3.12; deps are pinned in `requirements.txt` (torch 2.5.1, transformers 4.46.2, wandb 0.18.7, …).
186
+
187
+ ```bash
188
+ git clone git@github.com:seyedparsa/reasoning-by-superposition.git # the private backup
189
+ cd reasoning-by-superposition
190
+ conda create -n superposition python=3.12 && conda activate superposition
191
+ pip install -r requirements.txt
192
+ wandb login # required before training/eval (or set WANDB_MODE=offline)
193
+ ```
194
+
195
+ On the Berkeley box, the shared env already exists at `/accounts/projects/peter/gatmiry/sup_env`
196
+ (no need to rebuild) — just `export PATH="$ENV/bin:$PATH"`.
197
+
198
+ ### Data is NOT in git — regenerate it
199
+
200
+ `data/` is gitignored (regenerable). After cloning, generate what a config needs. E.g. the L=6
201
+ 2-arm-star final-only data:
202
+
203
+ ```bash
204
+ python generate_2arm_star.py # writes data/star_2arm_L6_{train,valid,test}_fo_{coconut,bfs}.json
205
+ ```
206
+
207
+ The `gen_*.py` scripts generate the specific sweep datasets (L=1 size sweep, L2/L3/L6 100k, etc.);
208
+ open the one matching your config's `train_path`. Check a config's `train_path`/`val_path` to see
209
+ which file it expects, then run the matching generator.
210
+
211
+ ### Train
212
+
213
+ Distributed launch (2 GPUs is the standard setup for the 768-dim model):
214
+
215
+ ```bash
216
+ torchrun --standalone --nproc_per_node 2 run.py args/star_fo_L6_100k_rev.yaml
217
+ # CLI overrides: append e.g. lr=3e-4 num_epochs=200
218
+ ```
219
+
220
+ On SLURM (Berkeley), use a launcher and set the node:
221
+ ```bash
222
+ sbatch --job-name=my-run run_repro_online_horton.sbatch args/star_fo_L6_100k_rev.yaml
223
+ ```
224
+ Launchers pin `--gpus-per-node 2` and the conda env by absolute path. Logs land in `logs/%x_%j.out`.
225
+ Checkpoints go to `ckpts/<name>/` (gitignored). Runs auto-resume from the latest checkpoint on
226
+ preemption (same wandb run id).
227
+
228
+ ### Evaluate
229
+
230
+ Set `only_eval: True` and `load_model_path: <ckpt>` in a config, then launch normally. There are
231
+ `eval_*.yaml` configs (e.g. `eval_prosqa_coconut_2l_8h_768d.yaml`) as templates.
232
+
233
+ ### Sanity-check before a big run
234
+
235
+ ```bash
236
+ python show_examples.py # print tokenized examples for a config's data
237
+ python verify_star.py # structural checks on generated star data
238
+ ```
239
+
240
+ ---
241
+
242
+ ## 5. Git workflow
243
+
244
+ ### Remotes
245
+
246
+ ```
247
+ origin https://github.com/Ber666/reasoning-by-superposition.git # upstream author (READ-ONLY for us)
248
+ mine git@github.com:seyedparsa/reasoning-by-superposition.git # our private backup (push here)
249
+ ```
250
+
251
+ **All our work lives on `mine`.** Never expect to push to `origin`. `mine/main` tracks our
252
+ `main` (upstream history + our 10 commits).
253
+
254
+ ### For a collaborator
255
+
256
+ 1. **Get access.** The `mine` repo is private — the owner (seyedparsa) must add you as a
257
+ collaborator on GitHub (repo Settings → Collaborators). Then clone `mine` (see Quick Start).
258
+ 2. **Branch for your work** rather than committing straight to `main`:
259
+ ```bash
260
+ git checkout -b feature/my-experiment
261
+ # ... work, commit ...
262
+ git push -u origin feature/my-experiment # 'origin' = your clone's default = the mine repo
263
+ ```
264
+ Open a PR into `main` on `github.com/seyedparsa/reasoning-by-superposition`.
265
+ 3. **Don't commit `data/` or `ckpts/`** — both are gitignored (regenerable / large). Commit code,
266
+ configs, and small model defs only.
267
+ 4. **Keeping the Berkeley working copy in sync.** On the shared box the remote is named `mine`:
268
+ ```bash
269
+ cd /scratch/users/gatmiry/reasoning-by-superposition
270
+ git add -A && git commit -m "..." && git push mine main
271
+ ```
272
+
273
+ ### Pulling upstream changes (rare)
274
+
275
+ ```bash
276
+ git fetch origin && git merge origin/main # if Ber666 ever updates
277
+ ```
278
+
279
+ ---
280
+
281
+ ## 6. Gotchas / lessons
282
+
283
+ - **fp32 only** — bf16 produces NaNs in the iterative latent feedback with ≥3 latent tokens.
284
+ Keep `bf16: False`.
285
+ - **Binary final-only memorizes.** Per-hop / chain supervision generalizes at ~14k examples with no
286
+ train/eval gap; binary final-only supervision (train→1.0, eval→chance) is far more data-hungry and
287
+ the data threshold rises steeply with graph size. This is the central finding driving the current
288
+ experiments.
289
+ - **`save_only_improve` + preemption livelock**: if val is stuck at chance, a preempted run can
290
+ resume from an early "best" checkpoint and never progress. Watch for it on flat-val runs.
291
+ - **Berkeley disk**: scratch is quota'd (~20 GB). Data (108 MB) + checkpoints (230 MB+) fill it fast;
292
+ clean concluded `ckpts/` and regenerable `data/` when low.
293
+ - **`data/` is not backed up** — only code + configs are in git. Regenerate from the gen scripts.
294
+ - **Results/metrics** live in wandb (`seyedparsa/coconut` project), not in the repo.
README.md ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Latent CoT curriculum + backtracking (2-arm star)
2
+
3
+ Fork / extension of [Reasoning by Superposition](https://arxiv.org/abs/2505.12514) ([original repo](https://github.com/Ber666/reasoning-by-superposition)).
4
+
5
+ We train [Coconut](https://arxiv.org/abs/2412.06769)-style continuous chain-of-thought on **2-arm star graph reachability** with:
6
+
7
+ - **CE-gated curriculum** over latent depth (promote when per-hop CE score clears a threshold)
8
+ - **Backtracking (BT)** when an earlier hop drops below threshold
9
+ - **Truncated BPTT** via `backprop_depth` (reported recipe: **W=2**)
10
+ - **Latent interventions** that test whether the answer depends on the last thought vs earlier ones
11
+
12
+ Repo: https://github.com/Avra98/reasoning-by-superposition-latent
13
+
14
+ ## Setup
15
+
16
+ ```bash
17
+ git clone https://github.com/Avra98/reasoning-by-superposition-latent.git
18
+ cd reasoning-by-superposition-latent
19
+ conda create -n superposition python=3.12
20
+ conda activate superposition
21
+ pip install -r requirements.txt
22
+ ```
23
+
24
+ ## Reproduce training (L=10 / 15 / 20, `backprop_depth=2`)
25
+
26
+ ### 1. Generate data
27
+
28
+ ```bash
29
+ # L=10 (14k train)
30
+ python generate_2arm_star.py --L 10 --n_train 14000 --n_valid 256 --seed 0
31
+
32
+ # L=15 (100k train)
33
+ python generate_2arm_star.py --L 15 --n_train 100000 --n_valid 256 --seed 0
34
+
35
+ # L=20 (100k train; rename to match the yaml paths)
36
+ python generate_2arm_star.py --L 20 --n_train 100000 --n_valid 256 --seed 0
37
+ mv data/star_2arm_L20_train_fo_bfs.json data/star_2arm_L20_100k_train_fo_bfs.json
38
+ mv data/star_2arm_L20_valid_fo_bfs.json data/star_2arm_L20_100k_valid_fo_bfs.json
39
+ mv data/star_2arm_L20_test_fo_bfs.json data/star_2arm_L20_100k_test_fo_bfs.json
40
+ ```
41
+
42
+ ### 2. Stage-0 warm-starts
43
+
44
+ ```bash
45
+ # L10 stage-0 (cold) → ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
46
+ CUDA_VISIBLE_DEVICES=0 torchrun --standalone --nnodes 1 --nproc_per_node 1 \
47
+ run.py args/star_coconut_L10_bfs_stage0.yaml
48
+
49
+ # L15 stage-0 warm from L10 → ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
50
+ CUDA_VISIBLE_DEVICES=0 torchrun --standalone --nnodes 1 --nproc_per_node 1 \
51
+ run.py args/star_coconut_L15_bfs_stage0_warm.yaml
52
+ ```
53
+
54
+ L20 curriculum warm-starts from the same L15 stage-0 checkpoint.
55
+
56
+ ### 3. Train W=2 + backtracking
57
+
58
+ | Depth | Config | Promote / BT gate | Warm-start |
59
+ |------:|--------|-------------------|------------|
60
+ | L=10 | `args/L10_w2_prom095_bt095.yaml` | CE @ 0.95 | L10 stage-0 |
61
+ | L=15 | `args/L15_w2_s1_prom095_bt095.yaml` | CE @ 0.95 | L15 stage-0 warm |
62
+ | L=20 | `args/L20_w2_s1_prom090_bt090.yaml` | CE @ 0.90 | L15 stage-0 warm |
63
+
64
+ ```bash
65
+ # L=10, backprop_depth=2
66
+ CUDA_VISIBLE_DEVICES=0 torchrun --standalone --nnodes 1 --nproc_per_node 1 \
67
+ --master_port 29510 \
68
+ run.py args/L10_w2_prom095_bt095.yaml
69
+
70
+ # L=15, backprop_depth=2
71
+ CUDA_VISIBLE_DEVICES=0 torchrun --standalone --nnodes 1 --nproc_per_node 1 \
72
+ --master_port 29515 \
73
+ run.py args/L15_w2_s1_prom095_bt095.yaml
74
+
75
+ # L=20, backprop_depth=2
76
+ CUDA_VISIBLE_DEVICES=0 torchrun --standalone --nnodes 1 --nproc_per_node 1 \
77
+ --master_port 29520 \
78
+ run.py args/L20_w2_s1_prom090_bt090.yaml
79
+ ```
80
+
81
+ Checkpoints: `ckpts/<name>/`. Optional launchers: `scripts/launch_L15_w2_w5_s1.sh`, `scripts/launch_L20_bt_cso_pair.sh`.
82
+
83
+ **L20 contrast (same CE@0.90 gate):** BT W=2 / BT W=5 finish the curriculum with high leaf accuracy; CSO (`args/L20_cso_prom090.yaml`) finishes the ladder but leaf accuracy stays near chance (~0.5).
84
+
85
+ ## Latent interventions (L=20)
86
+
87
+ We probe finished L20 checkpoints by editing continuous thoughts, then measuring **leaf accuracy** on 128 val graphs.
88
+
89
+ | Protocol | What we do |
90
+ |----------|------------|
91
+ | **Pin-last** | Keep the last thought intact; replace earlier thoughts with noise / other-graph donors |
92
+ | **Corrupt last** | Replace only the final thought |
93
+ | **Propagate** | Corrupt one mid-chain thought, then recompute all later thoughts |
94
+
95
+ **Takeaway:** BT concentrates the answer in the **last** latent — wiping L1…L19 barely hurts if L20 is pinned; corrupting L20 (or propagating mid-chain noise) collapses accuracy toward chance. CSO is weak and flat under every edit.
96
+
97
+ ### Summary numbers (128 graphs)
98
+
99
+ | Method | ckpt | clean | earlier→noise (last pinned) | last→donor | pin-last all-19 |
100
+ |--------|------|------:|----------------------------:|-----------:|----------------:|
101
+ | BT W=5 | `.../checkpoint_225` | 1.000 | 1.000 | 0.516 | 1.000 |
102
+ | BT W=2 | `.../checkpoint_225` | 0.930 | 0.922 | 0.430 | 0.930 |
103
+ | CSO | `.../checkpoint_200` | 0.531 | 0.516 | 0.531 | 0.531 |
104
+
105
+ ### Figures
106
+
107
+ **All protocols (pin-last-k, aggregates, per-slot pin / propagate):**
108
+
109
+ ![L20 all intervention methods](figs/interventions/L20_all_methods.png)
110
+
111
+ **Pin-last vs number of earlier latents corrupted:**
112
+
113
+ ![L20 pin-last k](figs/interventions/pinlast_k_L20.png)
114
+
115
+ **Same pin-last-k as a table:**
116
+
117
+ ![L20 pin-last table](figs/interventions/pinlast_k_L20_table.png)
118
+
119
+ **Earlier depths (L=10 / L=15) show the same BT last-thought concentration:**
120
+
121
+ ![L10 L15 pin-last](figs/interventions/pinlast_k_L10_L15.png)
122
+
123
+ ### Re-run interventions / regenerate plots
124
+
125
+ ```bash
126
+ # needs trained ckpts + val data on disk
127
+ python scripts/intervene_L20.py --ckpt ckpts/L20_w2_s1_prom090_bt090/checkpoint_225 --name "BT W=2"
128
+ python scripts/intervene_L20.py --ckpt ckpts/L20_w5_s1_prom090_bt090/checkpoint_225 --name "BT W=5"
129
+ python scripts/intervene_L20.py --ckpt ckpts/L20_cso_prom090/checkpoint_200 --name "CSO"
130
+
131
+ # rebuild README figures from saved JSON (no GPU needed)
132
+ python scripts/plot_interventions_readme.py
133
+ ```
134
+
135
+ Raw JSON: `figs/interventions/L20_*.json`, `figs/interventions/pinlast_k_*.json`.
136
+
137
+ ## Citation (base paper)
138
+
139
+ ```bibtex
140
+ @misc{zhu2025reasoning,
141
+ title = {Reasoning by Superposition: A Theoretical Perspective on Chain of Continuous Thought},
142
+ author = {Hanlin Zhu and Shibo Hao and Zhiting Hu and Jiantao Jiao and Stuart Russell and Yuandong Tian},
143
+ year = {2025},
144
+ eprint = {2505.12514},
145
+ archivePrefix = {arXiv},
146
+ primaryClass = {cs.LG}
147
+ }
148
+ ```
149
+
150
+ ## License
151
+
152
+ MIT — see LICENSE.
args/L10_w1_prom098_bt098.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L10, W=1 for the whole curriculum after stage-0 warm start.
2
+ # Same recipe that solved L10 under full BPTT, but backprop_depth: 1 and
3
+ # both gates at ce_score 0.98 (stricter than the 0.95 winner).
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "L10_w1_prom098_bt098"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 10
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 1
20
+ promote_metric: ce_score
21
+ promote_threshold: 0.98
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.98
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 5
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 128
35
+ eval_print_full: False
36
+
37
+ backprop_depth: 1
38
+ train_size: 0
39
+ save_only_improve: False
40
+ save_every: 25
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
44
+ seed: 0
45
+ resume: 0
46
+ bf16: False
47
+ train_path: data/star_2arm_L10_train_fo_bfs.json
48
+ val_path: data/star_2arm_L10_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 2000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/L10_w1_prom099_bt099.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L10, W=1, both gates at ce_score 0.99. Same warm start as the L10 winner.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L10_w1_prom099_bt099"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 10
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.99
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.99
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 128
33
+ eval_print_full: False
34
+
35
+ backprop_depth: 1
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L10_train_fo_bfs.json
46
+ val_path: data/star_2arm_L10_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 2000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L10_w2_prom095_bt095.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L10 W=2 — same recipe as the L10 full-BPTT winner, backprop_depth: 2.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L10_w2_prom095_bt095"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 10
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.95
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.95
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 128
33
+ eval_print_full: False
34
+
35
+ backprop_depth: 2
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L10_train_fo_bfs.json
46
+ val_path: data/star_2arm_L10_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 2000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L10_w2_prom098_bt098.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L10 W=2 with tighter ce@0.98 gates (ablation vs 0.95).
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L10_w2_prom098_bt098"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 10
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.98
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.98
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 128
33
+ eval_print_full: False
34
+
35
+ backprop_depth: 2
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L10-bfs-stage0/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L10_train_fo_bfs.json
46
+ val_path: data/star_2arm_L10_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 2000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_push_2L_ce95_100k.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Data lever: same as the mirror arm but all 100k graphs.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_push_2L_ce95_100k"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.95
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.95
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: null
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_push_2L_ce95_100k_w1.yaml ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # W=1 BPTT transfer test.
2
+ # Warm-start from 100k full-BPTT promote snapshot at stage 6 (hops 1..6 already
3
+ # meet ce_score@0.95). Remaining stages 6->15 train with backprop_depth: 1 so
4
+ # gradients only flow through the last latent recurrence step.
5
+ project: coconut
6
+ save_path: ckpts
7
+ name: "L15_push_2L_ce95_100k_w1"
8
+
9
+ only_eval: False
10
+ coconut: True
11
+ cot: False
12
+ no_thoughts: False
13
+ no_cot: False
14
+
15
+ c_thought: 1
16
+ max_latent_stage: 15
17
+ pad_latent_to_max: True
18
+
19
+ accuracy_staging: True
20
+ init_stage: 6
21
+ promote_metric: ce_score
22
+ promote_threshold: 0.95
23
+ promote_on_current_only: False
24
+ epochs_per_stage: 25
25
+
26
+ backtrack: True
27
+ backtrack_metric: ce_score
28
+ backtrack_detect_threshold: 0.95
29
+ remember_rate: 0.3
30
+ revert_staging: False
31
+
32
+ eval_every: 5
33
+ log_every: 5
34
+ perhop_val_samples: 256
35
+ perhop_train_samples: 64
36
+ eval_print_full: False
37
+
38
+ backprop_depth: 1
39
+ train_size: 0
40
+ save_only_improve: False
41
+ save_every: 25
42
+ uniform_prob: 0.1
43
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
44
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
45
+ seed: 0
46
+ resume: 0
47
+ bf16: False
48
+ train_path: data/star_2arm_L15_train_fo_bfs.json
49
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
50
+ reset_optimizer: False
51
+ batch_size_training: 128
52
+ debug: False
53
+ gradient_accumulation_steps: 1
54
+ num_epochs: 4000
55
+ lr: !!float "1e-4"
56
+ grad_clip: !!float "1.0"
57
+ warmup_steps: 200
58
+ weight_decay: 0.01
59
+ bfs_variant: True
args/L15_push_2L_ce95_100k_w1_prom098.yaml ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # W=1 BPTT transfer, tighter PROMOTE gate (0.98).
2
+ # Same start as L15_push_2L_ce95_100k_w1: full-BPTT ckpt_120, init_stage 6,
3
+ # backprop_depth 1. Only change: promote_threshold 0.95 -> 0.98 so stages
4
+ # only advance when balance is very tight. Backtrack stays at 0.95.
5
+ project: coconut
6
+ save_path: ckpts
7
+ name: "L15_push_2L_ce95_100k_w1_prom098"
8
+
9
+ only_eval: False
10
+ coconut: True
11
+ cot: False
12
+ no_thoughts: False
13
+ no_cot: False
14
+
15
+ c_thought: 1
16
+ max_latent_stage: 15
17
+ pad_latent_to_max: True
18
+
19
+ accuracy_staging: True
20
+ init_stage: 6
21
+ promote_metric: ce_score
22
+ promote_threshold: 0.98
23
+ promote_on_current_only: False
24
+ epochs_per_stage: 25
25
+
26
+ backtrack: True
27
+ backtrack_metric: ce_score
28
+ backtrack_detect_threshold: 0.95
29
+ remember_rate: 0.3
30
+ revert_staging: False
31
+
32
+ eval_every: 5
33
+ log_every: 5
34
+ perhop_val_samples: 256
35
+ perhop_train_samples: 64
36
+ eval_print_full: False
37
+
38
+ backprop_depth: 1
39
+ train_size: 0
40
+ save_only_improve: False
41
+ save_every: 25
42
+ uniform_prob: 0.1
43
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
44
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
45
+ seed: 0
46
+ resume: 0
47
+ bf16: False
48
+ train_path: data/star_2arm_L15_train_fo_bfs.json
49
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
50
+ reset_optimizer: False
51
+ batch_size_training: 128
52
+ debug: False
53
+ gradient_accumulation_steps: 1
54
+ num_epochs: 4000
55
+ lr: !!float "1e-4"
56
+ grad_clip: !!float "1.0"
57
+ warmup_steps: 200
58
+ weight_decay: 0.01
59
+ bfs_variant: True
args/L15_push_2L_ce95_50k.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Exact L10-winner recipe at L15: 2-layer, ce_score@0.95 on both gates, fp32, 50k.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_push_2L_ce95_50k"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.95
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.95
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: null
36
+ train_size: 50000
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_push_4L_ce90_50k.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Capacity + looser gate: 4-layer, ce_score@0.90, 50k.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_push_4L_ce90_50k"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.9
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.9
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: null
36
+ train_size: 50000
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-4layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_push_4L_ce95_50k.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Capacity lever: 4-layer via layer-expansion warm start, ce_score@0.95, 50k.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_push_4L_ce95_50k"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 1
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.95
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.95
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: null
36
+ train_size: 50000
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-4layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_push_s0_4L.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 4-layer L15 stage-0, pinned. Produces a native 4-layer stage-0 checkpoint in case layer-expansion warm start degrades hop-1.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_push_s0_4L"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 0
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 0
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.95
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: False
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.95
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: null
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-4layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_w1_prom098_bt098.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # W=1 from stage-6 full-BPTT ckpt. BOTH gates at ce_score 0.98.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_w1_prom098_bt098"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 6
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.98
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.98
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: 1
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_w1_prom099_bt099.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # W=1 from stage-6 full-BPTT ckpt. BOTH gates at ce_score 0.99.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_w1_prom099_bt099"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 6
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.99
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.99
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: 1
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_w2_prom095_bt095.yaml ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # W=2 truncated BPTT from stage-6 full-BPTT ckpt.
2
+ # Same gates as the completed W=1@0.95 run for a clean speed/quality compare.
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "L15_w2_prom095_bt095"
6
+
7
+ only_eval: False
8
+ coconut: True
9
+ cot: False
10
+ no_thoughts: False
11
+ no_cot: False
12
+
13
+ c_thought: 1
14
+ max_latent_stage: 15
15
+ pad_latent_to_max: True
16
+
17
+ accuracy_staging: True
18
+ init_stage: 6
19
+ promote_metric: ce_score
20
+ promote_threshold: 0.95
21
+ promote_on_current_only: False
22
+ epochs_per_stage: 25
23
+
24
+ backtrack: True
25
+ backtrack_metric: ce_score
26
+ backtrack_detect_threshold: 0.95
27
+ remember_rate: 0.3
28
+ revert_staging: False
29
+
30
+ eval_every: 5
31
+ log_every: 5
32
+ perhop_val_samples: 256
33
+ perhop_train_samples: 64
34
+ eval_print_full: False
35
+
36
+ backprop_depth: 2
37
+ train_size: 0
38
+ save_only_improve: False
39
+ save_every: 25
40
+ uniform_prob: 0.1
41
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
42
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
43
+ seed: 0
44
+ resume: 0
45
+ bf16: False
46
+ train_path: data/star_2arm_L15_train_fo_bfs.json
47
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
48
+ reset_optimizer: False
49
+ batch_size_training: 128
50
+ debug: False
51
+ gradient_accumulation_steps: 1
52
+ num_epochs: 4000
53
+ lr: !!float "1e-4"
54
+ grad_clip: !!float "1.0"
55
+ warmup_steps: 200
56
+ weight_decay: 0.01
57
+ bfs_variant: True
args/L15_w2_s1_prom095_bt095.yaml ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 W=2 from stage 1 (whole curriculum after stage0 warm-start).
2
+ # Same recipe as L15_push_2L_ce95_100k, backprop_depth: 2.
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "L15_w2_s1_prom095_bt095"
6
+
7
+ only_eval: False
8
+ coconut: True
9
+ cot: False
10
+ no_thoughts: False
11
+ no_cot: False
12
+
13
+ c_thought: 1
14
+ max_latent_stage: 15
15
+ pad_latent_to_max: True
16
+
17
+ accuracy_staging: True
18
+ init_stage: 1
19
+ promote_metric: ce_score
20
+ promote_threshold: 0.95
21
+ promote_on_current_only: False
22
+ epochs_per_stage: 25
23
+
24
+ backtrack: True
25
+ backtrack_metric: ce_score
26
+ backtrack_detect_threshold: 0.95
27
+ remember_rate: 0.3
28
+ revert_staging: False
29
+
30
+ eval_every: 5
31
+ log_every: 5
32
+ perhop_val_samples: 256
33
+ perhop_train_samples: 64
34
+ eval_print_full: False
35
+
36
+ backprop_depth: 2
37
+ train_size: 0
38
+ save_only_improve: False
39
+ save_every: 25
40
+ uniform_prob: 0.1
41
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
42
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
43
+ seed: 0
44
+ resume: 0
45
+ bf16: False
46
+ train_path: data/star_2arm_L15_train_fo_bfs.json
47
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
48
+ reset_optimizer: False
49
+ batch_size_training: 128
50
+ debug: False
51
+ gradient_accumulation_steps: 1
52
+ num_epochs: 4000
53
+ lr: !!float "1e-4"
54
+ grad_clip: !!float "1.0"
55
+ warmup_steps: 200
56
+ weight_decay: 0.01
57
+ bfs_variant: True
args/L15_w5_prom095_bt095.yaml ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # W=5 truncated BPTT from stage-6 full-BPTT ckpt.
2
+ # Same gates as the completed W=1@0.95 run for a clean speed/quality compare.
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "L15_w5_prom095_bt095"
6
+
7
+ only_eval: False
8
+ coconut: True
9
+ cot: False
10
+ no_thoughts: False
11
+ no_cot: False
12
+
13
+ c_thought: 1
14
+ max_latent_stage: 15
15
+ pad_latent_to_max: True
16
+
17
+ accuracy_staging: True
18
+ init_stage: 6
19
+ promote_metric: ce_score
20
+ promote_threshold: 0.95
21
+ promote_on_current_only: False
22
+ epochs_per_stage: 25
23
+
24
+ backtrack: True
25
+ backtrack_metric: ce_score
26
+ backtrack_detect_threshold: 0.95
27
+ remember_rate: 0.3
28
+ revert_staging: False
29
+
30
+ eval_every: 5
31
+ log_every: 5
32
+ perhop_val_samples: 256
33
+ perhop_train_samples: 64
34
+ eval_print_full: False
35
+
36
+ backprop_depth: 5
37
+ train_size: 0
38
+ save_only_improve: False
39
+ save_every: 25
40
+ uniform_prob: 0.1
41
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
42
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
43
+ seed: 0
44
+ resume: 0
45
+ bf16: False
46
+ train_path: data/star_2arm_L15_train_fo_bfs.json
47
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
48
+ reset_optimizer: False
49
+ batch_size_training: 128
50
+ debug: False
51
+ gradient_accumulation_steps: 1
52
+ num_epochs: 4000
53
+ lr: !!float "1e-4"
54
+ grad_clip: !!float "1.0"
55
+ warmup_steps: 200
56
+ weight_decay: 0.01
57
+ bfs_variant: True
args/L15_w5_prom098_bt098.yaml ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 W=5 from stage-6 full-BPTT ckpt, both gates at ce_score 0.98.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "L15_w5_prom098_bt098"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 6
18
+ promote_metric: ce_score
19
+ promote_threshold: 0.98
20
+ promote_on_current_only: False
21
+ epochs_per_stage: 25
22
+
23
+ backtrack: True
24
+ backtrack_metric: ce_score
25
+ backtrack_detect_threshold: 0.98
26
+ remember_rate: 0.3
27
+ revert_staging: False
28
+
29
+ eval_every: 5
30
+ log_every: 5
31
+ perhop_val_samples: 256
32
+ perhop_train_samples: 64
33
+ eval_print_full: False
34
+
35
+ backprop_depth: 5
36
+ train_size: 0
37
+ save_only_improve: False
38
+ save_every: 25
39
+ uniform_prob: 0.1
40
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
41
+ load_model_path: ckpts/L15_push_2L_ce95_100k/checkpoint_120
42
+ seed: 0
43
+ resume: 0
44
+ bf16: False
45
+ train_path: data/star_2arm_L15_train_fo_bfs.json
46
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
47
+ reset_optimizer: False
48
+ batch_size_training: 128
49
+ debug: False
50
+ gradient_accumulation_steps: 1
51
+ num_epochs: 4000
52
+ lr: !!float "1e-4"
53
+ grad_clip: !!float "1.0"
54
+ warmup_steps: 200
55
+ weight_decay: 0.01
56
+ bfs_variant: True
args/L15_w5_s1_prom095_bt095.yaml ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 W=5 from stage 1 (whole curriculum after stage0 warm-start).
2
+ # Same recipe as L15_push_2L_ce95_100k, backprop_depth: 5.
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "L15_w5_s1_prom095_bt095"
6
+
7
+ only_eval: False
8
+ coconut: True
9
+ cot: False
10
+ no_thoughts: False
11
+ no_cot: False
12
+
13
+ c_thought: 1
14
+ max_latent_stage: 15
15
+ pad_latent_to_max: True
16
+
17
+ accuracy_staging: True
18
+ init_stage: 1
19
+ promote_metric: ce_score
20
+ promote_threshold: 0.95
21
+ promote_on_current_only: False
22
+ epochs_per_stage: 25
23
+
24
+ backtrack: True
25
+ backtrack_metric: ce_score
26
+ backtrack_detect_threshold: 0.95
27
+ remember_rate: 0.3
28
+ revert_staging: False
29
+
30
+ eval_every: 5
31
+ log_every: 5
32
+ perhop_val_samples: 256
33
+ perhop_train_samples: 64
34
+ eval_print_full: False
35
+
36
+ backprop_depth: 5
37
+ train_size: 0
38
+ save_only_improve: False
39
+ save_every: 25
40
+ uniform_prob: 0.1
41
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
42
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
43
+ seed: 0
44
+ resume: 0
45
+ bf16: False
46
+ train_path: data/star_2arm_L15_train_fo_bfs.json
47
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
48
+ reset_optimizer: False
49
+ batch_size_training: 128
50
+ debug: False
51
+ gradient_accumulation_steps: 1
52
+ num_epochs: 4000
53
+ lr: !!float "1e-4"
54
+ grad_clip: !!float "1.0"
55
+ warmup_steps: 200
56
+ weight_decay: 0.01
57
+ bfs_variant: True
args/L20_cso_prom090.yaml ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L20 CSO (no backtracking) — fair contrast to L20_w2_s1_prom090_bt090.
2
+ # Same CE promote gate @ 0.90; promote_on_current_only so earlier hops can
3
+ # regress without blocking. No force-promote: if it stalls, that is the result
4
+ # (collapse / failure to finish curriculum).
5
+ project: coconut
6
+ save_path: ckpts
7
+ name: "L20_cso_prom090"
8
+
9
+ only_eval: False
10
+ coconut: True
11
+ cot: False
12
+ no_thoughts: False
13
+ no_cot: False
14
+
15
+ c_thought: 1
16
+ max_latent_stage: 20
17
+ pad_latent_to_max: True
18
+
19
+ accuracy_staging: True
20
+ init_stage: 20
21
+ promote_metric: ce_score
22
+ promote_threshold: 0.90
23
+ promote_on_current_only: True
24
+ epochs_per_stage: 25
25
+
26
+ backtrack: False
27
+ backtrack_metric: ce_score
28
+ backtrack_detect_threshold: 0.90
29
+ remember_rate: 0.3
30
+ revert_staging: False
31
+
32
+ eval_every: 5
33
+ log_every: 5
34
+ perhop_val_samples: 256
35
+ perhop_train_samples: 64
36
+ eval_print_full: False
37
+
38
+ backprop_depth: null
39
+ train_size: 0
40
+ save_only_improve: False
41
+ save_every: 5
42
+ uniform_prob: 0.1
43
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
44
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
45
+ seed: 0
46
+ resume: 0
47
+ bf16: False
48
+ train_path: data/star_2arm_L20_100k_train_fo_bfs.json
49
+ val_path: data/star_2arm_L20_100k_valid_fo_bfs.json
50
+ reset_optimizer: False
51
+ batch_size_training: 128
52
+ debug: False
53
+ gradient_accumulation_steps: 1
54
+ num_epochs: 5000
55
+ lr: !!float "1e-4"
56
+ grad_clip: !!float "1.0"
57
+ warmup_steps: 200
58
+ weight_decay: 0.01
59
+ bfs_variant: True
args/L20_w20_s1_prom090_bt090.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L20 BT — CE-gated from stage 1, FULL chain BPTT (W=null / all 20 latents).
2
+ # Same gates and warm-start as L20_w2_s1_prom090_bt090 / L20_w5_s1_prom090_bt090.
3
+ # backprop_depth: null => no truncated BPTT detach (coconut.py full BPTT).
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "L20_w20_s1_prom090_bt090"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 20
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 1
20
+ promote_metric: ce_score
21
+ promote_threshold: 0.90
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.90
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 5
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 0
39
+ save_only_improve: False
40
+ save_every: 5
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
44
+ seed: 0
45
+ resume: 0
46
+ bf16: False
47
+ train_path: data/star_2arm_L20_100k_train_fo_bfs.json
48
+ val_path: data/star_2arm_L20_100k_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 5000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/L20_w2_s1_prom090_bt090.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L20 BT — CE-gated from stage 1, W=2 throughout.
2
+ # Threshold lowered vs L15 (0.95 -> 0.90) so both BT and CSO share a fair,
3
+ # achievable gate at depth 20. Matched promote/backtrack on ce_score.
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "L20_w2_s1_prom090_bt090"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 20
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 1
20
+ promote_metric: ce_score
21
+ promote_threshold: 0.90
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.90
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 5
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: 2
38
+ train_size: 0
39
+ save_only_improve: False
40
+ save_every: 5
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
44
+ seed: 0
45
+ resume: 0
46
+ bf16: False
47
+ train_path: data/star_2arm_L20_100k_train_fo_bfs.json
48
+ val_path: data/star_2arm_L20_100k_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 5000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/L20_w5_s1_prom090_bt090.yaml ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L20 BT — CE-gated from stage 1, W=5 throughout.
2
+ # Same gates as L20_w2_s1_prom090_bt090 / L20_cso_prom090 (ce_score @ 0.90).
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "L20_w5_s1_prom090_bt090"
6
+
7
+ only_eval: False
8
+ coconut: True
9
+ cot: False
10
+ no_thoughts: False
11
+ no_cot: False
12
+
13
+ c_thought: 1
14
+ max_latent_stage: 20
15
+ pad_latent_to_max: True
16
+
17
+ accuracy_staging: True
18
+ init_stage: 20
19
+ promote_metric: ce_score
20
+ promote_threshold: 0.90
21
+ promote_on_current_only: False
22
+ epochs_per_stage: 25
23
+
24
+ backtrack: True
25
+ backtrack_metric: ce_score
26
+ backtrack_detect_threshold: 0.90
27
+ remember_rate: 0.3
28
+ revert_staging: False
29
+
30
+ eval_every: 5
31
+ log_every: 5
32
+ perhop_val_samples: 256
33
+ perhop_train_samples: 64
34
+ eval_print_full: False
35
+
36
+ backprop_depth: 5
37
+ train_size: 0
38
+ save_only_improve: False
39
+ save_every: 5
40
+ uniform_prob: 0.1
41
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
42
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_150
43
+ seed: 0
44
+ resume: 0
45
+ bf16: False
46
+ train_path: data/star_2arm_L20_100k_train_fo_bfs.json
47
+ val_path: data/star_2arm_L20_100k_valid_fo_bfs.json
48
+ reset_optimizer: False
49
+ batch_size_training: 128
50
+ debug: False
51
+ gradient_accumulation_steps: 1
52
+ num_epochs: 5000
53
+ lr: !!float "1e-4"
54
+ grad_clip: !!float "1.0"
55
+ warmup_steps: 200
56
+ weight_decay: 0.01
57
+ bfs_variant: True
args/diag_L15_ce_thr050.yaml ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAGNOSTIC arm 3/4 — CE-score gated BT at threshold 0.50.
2
+ # 0.9/0.7 were unreachable at depth; 0.5 ≈ "both arms still have real mass"
3
+ # (~93/7 split for |F|=2) while still demanding superposition-ish balance.
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag-L15-ce-thr050"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_threshold: 0.50
21
+ epochs_per_stage: 25
22
+ staging_metric: ce_score
23
+
24
+ eval_every: 2
25
+ log_every: 1
26
+ perhop_val_samples: 256
27
+ perhop_train_samples: 64
28
+
29
+ backprop_depth: null
30
+ backtrack: True
31
+ remember_rate: 0.3
32
+ backtrack_detect_threshold: 0.50
33
+ revert_staging: False
34
+
35
+ train_size: 50000
36
+ save_only_improve: False
37
+ save_every: 50
38
+ uniform_prob: 0.1
39
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
40
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
41
+ seed: 0
42
+ resume: 0
43
+ bf16: True
44
+ train_path: data/star_2arm_L15_train_fo_bfs.json
45
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
46
+ reset_optimizer: False
47
+ batch_size_training: 128
48
+ debug: False
49
+ gradient_accumulation_steps: 1
50
+ num_epochs: 3000
51
+ lr: !!float "1e-4"
52
+ grad_clip: !!float "1.0"
53
+ warmup_steps: 200
54
+ weight_decay: 0.01
55
+ bfs_variant: True
args/diag_L15_ce_thr090.yaml ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — both-arms score (ce_score) gate at 0.90 + backtracking.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "diag-L15-ce-thr090"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 0
18
+ promote_threshold: 0.90
19
+ epochs_per_stage: 25
20
+ staging_metric: ce_score
21
+
22
+ eval_every: 2
23
+ log_every: 1
24
+ perhop_val_samples: 256
25
+ perhop_train_samples: 64
26
+
27
+ backprop_depth: null
28
+ backtrack: True
29
+ remember_rate: 0.3
30
+ backtrack_detect_threshold: 0.90
31
+ revert_staging: False
32
+
33
+ train_size: 50000
34
+ save_only_improve: False
35
+ save_every: 50
36
+ uniform_prob: 0.1
37
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
38
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
39
+ seed: 0
40
+ resume: 0
41
+ bf16: True
42
+ train_path: data/star_2arm_L15_train_fo_bfs.json
43
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
44
+ reset_optimizer: False
45
+ batch_size_training: 128
46
+ debug: False
47
+ gradient_accumulation_steps: 1
48
+ num_epochs: 3000
49
+ lr: !!float "1e-4"
50
+ grad_clip: !!float "1.0"
51
+ warmup_steps: 200
52
+ weight_decay: 0.01
53
+ bfs_variant: True
args/diag_L15_ce_thr095.yaml ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — both-arms score (ce_score) gate at 0.95 + backtracking.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "diag-L15-ce-thr095"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 0
18
+ promote_threshold: 0.95
19
+ epochs_per_stage: 25
20
+ staging_metric: ce_score
21
+
22
+ eval_every: 2
23
+ log_every: 1
24
+ perhop_val_samples: 256
25
+ perhop_train_samples: 64
26
+
27
+ backprop_depth: null
28
+ backtrack: True
29
+ remember_rate: 0.3
30
+ backtrack_detect_threshold: 0.95
31
+ revert_staging: False
32
+
33
+ train_size: 50000
34
+ save_only_improve: False
35
+ save_every: 50
36
+ uniform_prob: 0.1
37
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
38
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
39
+ seed: 0
40
+ resume: 0
41
+ bf16: True
42
+ train_path: data/star_2arm_L15_train_fo_bfs.json
43
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
44
+ reset_optimizer: False
45
+ batch_size_training: 128
46
+ debug: False
47
+ gradient_accumulation_steps: 1
48
+ num_epochs: 3000
49
+ lr: !!float "1e-4"
50
+ grad_clip: !!float "1.0"
51
+ warmup_steps: 200
52
+ weight_decay: 0.01
53
+ bfs_variant: True
args/diag_L15_frontier_nobt_095.yaml ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAGNOSTIC arm 4/4 — frontier@0.95, NO backtracking, promote on current
2
+ # stage only. Isolates "how long does each stage take to hit 95%" without
3
+ # earlier-hop retention blocking promotion.
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag-L15-frontier-nobt-095"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_threshold: 0.95
21
+ epochs_per_stage: 25
22
+ staging_metric: frontier
23
+ promote_on_current_only: True
24
+
25
+ eval_every: 2
26
+ log_every: 1
27
+ perhop_val_samples: 256
28
+ perhop_train_samples: 64
29
+
30
+ backprop_depth: null
31
+ backtrack: False
32
+ revert_staging: False
33
+
34
+ train_size: 50000
35
+ save_only_improve: False
36
+ save_every: 50
37
+ uniform_prob: 0.1
38
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
39
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
40
+ seed: 0
41
+ resume: 0
42
+ bf16: True
43
+ train_path: data/star_2arm_L15_train_fo_bfs.json
44
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
45
+ reset_optimizer: False
46
+ batch_size_training: 128
47
+ debug: False
48
+ gradient_accumulation_steps: 1
49
+ num_epochs: 3000
50
+ lr: !!float "1e-4"
51
+ grad_clip: !!float "1.0"
52
+ warmup_steps: 200
53
+ weight_decay: 0.01
54
+ bfs_variant: True
args/diag_L15_frontier_thr085.yaml ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAGNOSTIC arm 2/4 — frontier-gated BT at threshold 0.85 (lower bar).
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "diag-L15-frontier-thr085"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 0
18
+ promote_threshold: 0.85
19
+ epochs_per_stage: 25
20
+ staging_metric: frontier
21
+
22
+ eval_every: 2
23
+ log_every: 1
24
+ perhop_val_samples: 256
25
+ perhop_train_samples: 64
26
+
27
+ backprop_depth: null
28
+ backtrack: True
29
+ remember_rate: 0.3
30
+ backtrack_detect_threshold: 0.85
31
+ revert_staging: False
32
+
33
+ train_size: 50000
34
+ save_only_improve: False
35
+ save_every: 50
36
+ uniform_prob: 0.1
37
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
38
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
39
+ seed: 0
40
+ resume: 0
41
+ bf16: True
42
+ train_path: data/star_2arm_L15_train_fo_bfs.json
43
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
44
+ reset_optimizer: False
45
+ batch_size_training: 128
46
+ debug: False
47
+ gradient_accumulation_steps: 1
48
+ num_epochs: 3000
49
+ lr: !!float "1e-4"
50
+ grad_clip: !!float "1.0"
51
+ warmup_steps: 200
52
+ weight_decay: 0.01
53
+ bfs_variant: True
args/diag_L15_frontier_thr090.yaml ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — hop-accuracy (frontier) gate at 0.90 + backtracking.
2
+ project: coconut
3
+ save_path: ckpts
4
+ name: "diag-L15-frontier-thr090"
5
+
6
+ only_eval: False
7
+ coconut: True
8
+ cot: False
9
+ no_thoughts: False
10
+ no_cot: False
11
+
12
+ c_thought: 1
13
+ max_latent_stage: 15
14
+ pad_latent_to_max: True
15
+
16
+ accuracy_staging: True
17
+ init_stage: 0
18
+ promote_threshold: 0.90
19
+ epochs_per_stage: 25
20
+ staging_metric: frontier
21
+
22
+ eval_every: 2
23
+ log_every: 1
24
+ perhop_val_samples: 256
25
+ perhop_train_samples: 64
26
+
27
+ backprop_depth: null
28
+ backtrack: True
29
+ remember_rate: 0.3
30
+ backtrack_detect_threshold: 0.90
31
+ revert_staging: False
32
+
33
+ train_size: 50000
34
+ save_only_improve: False
35
+ save_every: 50
36
+ uniform_prob: 0.1
37
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
38
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
39
+ seed: 0
40
+ resume: 0
41
+ bf16: True
42
+ train_path: data/star_2arm_L15_train_fo_bfs.json
43
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
44
+ reset_optimizer: False
45
+ batch_size_training: 128
46
+ debug: False
47
+ gradient_accumulation_steps: 1
48
+ num_epochs: 3000
49
+ lr: !!float "1e-4"
50
+ grad_clip: !!float "1.0"
51
+ warmup_steps: 200
52
+ weight_decay: 0.01
53
+ bfs_variant: True
args/diag_L15_frontier_thr095.yaml ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAGNOSTIC arm 1/4 — frontier-gated BT at promote/BT threshold 0.95.
2
+ # Start from stage 0 (warm hop-1 ckpt). Measure epochs-to-0.95 per stage.
3
+ # ce_score is always logged alongside frontier for BT-metric diagnosis.
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag-L15-frontier-thr095"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_threshold: 0.95
21
+ epochs_per_stage: 25
22
+ staging_metric: frontier
23
+
24
+ eval_every: 2
25
+ log_every: 1
26
+ perhop_val_samples: 256
27
+ perhop_train_samples: 64
28
+
29
+ backprop_depth: null
30
+ backtrack: True
31
+ remember_rate: 0.3
32
+ backtrack_detect_threshold: 0.95
33
+ revert_staging: False
34
+
35
+ train_size: 50000
36
+ save_only_improve: False
37
+ save_every: 50
38
+ uniform_prob: 0.1
39
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
40
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
41
+ seed: 0
42
+ resume: 0
43
+ bf16: True
44
+ train_path: data/star_2arm_L15_train_fo_bfs.json
45
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
46
+ reset_optimizer: False
47
+ batch_size_training: 128
48
+ debug: False
49
+ gradient_accumulation_steps: 1
50
+ num_epochs: 3000
51
+ lr: !!float "1e-4"
52
+ grad_clip: !!float "1.0"
53
+ warmup_steps: 200
54
+ weight_decay: 0.01
55
+ bfs_variant: True
args/diag_L15_frontier_thr099.yaml ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — hop-accuracy gate at 0.99 + backtracking.
2
+ # Hypothesis: 0.95 promotes too fast (~2 epochs/stage early on); 0.99 forces
3
+ # each latent stage to be nearly solved before advancing.
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag-L15-frontier-thr099"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_threshold: 0.99
21
+ epochs_per_stage: 25
22
+ staging_metric: frontier
23
+
24
+ eval_every: 2
25
+ log_every: 1
26
+ perhop_val_samples: 256
27
+ perhop_train_samples: 64
28
+
29
+ backprop_depth: null
30
+ backtrack: True
31
+ remember_rate: 0.3
32
+ backtrack_detect_threshold: 0.99
33
+ revert_staging: False
34
+
35
+ train_size: 50000
36
+ save_only_improve: False
37
+ save_every: 50
38
+ uniform_prob: 0.1
39
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
40
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
41
+ seed: 0
42
+ resume: 0
43
+ bf16: True
44
+ train_path: data/star_2arm_L15_train_fo_bfs.json
45
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
46
+ reset_optimizer: False
47
+ batch_size_training: 128
48
+ debug: False
49
+ gradient_accumulation_steps: 1
50
+ num_epochs: 3000
51
+ lr: !!float "1e-4"
52
+ grad_clip: !!float "1.0"
53
+ warmup_steps: 200
54
+ weight_decay: 0.01
55
+ bfs_variant: True
args/diag_L15_promCE090_btCE090.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): ce_score >= 0.9
3
+ # BACKTRACK (retrain earlier): ce_score < 0.9 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promCE090_btCE090"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: ce_score
21
+ promote_threshold: 0.9
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.9
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF095_btCE050.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.95
3
+ # BACKTRACK (retrain earlier): ce_score < 0.5 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF095_btCE050"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.95
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.5
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF095_btCE090.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.95
3
+ # BACKTRACK (retrain earlier): ce_score < 0.9 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF095_btCE090"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.95
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.9
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF095_btCE095.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.95
3
+ # BACKTRACK (retrain earlier): ce_score < 0.95 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF095_btCE095"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.95
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.95
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF095_btF095.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.95
3
+ # BACKTRACK (retrain earlier): frontier < 0.95 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF095_btF095"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.95
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: frontier
27
+ backtrack_detect_threshold: 0.95
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF095_btNONE.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.95
3
+ # BACKTRACK (retrain earlier): OFF
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF095_btNONE"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.95
22
+ promote_on_current_only: True
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: False
26
+ backtrack_metric: frontier
27
+ backtrack_detect_threshold: 0.95
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF099_btCE050.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.99
3
+ # BACKTRACK (retrain earlier): ce_score < 0.5 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF099_btCE050"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.99
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.5
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_L15_promF099_btCE090.yaml ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # L15 DIAG — SEPARATE gates
2
+ # PROMOTE (stage i -> i+1): frontier >= 0.99
3
+ # BACKTRACK (retrain earlier): ce_score < 0.9 triggers retrain
4
+ project: coconut
5
+ save_path: ckpts
6
+ name: "diag_L15_promF099_btCE090"
7
+
8
+ only_eval: False
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ max_latent_stage: 15
16
+ pad_latent_to_max: True
17
+
18
+ accuracy_staging: True
19
+ init_stage: 0
20
+ promote_metric: frontier
21
+ promote_threshold: 0.99
22
+ promote_on_current_only: False
23
+ epochs_per_stage: 25
24
+
25
+ backtrack: True
26
+ backtrack_metric: ce_score
27
+ backtrack_detect_threshold: 0.9
28
+ remember_rate: 0.3
29
+ revert_staging: False
30
+
31
+ eval_every: 10
32
+ log_every: 5
33
+ perhop_val_samples: 256
34
+ perhop_train_samples: 64
35
+ eval_print_full: False
36
+
37
+ backprop_depth: null
38
+ train_size: 50000
39
+ save_only_improve: False
40
+ save_every: 200
41
+ uniform_prob: 0.1
42
+ model_id: configs/symbol-2layer-8head-768dim-L20.json
43
+ load_model_path: ckpts/star-coconut-L15-bfs-stage0-warm/checkpoint_99
44
+ seed: 0
45
+ resume: 0
46
+ bf16: True
47
+ train_path: data/star_2arm_L15_train_fo_bfs.json
48
+ val_path: data/star_2arm_L15_valid_fo_bfs.json
49
+ reset_optimizer: False
50
+ batch_size_training: 128
51
+ debug: False
52
+ gradient_accumulation_steps: 1
53
+ num_epochs: 3000
54
+ lr: !!float "1e-4"
55
+ grad_clip: !!float "1.0"
56
+ warmup_steps: 200
57
+ weight_decay: 0.01
58
+ bfs_variant: True
args/diag_ctrl.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # need 2 gpus
2
+
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "star-cc-d1-14k-ctrl"
6
+
7
+ only_eval: False
8
+
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ epochs_per_stage: 25
16
+ max_latent_stage: 1
17
+ pad_latent_to_max: True
18
+
19
+ save_only_improve: True
20
+ uniform_prob: 0.0
21
+ model_id: configs/symbol-2layer-8head-768dim.json
22
+ load_model_path: None
23
+ seed: 0
24
+ resume: 0
25
+ bf16: False
26
+ train_path: data/star_d1_train14000_fo_coconut.json
27
+ val_path: data/star_d1_valid_fo_coconut.json
28
+ reset_optimizer: False
29
+ batch_size_training: 128
30
+ debug: False
31
+ gradient_accumulation_steps: 1
32
+ num_epochs: 80
33
+ lr: !!float "1e-4"
34
+ weight_decay: 0.01
35
+ bfs_variant: False
36
+ final_only: False
args/diag_lr1e3.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # need 2 gpus
2
+
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "star-fo-d1-14k-lr1e3"
6
+
7
+ only_eval: False
8
+
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ epochs_per_stage: 25
16
+ max_latent_stage: 1
17
+ pad_latent_to_max: True
18
+
19
+ save_only_improve: True
20
+ uniform_prob: 0.0
21
+ model_id: configs/symbol-2layer-8head-768dim.json
22
+ load_model_path: None
23
+ seed: 0
24
+ resume: 0
25
+ bf16: False
26
+ train_path: data/star_d1_train14000_fo_coconut.json
27
+ val_path: data/star_d1_valid_fo_coconut.json
28
+ reset_optimizer: False
29
+ batch_size_training: 128
30
+ debug: False
31
+ gradient_accumulation_steps: 1
32
+ num_epochs: 80
33
+ lr: !!float "1e-3"
34
+ weight_decay: 0.01
35
+ bfs_variant: False
36
+ final_only: True
args/diag_lr3e4.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # need 2 gpus
2
+
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "star-fo-d1-14k-lr3e4"
6
+
7
+ only_eval: False
8
+
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ epochs_per_stage: 25
16
+ max_latent_stage: 1
17
+ pad_latent_to_max: True
18
+
19
+ save_only_improve: True
20
+ uniform_prob: 0.0
21
+ model_id: configs/symbol-2layer-8head-768dim.json
22
+ load_model_path: None
23
+ seed: 0
24
+ resume: 0
25
+ bf16: False
26
+ train_path: data/star_d1_train14000_fo_coconut.json
27
+ val_path: data/star_d1_valid_fo_coconut.json
28
+ reset_optimizer: False
29
+ batch_size_training: 128
30
+ debug: False
31
+ gradient_accumulation_steps: 1
32
+ num_epochs: 80
33
+ lr: !!float "3e-4"
34
+ weight_decay: 0.01
35
+ bfs_variant: False
36
+ final_only: True
args/diag_small.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # need 2 gpus
2
+
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "star-fo-d1-14k-small"
6
+
7
+ only_eval: False
8
+
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ epochs_per_stage: 25
16
+ max_latent_stage: 1
17
+ pad_latent_to_max: True
18
+
19
+ save_only_improve: True
20
+ uniform_prob: 0.0
21
+ model_id: configs/symbol-2layer-8head-128dim.json
22
+ load_model_path: None
23
+ seed: 0
24
+ resume: 0
25
+ bf16: False
26
+ train_path: data/star_d1_train14000_fo_coconut.json
27
+ val_path: data/star_d1_valid_fo_coconut.json
28
+ reset_optimizer: False
29
+ batch_size_training: 128
30
+ debug: False
31
+ gradient_accumulation_steps: 1
32
+ num_epochs: 40
33
+ lr: !!float "1e-4"
34
+ weight_decay: 0.01
35
+ bfs_variant: False
36
+ final_only: True
args/diag_smallL.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # need 2 gpus
2
+
3
+ project: coconut
4
+ save_path: ckpts
5
+ name: "star-fo-d1-14k-smallL"
6
+
7
+ only_eval: False
8
+
9
+ coconut: True
10
+ cot: False
11
+ no_thoughts: False
12
+ no_cot: False
13
+
14
+ c_thought: 1
15
+ epochs_per_stage: 25
16
+ max_latent_stage: 1
17
+ pad_latent_to_max: True
18
+
19
+ save_only_improve: True
20
+ uniform_prob: 0.0
21
+ model_id: configs/symbol-2layer-8head-128dim.json
22
+ load_model_path: None
23
+ seed: 0
24
+ resume: 0
25
+ bf16: False
26
+ train_path: data/star_d1_train14000_fo_coconut.json
27
+ val_path: data/star_d1_valid_fo_coconut.json
28
+ reset_optimizer: False
29
+ batch_size_training: 128
30
+ debug: False
31
+ gradient_accumulation_steps: 1
32
+ num_epochs: 150
33
+ lr: !!float "1e-4"
34
+ weight_decay: 0.01
35
+ bfs_variant: False
36
+ final_only: True