Physicsru commited on
Commit
a1fcfb7
·
verified ·
1 Parent(s): 81f1603

For Minegishi: training documentation, reproducibility files and line plots 4/4

Browse files
Files changed (30) hide show
  1. for Minegishi/SHA256SUMS +104 -0
  2. for Minegishi/reproducibility/data/chain_loop_k8/train_atomic.json +0 -0
  3. for Minegishi/reproducibility/data/chain_loop_k8/train_d2.json +0 -0
  4. for Minegishi/reproducibility/data/chain_loop_k8/train_d3.json +1 -0
  5. for Minegishi/reproducibility/data/chain_loop_k8/train_w2.json +0 -0
  6. for Minegishi/reproducibility/data/chain_loop_k8/val_d2.json +1 -0
  7. for Minegishi/reproducibility/data/chain_loop_k8/val_w2.json +1 -0
  8. for Minegishi/reproducibility/data/chain_loop_k8/vocab.json +1 -0
  9. for Minegishi/reproducibility/eval_checkpoints.py +26 -0
  10. for Minegishi/reproducibility/extrapolation.py +289 -0
  11. for Minegishi/reproducibility/frozen_source_sha256.json +14 -0
  12. for Minegishi/reproducibility/load_checkpoint.py +35 -0
  13. for Minegishi/reproducibility/model/__init__.py +5 -0
  14. for Minegishi/reproducibility/model/gpt2.py +258 -0
  15. for Minegishi/reproducibility/model/latent_executor.py +213 -0
  16. for Minegishi/reproducibility/model/loop_gpt.py +216 -0
  17. for Minegishi/reproducibility/objectives.py +467 -0
  18. for Minegishi/reproducibility/probes.py +371 -0
  19. for Minegishi/reproducibility/recipe.py +20 -0
  20. for Minegishi/reproducibility/streams.py +27 -0
  21. for Minegishi/reproducibility/train_chain.py +435 -0
  22. for Minegishi/reproducibility/train_halt.py +996 -0
  23. for Minegishi/wait_denoising_fixed_t0/README.md +7 -0
  24. for Minegishi/wait_denoising_fixed_t0/checkpoints.json +237 -0
  25. for Minegishi/wait_denoising_fixed_t0/config.json +19 -0
  26. for Minegishi/wait_denoising_fixed_t0/final_eval.json +513 -0
  27. for Minegishi/wait_denoising_fixed_t0/manifest.json +188 -0
  28. for Minegishi/wait_denoising_fixed_t0/metrics.jsonl +23 -0
  29. for Minegishi/wait_denoising_fixed_t0/train_command.json +63 -0
  30. for Minegishi/wait_denoising_fixed_t0/train_log.jsonl +0 -0
for Minegishi/SHA256SUMS ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ 59f09ffdd9252ff8b9b21d93e77a778e52cce83763368d221c09fd19089abf67 README.md
2
+ 7e71a2c57ca912ad7d7c97f2346f351d718b1fb7bb4ad7002bfecb405c8ecfe0 analysis/RESULTS_zh.md
3
+ 9ab756490bd038ebc52ccb6037305e9f3a98c6d482bd66111e055857c1e5315a analysis/mechanism_summary.json
4
+ 8688d6465f0d9cbf23f89057eeeddee7aa1061fd221c017cfa6144079e2c8a28 analysis/plot_data.json
5
+ ff39620a2f2eb3a0a5d4240444e93d0976cb9245e2295ed38e914d4d18a30c5d analysis/protocol.json
6
+ 94d8a7b94667c880ef080689cf11bcb1806acc244c4cc68563ce0d7514da8302 checkpoint_index.json
7
+ 79aad031247e082548f869c0d331d80f508c2ccf884a98b7cb0d8bca16e715e2 figures/comparison_lines_d128.pdf
8
+ bc5621059fc8adbfe402a45cefacafca368f3c85e756aa8afee7b7a7f60d0bc3 figures/comparison_lines_d128.png
9
+ f849a1f4d4514f534bd72bb86525b41991001c2026daada8a06793d0793d4563 figures/comparison_lines_d128.svg
10
+ dfb8f8a2f670140384a8d5dc7193717ed50da6d15fc9f8d02491d189459e4497 figures/comparison_lines_d16.png
11
+ 62bb3bede48fd96a99566c1ba33f41f533b72c6e889eaf933657f274715d5307 figures/comparison_lines_d2.png
12
+ f2de623d67be685a16fe4db7131c97a09c037bd78566ececf821383497ba48c6 figures/comparison_lines_d32.png
13
+ 25c76733faae3d08480b54264f17bb209a2ee8b13d5b651152c6a864deefb233 figures/comparison_lines_d4.png
14
+ b3ee226596b966ff52399d6422441705d7422606dab22ff6ad5ee7055eda36b8 figures/comparison_lines_d64.png
15
+ 1e3a78c18dc49666fc81401fde247f8ab193bf03b0c18b8f5752f32b7879d0f2 figures/comparison_lines_d8.png
16
+ 62e0df1775cf79b5b5775a05bf6ded774b061f16339467c2c6b83bd27dc82332 figures/fixed_loop_error_curves.pdf
17
+ 00ba286ce961aa19f880ad3e13e8badf8c33393ed7a3c70f725daf00401876b0 figures/three_column_lines_all_depths.pdf
18
+ d6f34cfd11fa2a688e845594f9f01405b7a531464b67aeb572734c23b2341441 pure_wait_cs_random_t/README.md
19
+ 889bdf871f8ac9bac918348fb74c923a303a21d08bdc4137537af0feea97d466 pure_wait_cs_random_t/checkpoints.json
20
+ bcd81cb2fa52b7492682ba53d9a7d37826589dcab0bb45b5293a12352461a7a6 pure_wait_cs_random_t/ckpt_0001000.pt
21
+ 1c30c6a7c423437787176a45a591f1a2697fba869350bcf8d4702e08dfabe6c8 pure_wait_cs_random_t/ckpt_0002000.pt
22
+ b7c914f3f4d4bfc5e015604ed4b7a3cc6fd1401f1ff6f71e28a0a2672773e2c7 pure_wait_cs_random_t/ckpt_0005000.pt
23
+ d655c3be6a5529390af03141f72ac94acb13c6bafefe7e5b86e425b3b02254eb pure_wait_cs_random_t/ckpt_0010000.pt
24
+ bdd1bd1b285737eb756ffb03362024fd2abd253d4a1bd4218488a355a30613e3 pure_wait_cs_random_t/ckpt_0020000.pt
25
+ a07c56565f7e8167c104ac7f64ecf7e45c42cd5e1bbbd0b5b9ece863c47e6bef pure_wait_cs_random_t/ckpt_0030000.pt
26
+ b4a6ceac56316db92b9efc22b59536df5f611f4b879de3424a2a06b7ae2dc6a8 pure_wait_cs_random_t/ckpt_0040000.pt
27
+ 4c7edff25cdb2da0877178bad0c447460931bde8b914dcb2a6a264585c18b6f9 pure_wait_cs_random_t/ckpt_0050000.pt
28
+ eace02f9b72eec0a5dbe58ea7a96c71db2f4ea04d0ffdb56e863219fa64d39c9 pure_wait_cs_random_t/ckpt_0060000.pt
29
+ 4bc4e02c6fe8a79264411fb578f21de563d2079a0f8f6cbe7e1abe4a52999a1a pure_wait_cs_random_t/ckpt_0070000.pt
30
+ 492e306421ea28b94cf900033622633ca226089a8edc087934c248fe16012b1b pure_wait_cs_random_t/ckpt_0080000.pt
31
+ 8698c1108433fe3c38b6d519df122d76ab78a5be0c8e0f78b1b5fbd5af83c7cf pure_wait_cs_random_t/ckpt_0090000.pt
32
+ 728b00c9b1ebd5e4ddd9f714d8834baef48f008f3d66660775eae319183b58c9 pure_wait_cs_random_t/ckpt_0100000.pt
33
+ 86802449c3b6ffd8b354eed64bf0be79fa712c834178333218edffb26debbd5b pure_wait_cs_random_t/ckpt_0110000.pt
34
+ eadc632d68cf9d2e4457934b65f461b416affb35836874fcac74787a646d1ac9 pure_wait_cs_random_t/ckpt_0120000.pt
35
+ 440f34d69aba2d802480318a5ab97cdf2550a9bb4b9dfc3f0ae454077abb4925 pure_wait_cs_random_t/ckpt_0130000.pt
36
+ 9bc7ba9624294b4b91b133b2f18d87ae93292120ca5aa2913ccf95aa48192b49 pure_wait_cs_random_t/ckpt_0140000.pt
37
+ ab440a762c0db2d6d2bd00fca6e11a6275da8e6852e1636c15b1b60970eba247 pure_wait_cs_random_t/ckpt_0150000.pt
38
+ 78c238c7e4be4041be7fcc366e65aa37753d7ffec652411eece3bbbb900738cb pure_wait_cs_random_t/ckpt_0160000.pt
39
+ 743d16c538289bf6cf065f710c248e02d12fab22c5a206c164e57affd08168b7 pure_wait_cs_random_t/ckpt_0170000.pt
40
+ b0ff1c593de460f07f9e33c600e09bb719e84a5904b1da1471849d880df9d660 pure_wait_cs_random_t/ckpt_0180000.pt
41
+ add8fde1083bbebf1970cdd9c9535b47cca3834ec5031a08e59fd844374d2658 pure_wait_cs_random_t/ckpt_0190000.pt
42
+ 1f2b047da97b5628695fdb9446f4a75c2f6c1c225e350244866707a1dc482111 pure_wait_cs_random_t/ckpt_0200000.pt
43
+ 7c727a2a46fe5767503a2900919d8e7773995f83a01a1ef770c3442311ab9407 pure_wait_cs_random_t/config.json
44
+ a5253b83316becf2db1dbcc7a8d08870b5da413df71168ea969182acd159d2cd pure_wait_cs_random_t/final_eval.json
45
+ dbc44ea27d3468ea241ace1f840048c4f6975a1365e41812ef87762379582301 pure_wait_cs_random_t/manifest.json
46
+ 253aa941be1a3ccb1e146d775842c9185ec6b8c0b0ab04295d98d441141fd1b4 pure_wait_cs_random_t/metrics.jsonl
47
+ 1f88033cccf5bbdfdf4b3eceb7d4d6ffe0aef50e5539e6ab2670f9039452c941 pure_wait_cs_random_t/train_command.json
48
+ 33d684df0b5690dc38a5a4dae04320db03f98d2f9b6d32ce29d645e2258de900 pure_wait_cs_random_t/train_log.jsonl
49
+ 4e5f40b8dcd91769b0ada34c3050be09935f0d5e9225edfe95b9d2e7cb4e6e08 reproducibility/data/atomic_joint_2026-09-19/train_atomic.json
50
+ 506a7b3cb2e8aa1dbe74ef9fe21450d88a509b37e57b889c2b24922b0b369769 reproducibility/data/chain_loop_k8/audit.json
51
+ 0756f09e4ba515506f76881cc1bb17c5a9da8b2ee1ca387f20ce7f2caf9b9f13 reproducibility/data/chain_loop_k8/meta.json
52
+ a66eeb4380ef83fc3ae943bba0c7b9008915cab2252ce7d86cb05803784fb44d reproducibility/data/chain_loop_k8/tests.json
53
+ 3c49ed926a238159d6690918d1a0d4e7904c34af3e2f033df69aa627e443974f reproducibility/data/chain_loop_k8/train_atomic.json
54
+ e4263675cca07f9d815b97eb6dc054035eda2e49721e6c0d8e8047791391c2bc reproducibility/data/chain_loop_k8/train_d2.json
55
+ 4f53cda18c2baa0c0354bb5f9a3ecbe5ed12ab4d8e11ba873c2f11161202b945 reproducibility/data/chain_loop_k8/train_d3.json
56
+ c17bf311238b31781c3f438868038c11a4dd9a5b4dc5f2f3b9178b2e3e1d1ab4 reproducibility/data/chain_loop_k8/train_w2.json
57
+ 2184e30bb884158d8898ae080f8bed83537e39f5bb70b6246d1c9e41ab3971a4 reproducibility/data/chain_loop_k8/val_d2.json
58
+ 7bf2660c2f724f75a5d5f83c809ba8a71a10a8ffc5aa1b2f80d5e91745a06771 reproducibility/data/chain_loop_k8/val_w2.json
59
+ 4aaf5b155079d7a1a111fba39b04010624ad6724324ae5199c11ced14465888e reproducibility/data/chain_loop_k8/vocab.json
60
+ 6871b82c230566b824c0258cd40a0b5b8e6f72ae2578209ab5178a64fbc7871f reproducibility/eval_checkpoints.py
61
+ 7aa2779121f317384cc9e926ba36740033448f001e9d24853f7c4ee074da1962 reproducibility/extrapolation.py
62
+ 81a60bb97071d386a7494d322088752630b375c7855358654a1418b585100632 reproducibility/frozen_source_sha256.json
63
+ 3defa8d28e57bc3762835f2de3b66d91705c17154df3b7dd34ff2b3166d869a8 reproducibility/load_checkpoint.py
64
+ 2f2b8681adf78b17209129795e950a954e3e9b030e97bdcc9446aa5aa9999247 reproducibility/model/__init__.py
65
+ f97c63c37d88d34698ac239507033334679d70023671f8542bcb6508b364537f reproducibility/model/gpt2.py
66
+ 26d33ae9d6ddb49b2c7e467ff9ede4aa1c43517336c53a2f58377e155af006e6 reproducibility/model/latent_executor.py
67
+ 2245bd84895cb9cb587bf381e502d37cb52db6bd2ff382d892a4b31ffd20078b reproducibility/model/loop_gpt.py
68
+ f20f2507e995da5bccffa411edbfc8322b484f2f2d509d05ce374da282eab798 reproducibility/objectives.py
69
+ d32d3a7f5fa232c3cd7b23c28681e6cb10c869bdc45e303c95f7496522f8b10e reproducibility/probes.py
70
+ 3aae5bfe4189ec4323a59a935a8bb32c1a85a44e2c2d7f6cdd46915a5051cb41 reproducibility/recipe.py
71
+ 4bd4a83b21172b32079041117504020a5f8921da54340da28398b42539688c75 reproducibility/streams.py
72
+ a9fc664c3dd03c0cd11e0e32a0deca0ecc60ef1b9a37a16d23d696f8a6f380b9 reproducibility/train_chain.py
73
+ 813e8f785c04341a85be4a7a58d1c6663ae3a3d4948f9ad9d063a8a71ef4eb0b reproducibility/train_halt.py
74
+ 5c6c65d4ee944504ef7875acdfd1af6ee1f59f594880bb0b4454505e4642bc0f wait_denoising_fixed_t0/README.md
75
+ ed9da46086ad43f17fecf2c879a4d68ca2159c03418878364b5b269589ddcfaa wait_denoising_fixed_t0/checkpoints.json
76
+ b04cbfdb2a1e7a54e1c3857ef6d70d87542c02cd8f03baf4e874d8a998915e97 wait_denoising_fixed_t0/ckpt_0001000.pt
77
+ 09775249c367ccf150c3fd773f92b295743a30dc7db1bc910b9b1f29ba125f68 wait_denoising_fixed_t0/ckpt_0002000.pt
78
+ cbb6360b83c05c86ba437a856901801e36acf74776beefffe41db5a319305cd6 wait_denoising_fixed_t0/ckpt_0005000.pt
79
+ aefc1a2e8d13155cf68e5cdbe2b56648b666bede651114928563fe53aee388d8 wait_denoising_fixed_t0/ckpt_0010000.pt
80
+ 107b8652838126379d513366c4e684c2de800ed5582deae281f48b691966e505 wait_denoising_fixed_t0/ckpt_0020000.pt
81
+ 584e9f1b1fb3f1ce52fa7c95b5a8dc85c8696f4e9aa829c583dc415bd1144f79 wait_denoising_fixed_t0/ckpt_0030000.pt
82
+ 259bdfe027e50d18508c6e332add29ebca312fce38c21f6fce9bba2be6b86ce5 wait_denoising_fixed_t0/ckpt_0040000.pt
83
+ f8e4cd617629f536541f238f8e7df5e3b89f15db47c4c3025b5b7757a58a51fa wait_denoising_fixed_t0/ckpt_0050000.pt
84
+ f6dd0d27bcb48a93839be15dbccbf37b20db15b5bb5144ba8ef5eda581a0176b wait_denoising_fixed_t0/ckpt_0060000.pt
85
+ afe2b3f44168be01e54532886d2d1439bcb167f0f437c826c45e14668ce1caa0 wait_denoising_fixed_t0/ckpt_0070000.pt
86
+ 2800b11b460d0e5de02d7009c25a0e4d93b555c7c63f10a59e89739383043118 wait_denoising_fixed_t0/ckpt_0080000.pt
87
+ 61ba4ff1b03eed8710610004e9b1e68db3886813e2fbd68c83d58da444a9dafe wait_denoising_fixed_t0/ckpt_0090000.pt
88
+ e600093705a1477f6a092135438a4346dd1a4bc5c97430f2d179fdbb7b73a8d7 wait_denoising_fixed_t0/ckpt_0100000.pt
89
+ 1791c9e8a4be8f84d2fdd713b598439dd7ca3332ae0b9efdc20c53dbd86c07e6 wait_denoising_fixed_t0/ckpt_0110000.pt
90
+ c6767cfb80c10f38ca4bac059d5240211929dbe5a5cadb4b48ebb07a196566bf wait_denoising_fixed_t0/ckpt_0120000.pt
91
+ 1bf915b7a6cb858a51b534b9afd289421e0fdbb3c6f534831f0c4254143fb890 wait_denoising_fixed_t0/ckpt_0130000.pt
92
+ a00e672829198b3cba32d4241aa914b067cbd82f4f46f1a1219075f23ce22f8a wait_denoising_fixed_t0/ckpt_0140000.pt
93
+ c7b347a167d12485afddc0943dc92ce8326e3eb5a099ce5e4aafb4c2c0b81d8c wait_denoising_fixed_t0/ckpt_0150000.pt
94
+ bbd8ee69c630dd68e2ae25cb2eacb4bd1d3ffcadadd37a7bbc9938cc791ecb3a wait_denoising_fixed_t0/ckpt_0160000.pt
95
+ 76a4c1cd41d687ef06a1ccf6b67b2c56399c1ea2bf1103c51c144b3eddb6e20f wait_denoising_fixed_t0/ckpt_0170000.pt
96
+ d767ae2ced030111845d93e7f35587d7ecfc7d8ce43622d47df0d9c93638b191 wait_denoising_fixed_t0/ckpt_0180000.pt
97
+ a44ac51a1decb6919d48d1d02fcef7f4e525759760b913fffa12e064754d2c29 wait_denoising_fixed_t0/ckpt_0190000.pt
98
+ 6be182b9f912f741dfb0ba491879812e0813cc8adb264c76fa0cf5a9bbc64e22 wait_denoising_fixed_t0/ckpt_0200000.pt
99
+ 1dc960fbcbfb2132d1f246a6f5791c71e39c2b0e05c71d3e27c6a98f5bf66c6a wait_denoising_fixed_t0/config.json
100
+ c7684ab91f6c59cbee928221e67a3933953b55fa6b69fa410fc6181d9d714772 wait_denoising_fixed_t0/final_eval.json
101
+ d7379f92da8061a7c726b5bf711e7e1015816e7c595fe5a13a4988b836a210b5 wait_denoising_fixed_t0/manifest.json
102
+ d3f3494631778de417d383488b2ef5c929457de1114acc6935492d766dff69f3 wait_denoising_fixed_t0/metrics.jsonl
103
+ b6e35c69f5fec21a3c6da88147b5001517a3c60f9a1edf515a76eb0004da0a13 wait_denoising_fixed_t0/train_command.json
104
+ c42e3128392ab282563ebdda490fb3d06e3c79ef63f7131a1b2dbb51c4b963e4 wait_denoising_fixed_t0/train_log.jsonl
for Minegishi/reproducibility/data/chain_loop_k8/train_atomic.json ADDED
The diff for this file is too large to render. See raw diff
 
for Minegishi/reproducibility/data/chain_loop_k8/train_d2.json ADDED
The diff for this file is too large to render. See raw diff
 
for Minegishi/reproducibility/data/chain_loop_k8/train_d3.json ADDED
@@ -0,0 +1 @@
 
 
1
+ []
for Minegishi/reproducibility/data/chain_loop_k8/train_w2.json ADDED
The diff for this file is too large to render. See raw diff
 
for Minegishi/reproducibility/data/chain_loop_k8/val_d2.json ADDED
@@ -0,0 +1 @@
 
 
1
+ [{"kind": "d2", "x": 491, "r": 0, "s": 14}, {"kind": "d2", "x": 489, "r": 0, "s": 14}, {"kind": "d2", "x": 444, "r": 0, "s": 14}, {"kind": "d2", "x": 156, "r": 0, "s": 14}, {"kind": "d2", "x": 92, "r": 0, "s": 14}, {"kind": "d2", "x": 167, "r": 0, "s": 14}, {"kind": "d2", "x": 482, "r": 0, "s": 14}, {"kind": "d2", "x": 425, "r": 0, "s": 14}, {"kind": "d2", "x": 55, "r": 0, "s": 14}, {"kind": "d2", "x": 228, "r": 0, "s": 14}, {"kind": "d2", "x": 231, "r": 0, "s": 14}, {"kind": "d2", "x": 241, "r": 0, "s": 14}, {"kind": "d2", "x": 244, "r": 0, "s": 14}, {"kind": "d2", "x": 364, "r": 0, "s": 14}, {"kind": "d2", "x": 212, "r": 0, "s": 14}, {"kind": "d2", "x": 368, "r": 0, "s": 14}, {"kind": "d2", "x": 332, "r": 0, "s": 14}, {"kind": "d2", "x": 339, "r": 0, "s": 14}, {"kind": "d2", "x": 78, "r": 0, "s": 14}, {"kind": "d2", "x": 229, "r": 0, "s": 14}, {"kind": "d2", "x": 30, "r": 0, "s": 14}, {"kind": "d2", "x": 403, "r": 0, "s": 14}, {"kind": "d2", "x": 312, "r": 0, "s": 14}, {"kind": "d2", "x": 107, "r": 0, "s": 14}, {"kind": "d2", "x": 481, "r": 0, "s": 14}, {"kind": "d2", "x": 3, "r": 0, "s": 14}, {"kind": "d2", "x": 141, "r": 0, "s": 14}, {"kind": "d2", "x": 71, "r": 0, "s": 14}, {"kind": "d2", "x": 265, "r": 0, "s": 14}, {"kind": "d2", "x": 104, "r": 0, "s": 14}, {"kind": "d2", "x": 321, "r": 0, "s": 14}, {"kind": "d2", "x": 67, "r": 0, "s": 14}, {"kind": "d2", "x": 437, "r": 0, "s": 14}, {"kind": "d2", "x": 21, "r": 0, "s": 14}, {"kind": "d2", "x": 288, "r": 0, "s": 14}, {"kind": "d2", "x": 433, "r": 0, "s": 14}, {"kind": "d2", "x": 445, "r": 0, "s": 14}, {"kind": "d2", "x": 250, "r": 0, "s": 14}, {"kind": "d2", "x": 363, "r": 0, "s": 14}, {"kind": "d2", "x": 25, "r": 0, "s": 14}, {"kind": "d2", "x": 410, "r": 0, "s": 14}, {"kind": "d2", "x": 68, "r": 0, "s": 14}, {"kind": "d2", "x": 48, "r": 0, "s": 14}, {"kind": "d2", "x": 99, "r": 0, "s": 14}, {"kind": "d2", "x": 106, "r": 0, "s": 14}, {"kind": "d2", "x": 200, "r": 0, "s": 14}, {"kind": "d2", "x": 333, "r": 0, "s": 14}, {"kind": "d2", "x": 143, "r": 0, "s": 14}, {"kind": "d2", "x": 375, "r": 0, "s": 14}, {"kind": "d2", "x": 193, "r": 0, "s": 14}, {"kind": "d2", "x": 94, "r": 1, "s": 2}, {"kind": "d2", "x": 39, "r": 1, "s": 2}, {"kind": "d2", "x": 325, "r": 1, "s": 2}, {"kind": "d2", "x": 110, "r": 1, "s": 2}, {"kind": "d2", "x": 204, "r": 1, "s": 2}, {"kind": "d2", "x": 168, "r": 1, "s": 2}, {"kind": "d2", "x": 71, "r": 1, "s": 2}, {"kind": "d2", "x": 247, "r": 1, "s": 2}, {"kind": "d2", "x": 301, "r": 1, "s": 2}, {"kind": "d2", "x": 210, "r": 1, "s": 2}, {"kind": "d2", "x": 465, "r": 1, "s": 2}, {"kind": "d2", "x": 373, "r": 1, "s": 2}, {"kind": "d2", "x": 404, "r": 1, "s": 2}, {"kind": "d2", "x": 187, "r": 1, "s": 2}, {"kind": "d2", "x": 121, "r": 1, "s": 2}, {"kind": "d2", "x": 435, "r": 1, "s": 2}, {"kind": "d2", "x": 18, "r": 1, "s": 2}, {"kind": "d2", "x": 30, "r": 1, "s": 2}, {"kind": "d2", "x": 462, "r": 1, "s": 2}, {"kind": "d2", "x": 300, "r": 1, "s": 2}, {"kind": "d2", "x": 38, "r": 1, "s": 2}, {"kind": "d2", "x": 270, "r": 1, "s": 2}, {"kind": "d2", "x": 443, "r": 1, "s": 2}, {"kind": "d2", "x": 128, "r": 1, "s": 2}, {"kind": "d2", "x": 427, "r": 1, "s": 2}, {"kind": "d2", "x": 393, "r": 1, "s": 2}, {"kind": "d2", "x": 310, "r": 1, "s": 2}, {"kind": "d2", "x": 439, "r": 1, "s": 2}, {"kind": "d2", "x": 497, "r": 1, "s": 2}, {"kind": "d2", "x": 172, "r": 1, "s": 2}, {"kind": "d2", "x": 492, "r": 1, "s": 2}, {"kind": "d2", "x": 379, "r": 1, "s": 2}, {"kind": "d2", "x": 228, "r": 1, "s": 2}, {"kind": "d2", "x": 163, "r": 1, "s": 2}, {"kind": "d2", "x": 205, "r": 1, "s": 2}, {"kind": "d2", "x": 341, "r": 1, "s": 2}, {"kind": "d2", "x": 99, "r": 1, "s": 2}, {"kind": "d2", "x": 291, "r": 1, "s": 2}, {"kind": "d2", "x": 145, "r": 1, "s": 2}, {"kind": "d2", "x": 199, "r": 1, "s": 2}, {"kind": "d2", "x": 364, "r": 1, "s": 2}, {"kind": "d2", "x": 36, "r": 1, "s": 2}, {"kind": "d2", "x": 431, "r": 1, "s": 2}, {"kind": "d2", "x": 421, "r": 1, "s": 2}, {"kind": "d2", "x": 208, "r": 1, "s": 2}, {"kind": "d2", "x": 318, "r": 1, "s": 2}, {"kind": "d2", "x": 51, "r": 1, "s": 2}, {"kind": "d2", "x": 499, "r": 1, "s": 2}, {"kind": "d2", "x": 334, "r": 1, "s": 2}, {"kind": "d2", "x": 234, "r": 1, "s": 2}, {"kind": "d2", "x": 248, "r": 2, "s": 12}, {"kind": "d2", "x": 244, "r": 2, "s": 12}, {"kind": "d2", "x": 57, "r": 2, "s": 12}, {"kind": "d2", "x": 8, "r": 2, "s": 12}, {"kind": "d2", "x": 380, "r": 2, "s": 12}, {"kind": "d2", "x": 270, "r": 2, "s": 12}, {"kind": "d2", "x": 54, "r": 2, "s": 12}, {"kind": "d2", "x": 38, "r": 2, "s": 12}, {"kind": "d2", "x": 27, "r": 2, "s": 12}, {"kind": "d2", "x": 229, "r": 2, "s": 12}, {"kind": "d2", "x": 439, "r": 2, "s": 12}, {"kind": "d2", "x": 469, "r": 2, "s": 12}, {"kind": "d2", "x": 398, "r": 2, "s": 12}, {"kind": "d2", "x": 193, "r": 2, "s": 12}, {"kind": "d2", "x": 243, "r": 2, "s": 12}, {"kind": "d2", "x": 220, "r": 2, "s": 12}, {"kind": "d2", "x": 447, "r": 2, "s": 12}, {"kind": "d2", "x": 215, "r": 2, "s": 12}, {"kind": "d2", "x": 123, "r": 2, "s": 12}, {"kind": "d2", "x": 24, "r": 2, "s": 12}, {"kind": "d2", "x": 173, "r": 2, "s": 12}, {"kind": "d2", "x": 69, "r": 2, "s": 12}, {"kind": "d2", "x": 483, "r": 2, "s": 12}, {"kind": "d2", "x": 340, "r": 2, "s": 12}, {"kind": "d2", "x": 305, "r": 2, "s": 12}, {"kind": "d2", "x": 191, "r": 2, "s": 12}, {"kind": "d2", "x": 110, "r": 2, "s": 12}, {"kind": "d2", "x": 25, "r": 2, "s": 12}, {"kind": "d2", "x": 338, "r": 2, "s": 12}, {"kind": "d2", "x": 161, "r": 2, "s": 12}, {"kind": "d2", "x": 327, "r": 2, "s": 12}, {"kind": "d2", "x": 365, "r": 2, "s": 12}, {"kind": "d2", "x": 403, "r": 2, "s": 12}, {"kind": "d2", "x": 167, "r": 2, "s": 12}, {"kind": "d2", "x": 443, "r": 2, "s": 12}, {"kind": "d2", "x": 93, "r": 2, "s": 12}, {"kind": "d2", "x": 114, "r": 2, "s": 12}, {"kind": "d2", "x": 143, "r": 2, "s": 12}, {"kind": "d2", "x": 289, "r": 2, "s": 12}, {"kind": "d2", "x": 187, "r": 2, "s": 12}, {"kind": "d2", "x": 393, "r": 2, "s": 12}, {"kind": "d2", "x": 351, "r": 2, "s": 12}, {"kind": "d2", "x": 492, "r": 2, "s": 12}, {"kind": "d2", "x": 496, "r": 2, "s": 12}, {"kind": "d2", "x": 50, "r": 2, "s": 12}, {"kind": "d2", "x": 75, "r": 2, "s": 12}, {"kind": "d2", "x": 255, "r": 2, "s": 12}, {"kind": "d2", "x": 207, "r": 2, "s": 12}, {"kind": "d2", "x": 147, "r": 2, "s": 12}, {"kind": "d2", "x": 283, "r": 2, "s": 12}, {"kind": "d2", "x": 411, "r": 3, "s": 11}, {"kind": "d2", "x": 139, "r": 3, "s": 11}, {"kind": "d2", "x": 280, "r": 3, "s": 11}, {"kind": "d2", "x": 125, "r": 3, "s": 11}, {"kind": "d2", "x": 247, "r": 3, "s": 11}, {"kind": "d2", "x": 454, "r": 3, "s": 11}, {"kind": "d2", "x": 364, "r": 3, "s": 11}, {"kind": "d2", "x": 213, "r": 3, "s": 11}, {"kind": "d2", "x": 32, "r": 3, "s": 11}, {"kind": "d2", "x": 17, "r": 3, "s": 11}, {"kind": "d2", "x": 490, "r": 3, "s": 11}, {"kind": "d2", "x": 334, "r": 3, "s": 11}, {"kind": "d2", "x": 113, "r": 3, "s": 11}, {"kind": "d2", "x": 426, "r": 3, "s": 11}, {"kind": "d2", "x": 232, "r": 3, "s": 11}, {"kind": "d2", "x": 414, "r": 3, "s": 11}, {"kind": "d2", "x": 60, "r": 3, "s": 11}, {"kind": "d2", "x": 2, "r": 3, "s": 11}, {"kind": "d2", "x": 424, "r": 3, "s": 11}, {"kind": "d2", "x": 380, "r": 3, "s": 11}, {"kind": "d2", "x": 77, "r": 3, "s": 11}, {"kind": "d2", "x": 359, "r": 3, "s": 11}, {"kind": "d2", "x": 52, "r": 3, "s": 11}, {"kind": "d2", "x": 26, "r": 3, "s": 11}, {"kind": "d2", "x": 312, "r": 3, "s": 11}, {"kind": "d2", "x": 90, "r": 3, "s": 11}, {"kind": "d2", "x": 256, "r": 3, "s": 11}, {"kind": "d2", "x": 371, "r": 3, "s": 11}, {"kind": "d2", "x": 278, "r": 3, "s": 11}, {"kind": "d2", "x": 434, "r": 3, "s": 11}, {"kind": "d2", "x": 274, "r": 3, "s": 11}, {"kind": "d2", "x": 433, "r": 3, "s": 11}, {"kind": "d2", "x": 231, "r": 3, "s": 11}, {"kind": "d2", "x": 483, "r": 3, "s": 11}, {"kind": "d2", "x": 402, "r": 3, "s": 11}, {"kind": "d2", "x": 437, "r": 3, "s": 11}, {"kind": "d2", "x": 315, "r": 3, "s": 11}, {"kind": "d2", "x": 325, "r": 3, "s": 11}, {"kind": "d2", "x": 263, "r": 3, "s": 11}, {"kind": "d2", "x": 327, "r": 3, "s": 11}, {"kind": "d2", "x": 350, "r": 3, "s": 11}, {"kind": "d2", "x": 1, "r": 3, "s": 11}, {"kind": "d2", "x": 444, "r": 3, "s": 11}, {"kind": "d2", "x": 374, "r": 3, "s": 11}, {"kind": "d2", "x": 406, "r": 3, "s": 11}, {"kind": "d2", "x": 21, "r": 3, "s": 11}, {"kind": "d2", "x": 331, "r": 3, "s": 11}, {"kind": "d2", "x": 5, "r": 3, "s": 11}, {"kind": "d2", "x": 474, "r": 3, "s": 11}, {"kind": "d2", "x": 314, "r": 3, "s": 11}, {"kind": "d2", "x": 401, "r": 4, "s": 10}, {"kind": "d2", "x": 7, "r": 4, "s": 10}, {"kind": "d2", "x": 261, "r": 4, "s": 10}, {"kind": "d2", "x": 412, "r": 4, "s": 10}, {"kind": "d2", "x": 428, "r": 4, "s": 10}, {"kind": "d2", "x": 126, "r": 4, "s": 10}, {"kind": "d2", "x": 224, "r": 4, "s": 10}, {"kind": "d2", "x": 55, "r": 4, "s": 10}, {"kind": "d2", "x": 437, "r": 4, "s": 10}, {"kind": "d2", "x": 398, "r": 4, "s": 10}, {"kind": "d2", "x": 385, "r": 4, "s": 10}, {"kind": "d2", "x": 348, "r": 4, "s": 10}, {"kind": "d2", "x": 101, "r": 4, "s": 10}, {"kind": "d2", "x": 8, "r": 4, "s": 10}, {"kind": "d2", "x": 245, "r": 4, "s": 10}, {"kind": "d2", "x": 321, "r": 4, "s": 10}, {"kind": "d2", "x": 434, "r": 4, "s": 10}, {"kind": "d2", "x": 366, "r": 4, "s": 10}, {"kind": "d2", "x": 289, "r": 4, "s": 10}, {"kind": "d2", "x": 2, "r": 4, "s": 10}, {"kind": "d2", "x": 449, "r": 4, "s": 10}, {"kind": "d2", "x": 149, "r": 4, "s": 10}, {"kind": "d2", "x": 447, "r": 4, "s": 10}, {"kind": "d2", "x": 365, "r": 4, "s": 10}, {"kind": "d2", "x": 285, "r": 4, "s": 10}, {"kind": "d2", "x": 86, "r": 4, "s": 10}, {"kind": "d2", "x": 135, "r": 4, "s": 10}, {"kind": "d2", "x": 459, "r": 4, "s": 10}, {"kind": "d2", "x": 448, "r": 4, "s": 10}, {"kind": "d2", "x": 277, "r": 4, "s": 10}, {"kind": "d2", "x": 262, "r": 4, "s": 10}, {"kind": "d2", "x": 381, "r": 4, "s": 10}, {"kind": "d2", "x": 279, "r": 4, "s": 10}, {"kind": "d2", "x": 476, "r": 4, "s": 10}, {"kind": "d2", "x": 77, "r": 4, "s": 10}, {"kind": "d2", "x": 182, "r": 4, "s": 10}, {"kind": "d2", "x": 20, "r": 4, "s": 10}, {"kind": "d2", "x": 47, "r": 4, "s": 10}, {"kind": "d2", "x": 284, "r": 4, "s": 10}, {"kind": "d2", "x": 496, "r": 4, "s": 10}, {"kind": "d2", "x": 15, "r": 4, "s": 10}, {"kind": "d2", "x": 268, "r": 4, "s": 10}, {"kind": "d2", "x": 410, "r": 4, "s": 10}, {"kind": "d2", "x": 70, "r": 4, "s": 10}, {"kind": "d2", "x": 69, "r": 4, "s": 10}, {"kind": "d2", "x": 369, "r": 4, "s": 10}, {"kind": "d2", "x": 293, "r": 4, "s": 10}, {"kind": "d2", "x": 481, "r": 4, "s": 10}, {"kind": "d2", "x": 225, "r": 4, "s": 10}, {"kind": "d2", "x": 356, "r": 4, "s": 10}, {"kind": "d2", "x": 387, "r": 5, "s": 15}, {"kind": "d2", "x": 157, "r": 5, "s": 15}, {"kind": "d2", "x": 275, "r": 5, "s": 15}, {"kind": "d2", "x": 408, "r": 5, "s": 15}, {"kind": "d2", "x": 436, "r": 5, "s": 15}, {"kind": "d2", "x": 497, "r": 5, "s": 15}, {"kind": "d2", "x": 314, "r": 5, "s": 15}, {"kind": "d2", "x": 285, "r": 5, "s": 15}, {"kind": "d2", "x": 319, "r": 5, "s": 15}, {"kind": "d2", "x": 364, "r": 5, "s": 15}, {"kind": "d2", "x": 190, "r": 5, "s": 15}, {"kind": "d2", "x": 74, "r": 5, "s": 15}, {"kind": "d2", "x": 276, "r": 5, "s": 15}, {"kind": "d2", "x": 109, "r": 5, "s": 15}, {"kind": "d2", "x": 234, "r": 5, "s": 15}, {"kind": "d2", "x": 98, "r": 5, "s": 15}, {"kind": "d2", "x": 366, "r": 5, "s": 15}, {"kind": "d2", "x": 160, "r": 5, "s": 15}, {"kind": "d2", "x": 386, "r": 5, "s": 15}, {"kind": "d2", "x": 343, "r": 5, "s": 15}, {"kind": "d2", "x": 492, "r": 5, "s": 15}, {"kind": "d2", "x": 120, "r": 5, "s": 15}, {"kind": "d2", "x": 447, "r": 5, "s": 15}, {"kind": "d2", "x": 165, "r": 5, "s": 15}, {"kind": "d2", "x": 116, "r": 5, "s": 15}, {"kind": "d2", "x": 44, "r": 5, "s": 15}, {"kind": "d2", "x": 168, "r": 5, "s": 15}, {"kind": "d2", "x": 254, "r": 5, "s": 15}, {"kind": "d2", "x": 163, "r": 5, "s": 15}, {"kind": "d2", "x": 183, "r": 5, "s": 15}, {"kind": "d2", "x": 8, "r": 5, "s": 15}, {"kind": "d2", "x": 214, "r": 5, "s": 15}, {"kind": "d2", "x": 401, "r": 5, "s": 15}, {"kind": "d2", "x": 359, "r": 5, "s": 15}, {"kind": "d2", "x": 227, "r": 5, "s": 15}, {"kind": "d2", "x": 218, "r": 5, "s": 15}, {"kind": "d2", "x": 138, "r": 5, "s": 15}, {"kind": "d2", "x": 46, "r": 5, "s": 15}, {"kind": "d2", "x": 42, "r": 5, "s": 15}, {"kind": "d2", "x": 41, "r": 5, "s": 15}, {"kind": "d2", "x": 221, "r": 5, "s": 15}, {"kind": "d2", "x": 390, "r": 5, "s": 15}, {"kind": "d2", "x": 32, "r": 5, "s": 15}, {"kind": "d2", "x": 477, "r": 5, "s": 15}, {"kind": "d2", "x": 372, "r": 5, "s": 15}, {"kind": "d2", "x": 445, "r": 5, "s": 15}, {"kind": "d2", "x": 394, "r": 5, "s": 15}, {"kind": "d2", "x": 317, "r": 5, "s": 15}, {"kind": "d2", "x": 438, "r": 5, "s": 15}, {"kind": "d2", "x": 252, "r": 5, "s": 15}, {"kind": "d2", "x": 326, "r": 6, "s": 1}, {"kind": "d2", "x": 424, "r": 6, "s": 1}, {"kind": "d2", "x": 387, "r": 6, "s": 1}, {"kind": "d2", "x": 45, "r": 6, "s": 1}, {"kind": "d2", "x": 383, "r": 6, "s": 1}, {"kind": "d2", "x": 368, "r": 6, "s": 1}, {"kind": "d2", "x": 390, "r": 6, "s": 1}, {"kind": "d2", "x": 355, "r": 6, "s": 1}, {"kind": "d2", "x": 331, "r": 6, "s": 1}, {"kind": "d2", "x": 389, "r": 6, "s": 1}, {"kind": "d2", "x": 262, "r": 6, "s": 1}, {"kind": "d2", "x": 281, "r": 6, "s": 1}, {"kind": "d2", "x": 196, "r": 6, "s": 1}, {"kind": "d2", "x": 144, "r": 6, "s": 1}, {"kind": "d2", "x": 268, "r": 6, "s": 1}, {"kind": "d2", "x": 341, "r": 6, "s": 1}, {"kind": "d2", "x": 1, "r": 6, "s": 1}, {"kind": "d2", "x": 69, "r": 6, "s": 1}, {"kind": "d2", "x": 71, "r": 6, "s": 1}, {"kind": "d2", "x": 132, "r": 6, "s": 1}, {"kind": "d2", "x": 425, "r": 6, "s": 1}, {"kind": "d2", "x": 139, "r": 6, "s": 1}, {"kind": "d2", "x": 155, "r": 6, "s": 1}, {"kind": "d2", "x": 127, "r": 6, "s": 1}, {"kind": "d2", "x": 294, "r": 6, "s": 1}, {"kind": "d2", "x": 302, "r": 6, "s": 1}, {"kind": "d2", "x": 287, "r": 6, "s": 1}, {"kind": "d2", "x": 374, "r": 6, "s": 1}, {"kind": "d2", "x": 359, "r": 6, "s": 1}, {"kind": "d2", "x": 493, "r": 6, "s": 1}, {"kind": "d2", "x": 332, "r": 6, "s": 1}, {"kind": "d2", "x": 166, "r": 6, "s": 1}, {"kind": "d2", "x": 108, "r": 6, "s": 1}, {"kind": "d2", "x": 195, "r": 6, "s": 1}, {"kind": "d2", "x": 380, "r": 6, "s": 1}, {"kind": "d2", "x": 83, "r": 6, "s": 1}, {"kind": "d2", "x": 322, "r": 6, "s": 1}, {"kind": "d2", "x": 422, "r": 6, "s": 1}, {"kind": "d2", "x": 187, "r": 6, "s": 1}, {"kind": "d2", "x": 10, "r": 6, "s": 1}, {"kind": "d2", "x": 433, "r": 6, "s": 1}, {"kind": "d2", "x": 340, "r": 6, "s": 1}, {"kind": "d2", "x": 115, "r": 6, "s": 1}, {"kind": "d2", "x": 207, "r": 6, "s": 1}, {"kind": "d2", "x": 50, "r": 6, "s": 1}, {"kind": "d2", "x": 436, "r": 6, "s": 1}, {"kind": "d2", "x": 454, "r": 6, "s": 1}, {"kind": "d2", "x": 491, "r": 6, "s": 1}, {"kind": "d2", "x": 157, "r": 6, "s": 1}, {"kind": "d2", "x": 378, "r": 6, "s": 1}, {"kind": "d2", "x": 261, "r": 7, "s": 0}, {"kind": "d2", "x": 326, "r": 7, "s": 0}, {"kind": "d2", "x": 3, "r": 7, "s": 0}, {"kind": "d2", "x": 460, "r": 7, "s": 0}, {"kind": "d2", "x": 330, "r": 7, "s": 0}, {"kind": "d2", "x": 168, "r": 7, "s": 0}, {"kind": "d2", "x": 44, "r": 7, "s": 0}, {"kind": "d2", "x": 494, "r": 7, "s": 0}, {"kind": "d2", "x": 364, "r": 7, "s": 0}, {"kind": "d2", "x": 255, "r": 7, "s": 0}, {"kind": "d2", "x": 9, "r": 7, "s": 0}, {"kind": "d2", "x": 224, "r": 7, "s": 0}, {"kind": "d2", "x": 320, "r": 7, "s": 0}, {"kind": "d2", "x": 377, "r": 7, "s": 0}, {"kind": "d2", "x": 411, "r": 7, "s": 0}, {"kind": "d2", "x": 274, "r": 7, "s": 0}, {"kind": "d2", "x": 11, "r": 7, "s": 0}, {"kind": "d2", "x": 351, "r": 7, "s": 0}, {"kind": "d2", "x": 220, "r": 7, "s": 0}, {"kind": "d2", "x": 207, "r": 7, "s": 0}, {"kind": "d2", "x": 299, "r": 7, "s": 0}, {"kind": "d2", "x": 476, "r": 7, "s": 0}, {"kind": "d2", "x": 96, "r": 7, "s": 0}, {"kind": "d2", "x": 149, "r": 7, "s": 0}, {"kind": "d2", "x": 260, "r": 7, "s": 0}, {"kind": "d2", "x": 137, "r": 7, "s": 0}, {"kind": "d2", "x": 142, "r": 7, "s": 0}, {"kind": "d2", "x": 394, "r": 7, "s": 0}, {"kind": "d2", "x": 357, "r": 7, "s": 0}, {"kind": "d2", "x": 30, "r": 7, "s": 0}, {"kind": "d2", "x": 53, "r": 7, "s": 0}, {"kind": "d2", "x": 436, "r": 7, "s": 0}, {"kind": "d2", "x": 66, "r": 7, "s": 0}, {"kind": "d2", "x": 416, "r": 7, "s": 0}, {"kind": "d2", "x": 233, "r": 7, "s": 0}, {"kind": "d2", "x": 54, "r": 7, "s": 0}, {"kind": "d2", "x": 181, "r": 7, "s": 0}, {"kind": "d2", "x": 422, "r": 7, "s": 0}, {"kind": "d2", "x": 429, "r": 7, "s": 0}, {"kind": "d2", "x": 490, "r": 7, "s": 0}, {"kind": "d2", "x": 478, "r": 7, "s": 0}, {"kind": "d2", "x": 358, "r": 7, "s": 0}, {"kind": "d2", "x": 491, "r": 7, "s": 0}, {"kind": "d2", "x": 338, "r": 7, "s": 0}, {"kind": "d2", "x": 1, "r": 7, "s": 0}, {"kind": "d2", "x": 198, "r": 7, "s": 0}, {"kind": "d2", "x": 246, "r": 7, "s": 0}, {"kind": "d2", "x": 41, "r": 7, "s": 0}, {"kind": "d2", "x": 60, "r": 7, "s": 0}, {"kind": "d2", "x": 300, "r": 7, "s": 0}, {"kind": "d2", "x": 184, "r": 8, "s": 7}, {"kind": "d2", "x": 295, "r": 8, "s": 7}, {"kind": "d2", "x": 180, "r": 8, "s": 7}, {"kind": "d2", "x": 406, "r": 8, "s": 7}, {"kind": "d2", "x": 348, "r": 8, "s": 7}, {"kind": "d2", "x": 296, "r": 8, "s": 7}, {"kind": "d2", "x": 172, "r": 8, "s": 7}, {"kind": "d2", "x": 344, "r": 8, "s": 7}, {"kind": "d2", "x": 254, "r": 8, "s": 7}, {"kind": "d2", "x": 91, "r": 8, "s": 7}, {"kind": "d2", "x": 279, "r": 8, "s": 7}, {"kind": "d2", "x": 447, "r": 8, "s": 7}, {"kind": "d2", "x": 264, "r": 8, "s": 7}, {"kind": "d2", "x": 161, "r": 8, "s": 7}, {"kind": "d2", "x": 331, "r": 8, "s": 7}, {"kind": "d2", "x": 87, "r": 8, "s": 7}, {"kind": "d2", "x": 113, "r": 8, "s": 7}, {"kind": "d2", "x": 312, "r": 8, "s": 7}, {"kind": "d2", "x": 105, "r": 8, "s": 7}, {"kind": "d2", "x": 445, "r": 8, "s": 7}, {"kind": "d2", "x": 176, "r": 8, "s": 7}, {"kind": "d2", "x": 140, "r": 8, "s": 7}, {"kind": "d2", "x": 33, "r": 8, "s": 7}, {"kind": "d2", "x": 203, "r": 8, "s": 7}, {"kind": "d2", "x": 143, "r": 8, "s": 7}, {"kind": "d2", "x": 314, "r": 8, "s": 7}, {"kind": "d2", "x": 456, "r": 8, "s": 7}, {"kind": "d2", "x": 386, "r": 8, "s": 7}, {"kind": "d2", "x": 35, "r": 8, "s": 7}, {"kind": "d2", "x": 434, "r": 8, "s": 7}, {"kind": "d2", "x": 286, "r": 8, "s": 7}, {"kind": "d2", "x": 265, "r": 8, "s": 7}, {"kind": "d2", "x": 421, "r": 8, "s": 7}, {"kind": "d2", "x": 444, "r": 8, "s": 7}, {"kind": "d2", "x": 38, "r": 8, "s": 7}, {"kind": "d2", "x": 234, "r": 8, "s": 7}, {"kind": "d2", "x": 359, "r": 8, "s": 7}, {"kind": "d2", "x": 261, "r": 8, "s": 7}, {"kind": "d2", "x": 493, "r": 8, "s": 7}, {"kind": "d2", "x": 166, "r": 8, "s": 7}, {"kind": "d2", "x": 190, "r": 8, "s": 7}, {"kind": "d2", "x": 361, "r": 8, "s": 7}, {"kind": "d2", "x": 3, "r": 8, "s": 7}, {"kind": "d2", "x": 467, "r": 8, "s": 7}, {"kind": "d2", "x": 18, "r": 8, "s": 7}, {"kind": "d2", "x": 454, "r": 8, "s": 7}, {"kind": "d2", "x": 350, "r": 8, "s": 7}, {"kind": "d2", "x": 107, "r": 8, "s": 7}, {"kind": "d2", "x": 130, "r": 8, "s": 7}, {"kind": "d2", "x": 480, "r": 8, "s": 7}, {"kind": "d2", "x": 416, "r": 9, "s": 8}, {"kind": "d2", "x": 300, "r": 9, "s": 8}, {"kind": "d2", "x": 402, "r": 9, "s": 8}, {"kind": "d2", "x": 420, "r": 9, "s": 8}, {"kind": "d2", "x": 347, "r": 9, "s": 8}, {"kind": "d2", "x": 384, "r": 9, "s": 8}, {"kind": "d2", "x": 439, "r": 9, "s": 8}, {"kind": "d2", "x": 299, "r": 9, "s": 8}, {"kind": "d2", "x": 76, "r": 9, "s": 8}, {"kind": "d2", "x": 479, "r": 9, "s": 8}, {"kind": "d2", "x": 40, "r": 9, "s": 8}, {"kind": "d2", "x": 126, "r": 9, "s": 8}, {"kind": "d2", "x": 462, "r": 9, "s": 8}, {"kind": "d2", "x": 215, "r": 9, "s": 8}, {"kind": "d2", "x": 130, "r": 9, "s": 8}, {"kind": "d2", "x": 21, "r": 9, "s": 8}, {"kind": "d2", "x": 129, "r": 9, "s": 8}, {"kind": "d2", "x": 103, "r": 9, "s": 8}, {"kind": "d2", "x": 236, "r": 9, "s": 8}, {"kind": "d2", "x": 127, "r": 9, "s": 8}, {"kind": "d2", "x": 96, "r": 9, "s": 8}, {"kind": "d2", "x": 85, "r": 9, "s": 8}, {"kind": "d2", "x": 46, "r": 9, "s": 8}, {"kind": "d2", "x": 113, "r": 9, "s": 8}, {"kind": "d2", "x": 154, "r": 9, "s": 8}, {"kind": "d2", "x": 143, "r": 9, "s": 8}, {"kind": "d2", "x": 291, "r": 9, "s": 8}, {"kind": "d2", "x": 435, "r": 9, "s": 8}, {"kind": "d2", "x": 251, "r": 9, "s": 8}, {"kind": "d2", "x": 170, "r": 9, "s": 8}, {"kind": "d2", "x": 174, "r": 9, "s": 8}, {"kind": "d2", "x": 56, "r": 9, "s": 8}, {"kind": "d2", "x": 119, "r": 9, "s": 8}, {"kind": "d2", "x": 240, "r": 9, "s": 8}, {"kind": "d2", "x": 309, "r": 9, "s": 8}, {"kind": "d2", "x": 413, "r": 9, "s": 8}, {"kind": "d2", "x": 157, "r": 9, "s": 8}, {"kind": "d2", "x": 77, "r": 9, "s": 8}, {"kind": "d2", "x": 456, "r": 9, "s": 8}, {"kind": "d2", "x": 301, "r": 9, "s": 8}, {"kind": "d2", "x": 358, "r": 9, "s": 8}, {"kind": "d2", "x": 216, "r": 9, "s": 8}, {"kind": "d2", "x": 114, "r": 9, "s": 8}, {"kind": "d2", "x": 93, "r": 9, "s": 8}, {"kind": "d2", "x": 476, "r": 9, "s": 8}, {"kind": "d2", "x": 182, "r": 9, "s": 8}, {"kind": "d2", "x": 369, "r": 9, "s": 8}, {"kind": "d2", "x": 150, "r": 9, "s": 8}, {"kind": "d2", "x": 359, "r": 9, "s": 8}, {"kind": "d2", "x": 370, "r": 9, "s": 8}, {"kind": "d2", "x": 283, "r": 10, "s": 6}, {"kind": "d2", "x": 55, "r": 10, "s": 6}, {"kind": "d2", "x": 427, "r": 10, "s": 6}, {"kind": "d2", "x": 138, "r": 10, "s": 6}, {"kind": "d2", "x": 202, "r": 10, "s": 6}, {"kind": "d2", "x": 436, "r": 10, "s": 6}, {"kind": "d2", "x": 401, "r": 10, "s": 6}, {"kind": "d2", "x": 94, "r": 10, "s": 6}, {"kind": "d2", "x": 219, "r": 10, "s": 6}, {"kind": "d2", "x": 199, "r": 10, "s": 6}, {"kind": "d2", "x": 58, "r": 10, "s": 6}, {"kind": "d2", "x": 449, "r": 10, "s": 6}, {"kind": "d2", "x": 25, "r": 10, "s": 6}, {"kind": "d2", "x": 332, "r": 10, "s": 6}, {"kind": "d2", "x": 383, "r": 10, "s": 6}, {"kind": "d2", "x": 371, "r": 10, "s": 6}, {"kind": "d2", "x": 274, "r": 10, "s": 6}, {"kind": "d2", "x": 478, "r": 10, "s": 6}, {"kind": "d2", "x": 378, "r": 10, "s": 6}, {"kind": "d2", "x": 54, "r": 10, "s": 6}, {"kind": "d2", "x": 250, "r": 10, "s": 6}, {"kind": "d2", "x": 360, "r": 10, "s": 6}, {"kind": "d2", "x": 182, "r": 10, "s": 6}, {"kind": "d2", "x": 390, "r": 10, "s": 6}, {"kind": "d2", "x": 144, "r": 10, "s": 6}, {"kind": "d2", "x": 299, "r": 10, "s": 6}, {"kind": "d2", "x": 61, "r": 10, "s": 6}, {"kind": "d2", "x": 17, "r": 10, "s": 6}, {"kind": "d2", "x": 309, "r": 10, "s": 6}, {"kind": "d2", "x": 159, "r": 10, "s": 6}, {"kind": "d2", "x": 281, "r": 10, "s": 6}, {"kind": "d2", "x": 96, "r": 10, "s": 6}, {"kind": "d2", "x": 308, "r": 10, "s": 6}, {"kind": "d2", "x": 185, "r": 10, "s": 6}, {"kind": "d2", "x": 26, "r": 10, "s": 6}, {"kind": "d2", "x": 415, "r": 10, "s": 6}, {"kind": "d2", "x": 413, "r": 10, "s": 6}, {"kind": "d2", "x": 307, "r": 10, "s": 6}, {"kind": "d2", "x": 408, "r": 10, "s": 6}, {"kind": "d2", "x": 86, "r": 10, "s": 6}, {"kind": "d2", "x": 279, "r": 10, "s": 6}, {"kind": "d2", "x": 81, "r": 10, "s": 6}, {"kind": "d2", "x": 57, "r": 10, "s": 6}, {"kind": "d2", "x": 217, "r": 10, "s": 6}, {"kind": "d2", "x": 428, "r": 10, "s": 6}, {"kind": "d2", "x": 468, "r": 10, "s": 6}, {"kind": "d2", "x": 154, "r": 10, "s": 6}, {"kind": "d2", "x": 47, "r": 10, "s": 6}, {"kind": "d2", "x": 171, "r": 10, "s": 6}, {"kind": "d2", "x": 241, "r": 10, "s": 6}, {"kind": "d2", "x": 309, "r": 11, "s": 18}, {"kind": "d2", "x": 403, "r": 11, "s": 18}, {"kind": "d2", "x": 211, "r": 11, "s": 18}, {"kind": "d2", "x": 58, "r": 11, "s": 18}, {"kind": "d2", "x": 404, "r": 11, "s": 18}, {"kind": "d2", "x": 140, "r": 11, "s": 18}, {"kind": "d2", "x": 462, "r": 11, "s": 18}, {"kind": "d2", "x": 161, "r": 11, "s": 18}, {"kind": "d2", "x": 461, "r": 11, "s": 18}, {"kind": "d2", "x": 196, "r": 11, "s": 18}, {"kind": "d2", "x": 218, "r": 11, "s": 18}, {"kind": "d2", "x": 224, "r": 11, "s": 18}, {"kind": "d2", "x": 250, "r": 11, "s": 18}, {"kind": "d2", "x": 300, "r": 11, "s": 18}, {"kind": "d2", "x": 463, "r": 11, "s": 18}, {"kind": "d2", "x": 466, "r": 11, "s": 18}, {"kind": "d2", "x": 350, "r": 11, "s": 18}, {"kind": "d2", "x": 202, "r": 11, "s": 18}, {"kind": "d2", "x": 220, "r": 11, "s": 18}, {"kind": "d2", "x": 408, "r": 11, "s": 18}, {"kind": "d2", "x": 217, "r": 11, "s": 18}, {"kind": "d2", "x": 286, "r": 11, "s": 18}, {"kind": "d2", "x": 130, "r": 11, "s": 18}, {"kind": "d2", "x": 337, "r": 11, "s": 18}, {"kind": "d2", "x": 351, "r": 11, "s": 18}, {"kind": "d2", "x": 197, "r": 11, "s": 18}, {"kind": "d2", "x": 281, "r": 11, "s": 18}, {"kind": "d2", "x": 440, "r": 11, "s": 18}, {"kind": "d2", "x": 57, "r": 11, "s": 18}, {"kind": "d2", "x": 412, "r": 11, "s": 18}, {"kind": "d2", "x": 195, "r": 11, "s": 18}, {"kind": "d2", "x": 336, "r": 11, "s": 18}, {"kind": "d2", "x": 157, "r": 11, "s": 18}, {"kind": "d2", "x": 387, "r": 11, "s": 18}, {"kind": "d2", "x": 107, "r": 11, "s": 18}, {"kind": "d2", "x": 274, "r": 11, "s": 18}, {"kind": "d2", "x": 490, "r": 11, "s": 18}, {"kind": "d2", "x": 292, "r": 11, "s": 18}, {"kind": "d2", "x": 60, "r": 11, "s": 18}, {"kind": "d2", "x": 457, "r": 11, "s": 18}, {"kind": "d2", "x": 471, "r": 11, "s": 18}, {"kind": "d2", "x": 182, "r": 11, "s": 18}, {"kind": "d2", "x": 273, "r": 11, "s": 18}, {"kind": "d2", "x": 353, "r": 11, "s": 18}, {"kind": "d2", "x": 151, "r": 11, "s": 18}, {"kind": "d2", "x": 458, "r": 11, "s": 18}, {"kind": "d2", "x": 44, "r": 11, "s": 18}, {"kind": "d2", "x": 388, "r": 11, "s": 18}, {"kind": "d2", "x": 327, "r": 11, "s": 18}, {"kind": "d2", "x": 184, "r": 11, "s": 18}, {"kind": "d2", "x": 392, "r": 12, "s": 5}, {"kind": "d2", "x": 298, "r": 12, "s": 5}, {"kind": "d2", "x": 48, "r": 12, "s": 5}, {"kind": "d2", "x": 31, "r": 12, "s": 5}, {"kind": "d2", "x": 111, "r": 12, "s": 5}, {"kind": "d2", "x": 129, "r": 12, "s": 5}, {"kind": "d2", "x": 297, "r": 12, "s": 5}, {"kind": "d2", "x": 423, "r": 12, "s": 5}, {"kind": "d2", "x": 69, "r": 12, "s": 5}, {"kind": "d2", "x": 246, "r": 12, "s": 5}, {"kind": "d2", "x": 98, "r": 12, "s": 5}, {"kind": "d2", "x": 86, "r": 12, "s": 5}, {"kind": "d2", "x": 118, "r": 12, "s": 5}, {"kind": "d2", "x": 92, "r": 12, "s": 5}, {"kind": "d2", "x": 486, "r": 12, "s": 5}, {"kind": "d2", "x": 479, "r": 12, "s": 5}, {"kind": "d2", "x": 335, "r": 12, "s": 5}, {"kind": "d2", "x": 327, "r": 12, "s": 5}, {"kind": "d2", "x": 352, "r": 12, "s": 5}, {"kind": "d2", "x": 197, "r": 12, "s": 5}, {"kind": "d2", "x": 74, "r": 12, "s": 5}, {"kind": "d2", "x": 435, "r": 12, "s": 5}, {"kind": "d2", "x": 202, "r": 12, "s": 5}, {"kind": "d2", "x": 331, "r": 12, "s": 5}, {"kind": "d2", "x": 58, "r": 12, "s": 5}, {"kind": "d2", "x": 230, "r": 12, "s": 5}, {"kind": "d2", "x": 84, "r": 12, "s": 5}, {"kind": "d2", "x": 213, "r": 12, "s": 5}, {"kind": "d2", "x": 383, "r": 12, "s": 5}, {"kind": "d2", "x": 198, "r": 12, "s": 5}, {"kind": "d2", "x": 319, "r": 12, "s": 5}, {"kind": "d2", "x": 338, "r": 12, "s": 5}, {"kind": "d2", "x": 216, "r": 12, "s": 5}, {"kind": "d2", "x": 156, "r": 12, "s": 5}, {"kind": "d2", "x": 248, "r": 12, "s": 5}, {"kind": "d2", "x": 470, "r": 12, "s": 5}, {"kind": "d2", "x": 283, "r": 12, "s": 5}, {"kind": "d2", "x": 89, "r": 12, "s": 5}, {"kind": "d2", "x": 209, "r": 12, "s": 5}, {"kind": "d2", "x": 487, "r": 12, "s": 5}, {"kind": "d2", "x": 37, "r": 12, "s": 5}, {"kind": "d2", "x": 495, "r": 12, "s": 5}, {"kind": "d2", "x": 313, "r": 12, "s": 5}, {"kind": "d2", "x": 32, "r": 12, "s": 5}, {"kind": "d2", "x": 154, "r": 12, "s": 5}, {"kind": "d2", "x": 196, "r": 12, "s": 5}, {"kind": "d2", "x": 73, "r": 12, "s": 5}, {"kind": "d2", "x": 458, "r": 12, "s": 5}, {"kind": "d2", "x": 420, "r": 12, "s": 5}, {"kind": "d2", "x": 498, "r": 12, "s": 5}, {"kind": "d2", "x": 271, "r": 13, "s": 9}, {"kind": "d2", "x": 86, "r": 13, "s": 9}, {"kind": "d2", "x": 217, "r": 13, "s": 9}, {"kind": "d2", "x": 236, "r": 13, "s": 9}, {"kind": "d2", "x": 183, "r": 13, "s": 9}, {"kind": "d2", "x": 478, "r": 13, "s": 9}, {"kind": "d2", "x": 139, "r": 13, "s": 9}, {"kind": "d2", "x": 248, "r": 13, "s": 9}, {"kind": "d2", "x": 203, "r": 13, "s": 9}, {"kind": "d2", "x": 317, "r": 13, "s": 9}, {"kind": "d2", "x": 253, "r": 13, "s": 9}, {"kind": "d2", "x": 215, "r": 13, "s": 9}, {"kind": "d2", "x": 457, "r": 13, "s": 9}, {"kind": "d2", "x": 172, "r": 13, "s": 9}, {"kind": "d2", "x": 264, "r": 13, "s": 9}, {"kind": "d2", "x": 335, "r": 13, "s": 9}, {"kind": "d2", "x": 123, "r": 13, "s": 9}, {"kind": "d2", "x": 258, "r": 13, "s": 9}, {"kind": "d2", "x": 409, "r": 13, "s": 9}, {"kind": "d2", "x": 375, "r": 13, "s": 9}, {"kind": "d2", "x": 477, "r": 13, "s": 9}, {"kind": "d2", "x": 404, "r": 13, "s": 9}, {"kind": "d2", "x": 422, "r": 13, "s": 9}, {"kind": "d2", "x": 129, "r": 13, "s": 9}, {"kind": "d2", "x": 194, "r": 13, "s": 9}, {"kind": "d2", "x": 49, "r": 13, "s": 9}, {"kind": "d2", "x": 334, "r": 13, "s": 9}, {"kind": "d2", "x": 55, "r": 13, "s": 9}, {"kind": "d2", "x": 489, "r": 13, "s": 9}, {"kind": "d2", "x": 252, "r": 13, "s": 9}, {"kind": "d2", "x": 352, "r": 13, "s": 9}, {"kind": "d2", "x": 272, "r": 13, "s": 9}, {"kind": "d2", "x": 99, "r": 13, "s": 9}, {"kind": "d2", "x": 287, "r": 13, "s": 9}, {"kind": "d2", "x": 113, "r": 13, "s": 9}, {"kind": "d2", "x": 163, "r": 13, "s": 9}, {"kind": "d2", "x": 358, "r": 13, "s": 9}, {"kind": "d2", "x": 179, "r": 13, "s": 9}, {"kind": "d2", "x": 73, "r": 13, "s": 9}, {"kind": "d2", "x": 475, "r": 13, "s": 9}, {"kind": "d2", "x": 278, "r": 13, "s": 9}, {"kind": "d2", "x": 220, "r": 13, "s": 9}, {"kind": "d2", "x": 314, "r": 13, "s": 9}, {"kind": "d2", "x": 24, "r": 13, "s": 9}, {"kind": "d2", "x": 294, "r": 13, "s": 9}, {"kind": "d2", "x": 84, "r": 13, "s": 9}, {"kind": "d2", "x": 328, "r": 13, "s": 9}, {"kind": "d2", "x": 325, "r": 13, "s": 9}, {"kind": "d2", "x": 306, "r": 13, "s": 9}, {"kind": "d2", "x": 498, "r": 13, "s": 9}, {"kind": "d2", "x": 473, "r": 14, "s": 4}, {"kind": "d2", "x": 400, "r": 14, "s": 4}, {"kind": "d2", "x": 286, "r": 14, "s": 4}, {"kind": "d2", "x": 44, "r": 14, "s": 4}, {"kind": "d2", "x": 254, "r": 14, "s": 4}, {"kind": "d2", "x": 50, "r": 14, "s": 4}, {"kind": "d2", "x": 74, "r": 14, "s": 4}, {"kind": "d2", "x": 410, "r": 14, "s": 4}, {"kind": "d2", "x": 77, "r": 14, "s": 4}, {"kind": "d2", "x": 235, "r": 14, "s": 4}, {"kind": "d2", "x": 382, "r": 14, "s": 4}, {"kind": "d2", "x": 12, "r": 14, "s": 4}, {"kind": "d2", "x": 59, "r": 14, "s": 4}, {"kind": "d2", "x": 125, "r": 14, "s": 4}, {"kind": "d2", "x": 181, "r": 14, "s": 4}, {"kind": "d2", "x": 65, "r": 14, "s": 4}, {"kind": "d2", "x": 191, "r": 14, "s": 4}, {"kind": "d2", "x": 208, "r": 14, "s": 4}, {"kind": "d2", "x": 476, "r": 14, "s": 4}, {"kind": "d2", "x": 240, "r": 14, "s": 4}, {"kind": "d2", "x": 285, "r": 14, "s": 4}, {"kind": "d2", "x": 152, "r": 14, "s": 4}, {"kind": "d2", "x": 215, "r": 14, "s": 4}, {"kind": "d2", "x": 304, "r": 14, "s": 4}, {"kind": "d2", "x": 108, "r": 14, "s": 4}, {"kind": "d2", "x": 478, "r": 14, "s": 4}, {"kind": "d2", "x": 110, "r": 14, "s": 4}, {"kind": "d2", "x": 126, "r": 14, "s": 4}, {"kind": "d2", "x": 120, "r": 14, "s": 4}, {"kind": "d2", "x": 237, "r": 14, "s": 4}, {"kind": "d2", "x": 499, "r": 14, "s": 4}, {"kind": "d2", "x": 178, "r": 14, "s": 4}, {"kind": "d2", "x": 166, "r": 14, "s": 4}, {"kind": "d2", "x": 142, "r": 14, "s": 4}, {"kind": "d2", "x": 261, "r": 14, "s": 4}, {"kind": "d2", "x": 253, "r": 14, "s": 4}, {"kind": "d2", "x": 274, "r": 14, "s": 4}, {"kind": "d2", "x": 349, "r": 14, "s": 4}, {"kind": "d2", "x": 277, "r": 14, "s": 4}, {"kind": "d2", "x": 51, "r": 14, "s": 4}, {"kind": "d2", "x": 415, "r": 14, "s": 4}, {"kind": "d2", "x": 168, "r": 14, "s": 4}, {"kind": "d2", "x": 2, "r": 14, "s": 4}, {"kind": "d2", "x": 433, "r": 14, "s": 4}, {"kind": "d2", "x": 118, "r": 14, "s": 4}, {"kind": "d2", "x": 376, "r": 14, "s": 4}, {"kind": "d2", "x": 318, "r": 14, "s": 4}, {"kind": "d2", "x": 294, "r": 14, "s": 4}, {"kind": "d2", "x": 403, "r": 14, "s": 4}, {"kind": "d2", "x": 200, "r": 14, "s": 4}, {"kind": "d2", "x": 70, "r": 15, "s": 3}, {"kind": "d2", "x": 326, "r": 15, "s": 3}, {"kind": "d2", "x": 384, "r": 15, "s": 3}, {"kind": "d2", "x": 118, "r": 15, "s": 3}, {"kind": "d2", "x": 425, "r": 15, "s": 3}, {"kind": "d2", "x": 446, "r": 15, "s": 3}, {"kind": "d2", "x": 54, "r": 15, "s": 3}, {"kind": "d2", "x": 287, "r": 15, "s": 3}, {"kind": "d2", "x": 16, "r": 15, "s": 3}, {"kind": "d2", "x": 156, "r": 15, "s": 3}, {"kind": "d2", "x": 11, "r": 15, "s": 3}, {"kind": "d2", "x": 434, "r": 15, "s": 3}, {"kind": "d2", "x": 224, "r": 15, "s": 3}, {"kind": "d2", "x": 141, "r": 15, "s": 3}, {"kind": "d2", "x": 37, "r": 15, "s": 3}, {"kind": "d2", "x": 182, "r": 15, "s": 3}, {"kind": "d2", "x": 128, "r": 15, "s": 3}, {"kind": "d2", "x": 139, "r": 15, "s": 3}, {"kind": "d2", "x": 448, "r": 15, "s": 3}, {"kind": "d2", "x": 223, "r": 15, "s": 3}, {"kind": "d2", "x": 450, "r": 15, "s": 3}, {"kind": "d2", "x": 480, "r": 15, "s": 3}, {"kind": "d2", "x": 456, "r": 15, "s": 3}, {"kind": "d2", "x": 362, "r": 15, "s": 3}, {"kind": "d2", "x": 179, "r": 15, "s": 3}, {"kind": "d2", "x": 474, "r": 15, "s": 3}, {"kind": "d2", "x": 469, "r": 15, "s": 3}, {"kind": "d2", "x": 95, "r": 15, "s": 3}, {"kind": "d2", "x": 32, "r": 15, "s": 3}, {"kind": "d2", "x": 77, "r": 15, "s": 3}, {"kind": "d2", "x": 348, "r": 15, "s": 3}, {"kind": "d2", "x": 488, "r": 15, "s": 3}, {"kind": "d2", "x": 199, "r": 15, "s": 3}, {"kind": "d2", "x": 80, "r": 15, "s": 3}, {"kind": "d2", "x": 447, "r": 15, "s": 3}, {"kind": "d2", "x": 323, "r": 15, "s": 3}, {"kind": "d2", "x": 200, "r": 15, "s": 3}, {"kind": "d2", "x": 72, "r": 15, "s": 3}, {"kind": "d2", "x": 300, "r": 15, "s": 3}, {"kind": "d2", "x": 106, "r": 15, "s": 3}, {"kind": "d2", "x": 274, "r": 15, "s": 3}, {"kind": "d2", "x": 231, "r": 15, "s": 3}, {"kind": "d2", "x": 365, "r": 15, "s": 3}, {"kind": "d2", "x": 267, "r": 15, "s": 3}, {"kind": "d2", "x": 256, "r": 15, "s": 3}, {"kind": "d2", "x": 53, "r": 15, "s": 3}, {"kind": "d2", "x": 247, "r": 15, "s": 3}, {"kind": "d2", "x": 401, "r": 15, "s": 3}, {"kind": "d2", "x": 276, "r": 15, "s": 3}, {"kind": "d2", "x": 335, "r": 15, "s": 3}, {"kind": "d2", "x": 79, "r": 16, "s": 17}, {"kind": "d2", "x": 34, "r": 16, "s": 17}, {"kind": "d2", "x": 478, "r": 16, "s": 17}, {"kind": "d2", "x": 117, "r": 16, "s": 17}, {"kind": "d2", "x": 101, "r": 16, "s": 17}, {"kind": "d2", "x": 236, "r": 16, "s": 17}, {"kind": "d2", "x": 4, "r": 16, "s": 17}, {"kind": "d2", "x": 45, "r": 16, "s": 17}, {"kind": "d2", "x": 57, "r": 16, "s": 17}, {"kind": "d2", "x": 120, "r": 16, "s": 17}, {"kind": "d2", "x": 211, "r": 16, "s": 17}, {"kind": "d2", "x": 230, "r": 16, "s": 17}, {"kind": "d2", "x": 182, "r": 16, "s": 17}, {"kind": "d2", "x": 366, "r": 16, "s": 17}, {"kind": "d2", "x": 247, "r": 16, "s": 17}, {"kind": "d2", "x": 170, "r": 16, "s": 17}, {"kind": "d2", "x": 174, "r": 16, "s": 17}, {"kind": "d2", "x": 294, "r": 16, "s": 17}, {"kind": "d2", "x": 345, "r": 16, "s": 17}, {"kind": "d2", "x": 183, "r": 16, "s": 17}, {"kind": "d2", "x": 468, "r": 16, "s": 17}, {"kind": "d2", "x": 299, "r": 16, "s": 17}, {"kind": "d2", "x": 141, "r": 16, "s": 17}, {"kind": "d2", "x": 140, "r": 16, "s": 17}, {"kind": "d2", "x": 308, "r": 16, "s": 17}, {"kind": "d2", "x": 305, "r": 16, "s": 17}, {"kind": "d2", "x": 421, "r": 16, "s": 17}, {"kind": "d2", "x": 307, "r": 16, "s": 17}, {"kind": "d2", "x": 234, "r": 16, "s": 17}, {"kind": "d2", "x": 29, "r": 16, "s": 17}, {"kind": "d2", "x": 21, "r": 16, "s": 17}, {"kind": "d2", "x": 36, "r": 16, "s": 17}, {"kind": "d2", "x": 363, "r": 16, "s": 17}, {"kind": "d2", "x": 496, "r": 16, "s": 17}, {"kind": "d2", "x": 46, "r": 16, "s": 17}, {"kind": "d2", "x": 431, "r": 16, "s": 17}, {"kind": "d2", "x": 24, "r": 16, "s": 17}, {"kind": "d2", "x": 53, "r": 16, "s": 17}, {"kind": "d2", "x": 488, "r": 16, "s": 17}, {"kind": "d2", "x": 445, "r": 16, "s": 17}, {"kind": "d2", "x": 30, "r": 16, "s": 17}, {"kind": "d2", "x": 41, "r": 16, "s": 17}, {"kind": "d2", "x": 465, "r": 16, "s": 17}, {"kind": "d2", "x": 340, "r": 16, "s": 17}, {"kind": "d2", "x": 440, "r": 16, "s": 17}, {"kind": "d2", "x": 470, "r": 16, "s": 17}, {"kind": "d2", "x": 387, "r": 16, "s": 17}, {"kind": "d2", "x": 68, "r": 16, "s": 17}, {"kind": "d2", "x": 59, "r": 16, "s": 17}, {"kind": "d2", "x": 413, "r": 16, "s": 17}, {"kind": "d2", "x": 1, "r": 17, "s": 13}, {"kind": "d2", "x": 41, "r": 17, "s": 13}, {"kind": "d2", "x": 78, "r": 17, "s": 13}, {"kind": "d2", "x": 337, "r": 17, "s": 13}, {"kind": "d2", "x": 461, "r": 17, "s": 13}, {"kind": "d2", "x": 446, "r": 17, "s": 13}, {"kind": "d2", "x": 247, "r": 17, "s": 13}, {"kind": "d2", "x": 283, "r": 17, "s": 13}, {"kind": "d2", "x": 94, "r": 17, "s": 13}, {"kind": "d2", "x": 90, "r": 17, "s": 13}, {"kind": "d2", "x": 343, "r": 17, "s": 13}, {"kind": "d2", "x": 144, "r": 17, "s": 13}, {"kind": "d2", "x": 119, "r": 17, "s": 13}, {"kind": "d2", "x": 22, "r": 17, "s": 13}, {"kind": "d2", "x": 12, "r": 17, "s": 13}, {"kind": "d2", "x": 92, "r": 17, "s": 13}, {"kind": "d2", "x": 448, "r": 17, "s": 13}, {"kind": "d2", "x": 139, "r": 17, "s": 13}, {"kind": "d2", "x": 170, "r": 17, "s": 13}, {"kind": "d2", "x": 227, "r": 17, "s": 13}, {"kind": "d2", "x": 290, "r": 17, "s": 13}, {"kind": "d2", "x": 299, "r": 17, "s": 13}, {"kind": "d2", "x": 338, "r": 17, "s": 13}, {"kind": "d2", "x": 38, "r": 17, "s": 13}, {"kind": "d2", "x": 35, "r": 17, "s": 13}, {"kind": "d2", "x": 273, "r": 17, "s": 13}, {"kind": "d2", "x": 376, "r": 17, "s": 13}, {"kind": "d2", "x": 24, "r": 17, "s": 13}, {"kind": "d2", "x": 259, "r": 17, "s": 13}, {"kind": "d2", "x": 73, "r": 17, "s": 13}, {"kind": "d2", "x": 75, "r": 17, "s": 13}, {"kind": "d2", "x": 159, "r": 17, "s": 13}, {"kind": "d2", "x": 361, "r": 17, "s": 13}, {"kind": "d2", "x": 13, "r": 17, "s": 13}, {"kind": "d2", "x": 407, "r": 17, "s": 13}, {"kind": "d2", "x": 107, "r": 17, "s": 13}, {"kind": "d2", "x": 271, "r": 17, "s": 13}, {"kind": "d2", "x": 389, "r": 17, "s": 13}, {"kind": "d2", "x": 482, "r": 17, "s": 13}, {"kind": "d2", "x": 317, "r": 17, "s": 13}, {"kind": "d2", "x": 86, "r": 17, "s": 13}, {"kind": "d2", "x": 207, "r": 17, "s": 13}, {"kind": "d2", "x": 42, "r": 17, "s": 13}, {"kind": "d2", "x": 125, "r": 17, "s": 13}, {"kind": "d2", "x": 45, "r": 17, "s": 13}, {"kind": "d2", "x": 133, "r": 17, "s": 13}, {"kind": "d2", "x": 468, "r": 17, "s": 13}, {"kind": "d2", "x": 161, "r": 17, "s": 13}, {"kind": "d2", "x": 267, "r": 17, "s": 13}, {"kind": "d2", "x": 425, "r": 17, "s": 13}, {"kind": "d2", "x": 442, "r": 18, "s": 19}, {"kind": "d2", "x": 252, "r": 18, "s": 19}, {"kind": "d2", "x": 217, "r": 18, "s": 19}, {"kind": "d2", "x": 419, "r": 18, "s": 19}, {"kind": "d2", "x": 487, "r": 18, "s": 19}, {"kind": "d2", "x": 239, "r": 18, "s": 19}, {"kind": "d2", "x": 164, "r": 18, "s": 19}, {"kind": "d2", "x": 11, "r": 18, "s": 19}, {"kind": "d2", "x": 40, "r": 18, "s": 19}, {"kind": "d2", "x": 126, "r": 18, "s": 19}, {"kind": "d2", "x": 155, "r": 18, "s": 19}, {"kind": "d2", "x": 388, "r": 18, "s": 19}, {"kind": "d2", "x": 258, "r": 18, "s": 19}, {"kind": "d2", "x": 249, "r": 18, "s": 19}, {"kind": "d2", "x": 488, "r": 18, "s": 19}, {"kind": "d2", "x": 156, "r": 18, "s": 19}, {"kind": "d2", "x": 296, "r": 18, "s": 19}, {"kind": "d2", "x": 86, "r": 18, "s": 19}, {"kind": "d2", "x": 46, "r": 18, "s": 19}, {"kind": "d2", "x": 79, "r": 18, "s": 19}, {"kind": "d2", "x": 255, "r": 18, "s": 19}, {"kind": "d2", "x": 262, "r": 18, "s": 19}, {"kind": "d2", "x": 423, "r": 18, "s": 19}, {"kind": "d2", "x": 145, "r": 18, "s": 19}, {"kind": "d2", "x": 85, "r": 18, "s": 19}, {"kind": "d2", "x": 2, "r": 18, "s": 19}, {"kind": "d2", "x": 8, "r": 18, "s": 19}, {"kind": "d2", "x": 65, "r": 18, "s": 19}, {"kind": "d2", "x": 163, "r": 18, "s": 19}, {"kind": "d2", "x": 140, "r": 18, "s": 19}, {"kind": "d2", "x": 342, "r": 18, "s": 19}, {"kind": "d2", "x": 347, "r": 18, "s": 19}, {"kind": "d2", "x": 15, "r": 18, "s": 19}, {"kind": "d2", "x": 149, "r": 18, "s": 19}, {"kind": "d2", "x": 311, "r": 18, "s": 19}, {"kind": "d2", "x": 160, "r": 18, "s": 19}, {"kind": "d2", "x": 251, "r": 18, "s": 19}, {"kind": "d2", "x": 375, "r": 18, "s": 19}, {"kind": "d2", "x": 104, "r": 18, "s": 19}, {"kind": "d2", "x": 190, "r": 18, "s": 19}, {"kind": "d2", "x": 246, "r": 18, "s": 19}, {"kind": "d2", "x": 135, "r": 18, "s": 19}, {"kind": "d2", "x": 431, "r": 18, "s": 19}, {"kind": "d2", "x": 308, "r": 18, "s": 19}, {"kind": "d2", "x": 254, "r": 18, "s": 19}, {"kind": "d2", "x": 316, "r": 18, "s": 19}, {"kind": "d2", "x": 60, "r": 18, "s": 19}, {"kind": "d2", "x": 443, "r": 18, "s": 19}, {"kind": "d2", "x": 309, "r": 18, "s": 19}, {"kind": "d2", "x": 147, "r": 18, "s": 19}, {"kind": "d2", "x": 156, "r": 19, "s": 16}, {"kind": "d2", "x": 439, "r": 19, "s": 16}, {"kind": "d2", "x": 291, "r": 19, "s": 16}, {"kind": "d2", "x": 397, "r": 19, "s": 16}, {"kind": "d2", "x": 254, "r": 19, "s": 16}, {"kind": "d2", "x": 317, "r": 19, "s": 16}, {"kind": "d2", "x": 292, "r": 19, "s": 16}, {"kind": "d2", "x": 334, "r": 19, "s": 16}, {"kind": "d2", "x": 172, "r": 19, "s": 16}, {"kind": "d2", "x": 130, "r": 19, "s": 16}, {"kind": "d2", "x": 36, "r": 19, "s": 16}, {"kind": "d2", "x": 465, "r": 19, "s": 16}, {"kind": "d2", "x": 182, "r": 19, "s": 16}, {"kind": "d2", "x": 430, "r": 19, "s": 16}, {"kind": "d2", "x": 171, "r": 19, "s": 16}, {"kind": "d2", "x": 395, "r": 19, "s": 16}, {"kind": "d2", "x": 149, "r": 19, "s": 16}, {"kind": "d2", "x": 262, "r": 19, "s": 16}, {"kind": "d2", "x": 205, "r": 19, "s": 16}, {"kind": "d2", "x": 427, "r": 19, "s": 16}, {"kind": "d2", "x": 369, "r": 19, "s": 16}, {"kind": "d2", "x": 398, "r": 19, "s": 16}, {"kind": "d2", "x": 312, "r": 19, "s": 16}, {"kind": "d2", "x": 288, "r": 19, "s": 16}, {"kind": "d2", "x": 303, "r": 19, "s": 16}, {"kind": "d2", "x": 304, "r": 19, "s": 16}, {"kind": "d2", "x": 355, "r": 19, "s": 16}, {"kind": "d2", "x": 390, "r": 19, "s": 16}, {"kind": "d2", "x": 46, "r": 19, "s": 16}, {"kind": "d2", "x": 487, "r": 19, "s": 16}, {"kind": "d2", "x": 349, "r": 19, "s": 16}, {"kind": "d2", "x": 38, "r": 19, "s": 16}, {"kind": "d2", "x": 198, "r": 19, "s": 16}, {"kind": "d2", "x": 434, "r": 19, "s": 16}, {"kind": "d2", "x": 327, "r": 19, "s": 16}, {"kind": "d2", "x": 301, "r": 19, "s": 16}, {"kind": "d2", "x": 14, "r": 19, "s": 16}, {"kind": "d2", "x": 446, "r": 19, "s": 16}, {"kind": "d2", "x": 447, "r": 19, "s": 16}, {"kind": "d2", "x": 431, "r": 19, "s": 16}, {"kind": "d2", "x": 201, "r": 19, "s": 16}, {"kind": "d2", "x": 56, "r": 19, "s": 16}, {"kind": "d2", "x": 25, "r": 19, "s": 16}, {"kind": "d2", "x": 231, "r": 19, "s": 16}, {"kind": "d2", "x": 376, "r": 19, "s": 16}, {"kind": "d2", "x": 350, "r": 19, "s": 16}, {"kind": "d2", "x": 455, "r": 19, "s": 16}, {"kind": "d2", "x": 229, "r": 19, "s": 16}, {"kind": "d2", "x": 396, "r": 19, "s": 16}, {"kind": "d2", "x": 372, "r": 19, "s": 16}]
for Minegishi/reproducibility/data/chain_loop_k8/val_w2.json ADDED
@@ -0,0 +1 @@
 
 
1
+ [{"kind": "w2", "q1": [71, 15], "q2": [277, 4]}, {"kind": "w2", "q1": [161, 1], "q2": [405, 12]}, {"kind": "w2", "q1": [328, 13], "q2": [383, 11]}, {"kind": "w2", "q1": [165, 13], "q2": [191, 1]}, {"kind": "w2", "q1": [23, 15], "q2": [450, 2]}, {"kind": "w2", "q1": [128, 19], "q2": [145, 14]}, {"kind": "w2", "q1": [449, 4], "q2": [6, 10]}, {"kind": "w2", "q1": [255, 10], "q2": [139, 15]}, {"kind": "w2", "q1": [141, 4], "q2": [354, 1]}, {"kind": "w2", "q1": [194, 5], "q2": [246, 5]}, {"kind": "w2", "q1": [90, 18], "q2": [318, 6]}, {"kind": "w2", "q1": [348, 1], "q2": [71, 19]}, {"kind": "w2", "q1": [492, 6], "q2": [420, 3]}, {"kind": "w2", "q1": [236, 13], "q2": [252, 10]}, {"kind": "w2", "q1": [475, 14], "q2": [426, 5]}, {"kind": "w2", "q1": [327, 11], "q2": [170, 7]}, {"kind": "w2", "q1": [281, 19], "q2": [397, 9]}, {"kind": "w2", "q1": [411, 9], "q2": [139, 13]}, {"kind": "w2", "q1": [55, 11], "q2": [64, 4]}, {"kind": "w2", "q1": [457, 3], "q2": [69, 0]}]
for Minegishi/reproducibility/data/chain_loop_k8/vocab.json ADDED
@@ -0,0 +1 @@
 
 
1
+ ["<pad>", "<e_0>", "<e_1>", "<e_2>", "<e_3>", "<e_4>", "<e_5>", "<e_6>", "<e_7>", "<e_8>", "<e_9>", "<e_10>", "<e_11>", "<e_12>", "<e_13>", "<e_14>", "<e_15>", "<e_16>", "<e_17>", "<e_18>", "<e_19>", "<e_20>", "<e_21>", "<e_22>", "<e_23>", "<e_24>", "<e_25>", "<e_26>", "<e_27>", "<e_28>", "<e_29>", "<e_30>", "<e_31>", "<e_32>", "<e_33>", "<e_34>", "<e_35>", "<e_36>", "<e_37>", "<e_38>", "<e_39>", "<e_40>", "<e_41>", "<e_42>", "<e_43>", "<e_44>", "<e_45>", "<e_46>", "<e_47>", "<e_48>", "<e_49>", "<e_50>", "<e_51>", "<e_52>", "<e_53>", "<e_54>", "<e_55>", "<e_56>", "<e_57>", "<e_58>", "<e_59>", "<e_60>", "<e_61>", "<e_62>", "<e_63>", "<e_64>", "<e_65>", "<e_66>", "<e_67>", "<e_68>", "<e_69>", "<e_70>", "<e_71>", "<e_72>", "<e_73>", "<e_74>", "<e_75>", "<e_76>", "<e_77>", "<e_78>", "<e_79>", "<e_80>", "<e_81>", "<e_82>", "<e_83>", "<e_84>", "<e_85>", "<e_86>", "<e_87>", "<e_88>", "<e_89>", "<e_90>", "<e_91>", "<e_92>", "<e_93>", "<e_94>", "<e_95>", "<e_96>", "<e_97>", "<e_98>", "<e_99>", "<e_100>", "<e_101>", "<e_102>", "<e_103>", "<e_104>", "<e_105>", "<e_106>", "<e_107>", "<e_108>", "<e_109>", "<e_110>", "<e_111>", "<e_112>", "<e_113>", "<e_114>", "<e_115>", "<e_116>", "<e_117>", "<e_118>", "<e_119>", "<e_120>", "<e_121>", "<e_122>", "<e_123>", "<e_124>", "<e_125>", "<e_126>", "<e_127>", "<e_128>", "<e_129>", "<e_130>", "<e_131>", "<e_132>", "<e_133>", "<e_134>", "<e_135>", "<e_136>", "<e_137>", "<e_138>", "<e_139>", "<e_140>", "<e_141>", "<e_142>", "<e_143>", "<e_144>", "<e_145>", "<e_146>", "<e_147>", "<e_148>", "<e_149>", "<e_150>", "<e_151>", "<e_152>", "<e_153>", "<e_154>", "<e_155>", "<e_156>", "<e_157>", "<e_158>", "<e_159>", "<e_160>", "<e_161>", "<e_162>", "<e_163>", "<e_164>", "<e_165>", "<e_166>", "<e_167>", "<e_168>", "<e_169>", "<e_170>", "<e_171>", "<e_172>", "<e_173>", "<e_174>", "<e_175>", "<e_176>", "<e_177>", "<e_178>", "<e_179>", "<e_180>", "<e_181>", "<e_182>", "<e_183>", "<e_184>", "<e_185>", "<e_186>", "<e_187>", "<e_188>", "<e_189>", "<e_190>", "<e_191>", "<e_192>", "<e_193>", "<e_194>", "<e_195>", "<e_196>", "<e_197>", "<e_198>", "<e_199>", "<e_200>", "<e_201>", "<e_202>", "<e_203>", "<e_204>", "<e_205>", "<e_206>", "<e_207>", "<e_208>", "<e_209>", "<e_210>", "<e_211>", "<e_212>", "<e_213>", "<e_214>", "<e_215>", "<e_216>", "<e_217>", "<e_218>", "<e_219>", "<e_220>", "<e_221>", "<e_222>", "<e_223>", "<e_224>", "<e_225>", "<e_226>", "<e_227>", "<e_228>", "<e_229>", "<e_230>", "<e_231>", "<e_232>", "<e_233>", "<e_234>", "<e_235>", "<e_236>", "<e_237>", "<e_238>", "<e_239>", "<e_240>", "<e_241>", "<e_242>", "<e_243>", "<e_244>", "<e_245>", "<e_246>", "<e_247>", "<e_248>", "<e_249>", "<e_250>", "<e_251>", "<e_252>", "<e_253>", "<e_254>", "<e_255>", "<e_256>", "<e_257>", "<e_258>", "<e_259>", "<e_260>", "<e_261>", "<e_262>", "<e_263>", "<e_264>", "<e_265>", "<e_266>", "<e_267>", "<e_268>", "<e_269>", "<e_270>", "<e_271>", "<e_272>", "<e_273>", "<e_274>", "<e_275>", "<e_276>", "<e_277>", "<e_278>", "<e_279>", "<e_280>", "<e_281>", "<e_282>", "<e_283>", "<e_284>", "<e_285>", "<e_286>", "<e_287>", "<e_288>", "<e_289>", "<e_290>", "<e_291>", "<e_292>", "<e_293>", "<e_294>", "<e_295>", "<e_296>", "<e_297>", "<e_298>", "<e_299>", "<e_300>", "<e_301>", "<e_302>", "<e_303>", "<e_304>", "<e_305>", "<e_306>", "<e_307>", "<e_308>", "<e_309>", "<e_310>", "<e_311>", "<e_312>", "<e_313>", "<e_314>", "<e_315>", "<e_316>", "<e_317>", "<e_318>", "<e_319>", "<e_320>", "<e_321>", "<e_322>", "<e_323>", "<e_324>", "<e_325>", "<e_326>", "<e_327>", "<e_328>", "<e_329>", "<e_330>", "<e_331>", "<e_332>", "<e_333>", "<e_334>", "<e_335>", "<e_336>", "<e_337>", "<e_338>", "<e_339>", "<e_340>", "<e_341>", "<e_342>", "<e_343>", "<e_344>", "<e_345>", "<e_346>", "<e_347>", "<e_348>", "<e_349>", "<e_350>", "<e_351>", "<e_352>", "<e_353>", "<e_354>", "<e_355>", "<e_356>", "<e_357>", "<e_358>", "<e_359>", "<e_360>", "<e_361>", "<e_362>", "<e_363>", "<e_364>", "<e_365>", "<e_366>", "<e_367>", "<e_368>", "<e_369>", "<e_370>", "<e_371>", "<e_372>", "<e_373>", "<e_374>", "<e_375>", "<e_376>", "<e_377>", "<e_378>", "<e_379>", "<e_380>", "<e_381>", "<e_382>", "<e_383>", "<e_384>", "<e_385>", "<e_386>", "<e_387>", "<e_388>", "<e_389>", "<e_390>", "<e_391>", "<e_392>", "<e_393>", "<e_394>", "<e_395>", "<e_396>", "<e_397>", "<e_398>", "<e_399>", "<e_400>", "<e_401>", "<e_402>", "<e_403>", "<e_404>", "<e_405>", "<e_406>", "<e_407>", "<e_408>", "<e_409>", "<e_410>", "<e_411>", "<e_412>", "<e_413>", "<e_414>", "<e_415>", "<e_416>", "<e_417>", "<e_418>", "<e_419>", "<e_420>", "<e_421>", "<e_422>", "<e_423>", "<e_424>", "<e_425>", "<e_426>", "<e_427>", "<e_428>", "<e_429>", "<e_430>", "<e_431>", "<e_432>", "<e_433>", "<e_434>", "<e_435>", "<e_436>", "<e_437>", "<e_438>", "<e_439>", "<e_440>", "<e_441>", "<e_442>", "<e_443>", "<e_444>", "<e_445>", "<e_446>", "<e_447>", "<e_448>", "<e_449>", "<e_450>", "<e_451>", "<e_452>", "<e_453>", "<e_454>", "<e_455>", "<e_456>", "<e_457>", "<e_458>", "<e_459>", "<e_460>", "<e_461>", "<e_462>", "<e_463>", "<e_464>", "<e_465>", "<e_466>", "<e_467>", "<e_468>", "<e_469>", "<e_470>", "<e_471>", "<e_472>", "<e_473>", "<e_474>", "<e_475>", "<e_476>", "<e_477>", "<e_478>", "<e_479>", "<e_480>", "<e_481>", "<e_482>", "<e_483>", "<e_484>", "<e_485>", "<e_486>", "<e_487>", "<e_488>", "<e_489>", "<e_490>", "<e_491>", "<e_492>", "<e_493>", "<e_494>", "<e_495>", "<e_496>", "<e_497>", "<e_498>", "<e_499>", "<r_0>", "<r_1>", "<r_2>", "<r_3>", "<r_4>", "<r_5>", "<r_6>", "<r_7>", "<r_8>", "<r_9>", "<r_10>", "<r_11>", "<r_12>", "<r_13>", "<r_14>", "<r_15>", "<r_16>", "<r_17>", "<r_18>", "<r_19>", "Q", "ANS", "END"]
for Minegishi/reproducibility/eval_checkpoints.py ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Publish a named, immutable model snapshot before each training-time evaluation."""
2
+ import hashlib
3
+ import json
4
+ import os
5
+ from pathlib import Path
6
+ import torch
7
+
8
+ def save_eval_checkpoint(model, update, save_dir, rope_base, max_pos):
9
+ path=Path(save_dir)/f'ckpt_{update:07d}.pt'
10
+ payload=dict(model=model.state_dict(),update=int(update),rope_base=rope_base,max_pos=max_pos)
11
+ if path.exists():
12
+ old=torch.load(path,map_location='cpu',weights_only=False)
13
+ if old.get('update')!=update or old.get('rope_base')!=rope_base or old.get('max_pos')!=max_pos:
14
+ raise RuntimeError('Existing evaluation checkpoint has different metadata')
15
+ if old['model'].keys()!=payload['model'].keys() or any(not torch.equal(old['model'][k],v.detach().cpu()) for k,v in payload['model'].items()):
16
+ raise RuntimeError('Refusing to overwrite a different evaluation checkpoint')
17
+ else:
18
+ temporary=path.with_suffix('.pt.tmp')
19
+ with temporary.open('xb') as stream:
20
+ torch.save(payload,stream);stream.flush();os.fsync(stream.fileno())
21
+ os.replace(temporary,path)
22
+ digest=hashlib.sha256()
23
+ with path.open('rb') as stream:
24
+ for block in iter(lambda:stream.read(1<<20),b''):digest.update(block)
25
+ return dict(file=path.name,update=int(update),sha256=digest.hexdigest(),bytes=path.stat().st_size,
26
+ rope_base=rope_base,max_pos=max_pos,kind='model_weights_for_exact_evaluation')
for Minegishi/reproducibility/extrapolation.py ADDED
@@ -0,0 +1,289 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Evaluation-only depth and input-position extrapolation for final checkpoints.
2
+
3
+ The generated labels are used only to score predictions. No training, SNAP,
4
+ teacher state reset, or gold intermediate-state injection is performed.
5
+ """
6
+ import argparse
7
+ import collections
8
+ import hashlib
9
+ import json
10
+ from pathlib import Path
11
+ import random
12
+ import time
13
+
14
+ import numpy as np
15
+ import torch
16
+
17
+ from probes import SmallEvalSet, load_model, named_seed, read_json, save_json
18
+ from train_halt import eval_fixed
19
+
20
+
21
+ CATEGORIES = ("all_seen", "one_new", "random")
22
+
23
+
24
+ def relation_graph(meta):
25
+ nr = int(meta["num_relations"])
26
+ seen = {tuple(map(int, edge)) for edge in meta["train_edges"]}
27
+ if any(not (0 <= a < nr and 0 <= b < nr) for a, b in seen):
28
+ raise ValueError("Invalid relation ID in train_edges")
29
+ successors = [sorted(b for a, b in seen if a == r) for r in range(nr)]
30
+ if any(not row for row in successors):
31
+ raise ValueError("Every relation must have a seen outgoing edge")
32
+ # Preserve the original test split when available: its unseen edges exclude
33
+ # val_pairs. The fallback is the exact complement of train_edges.
34
+ pool = {tuple(map(int, edge)) for edge in meta.get(
35
+ "test_pool", [[a, b] for a in range(nr) for b in range(nr) if (a, b) not in seen])}
36
+ if pool & seen or any(not (0 <= a < nr and 0 <= b < nr) for a, b in pool):
37
+ raise ValueError("test_pool must contain valid, unseen relation pairs")
38
+ unseen_successors = [sorted(b for a, b in pool if a == r) for r in range(nr)]
39
+ if any(not row for row in unseen_successors):
40
+ raise ValueError("Every relation needs an unseen outgoing test edge for the fixed-jump sampler")
41
+ return seen, successors, unseen_successors
42
+
43
+
44
+ def make_chain_items(meta, depth, n, seed, category, namespace):
45
+ """Fixed-seed graph walks, exactly-one-new walks, or iid uniform relations."""
46
+ if category not in CATEGORIES or depth < 2:
47
+ raise ValueError("Unsupported chain category or depth")
48
+ seen, successors, unseen_successors = relation_graph(meta)
49
+ ne, nr = int(meta["num_entities"]), int(meta["num_relations"])
50
+ mapping = meta["rel_map"]
51
+ if len(mapping) != nr or any(len(row) != ne for row in mapping):
52
+ raise ValueError("rel_map must have shape (num_relations, num_entities)")
53
+ rng = random.Random(named_seed(seed, f"{namespace}|d{depth}|{category}"))
54
+ items = []
55
+ for index in range(n):
56
+ x = rng.randrange(ne)
57
+ jump = rng.randrange(1, depth) if category == "one_new" else None
58
+ relations = [rng.randrange(nr)]
59
+ for position in range(1, depth):
60
+ previous = relations[-1]
61
+ if category == "random":
62
+ relation = rng.randrange(nr)
63
+ elif position == jump:
64
+ relation = rng.choice(unseen_successors[previous])
65
+ else:
66
+ relation = rng.choice(successors[previous])
67
+ relations.append(relation)
68
+ unseen_positions = [i for i in range(1, depth)
69
+ if (relations[i - 1], relations[i]) not in seen]
70
+ if category == "all_seen" and unseen_positions:
71
+ raise AssertionError("all_seen contains an unseen adjacency")
72
+ if category == "one_new" and unseen_positions != [jump]:
73
+ raise AssertionError("one_new must have exactly one unseen adjacency occurrence")
74
+ gold = x
75
+ for relation in relations:
76
+ gold = int(mapping[relation][gold])
77
+ content = json.dumps([x, relations], separators=(",", ":"))
78
+ fingerprint = hashlib.sha256(content.encode()).hexdigest()[:12]
79
+ items.append(dict(id=f"{namespace}:d{depth}:{category}:{index:06d}:{fingerprint}",
80
+ cat=category, d=depth, x=x, rels=relations, gold=gold,
81
+ unseen_count=len(unseen_positions), unseen_positions=unseen_positions))
82
+ if len(items) != n or any(item["cat"] != category for item in items):
83
+ raise AssertionError("Category/sample count mismatch")
84
+ return items
85
+
86
+
87
+ def dataset_audit(items):
88
+ content = json.dumps([[item["id"], item["x"], item["rels"], item["gold"]]
89
+ for item in items], separators=(",", ":"))
90
+ return dict(n=len(items), category_counts=dict(collections.Counter(item["cat"] for item in items)),
91
+ unseen_adjacency_count_histogram=dict(sorted(collections.Counter(
92
+ item["unseen_count"] for item in items).items())),
93
+ content_sha256=hashlib.sha256(content.encode()).hexdigest(),
94
+ duplicate_input_count=len(items) - len({(item["x"], tuple(item["rels"])) for item in items}))
95
+
96
+
97
+ def scored_result(items, scores):
98
+ arrays = {str(budget): np.asarray(correct, dtype=bool) for budget, correct in scores.items()}
99
+ if any(len(correct) != len(items) for correct in arrays.values()):
100
+ raise AssertionError("Prediction/sample count mismatch")
101
+ return dict(audit=dataset_audit(items), ids=[item["id"] for item in items],
102
+ gold_entity_ids=[item["gold"] for item in items],
103
+ accuracy={budget: float(correct.mean()) for budget, correct in arrays.items()},
104
+ scores={budget: correct.tolist() for budget, correct in arrays.items()},
105
+ score_definition="Per-example Boolean correctness, in the same order as ids; not confidence scores.")
106
+
107
+
108
+ class PrefixEvalSet:
109
+ """Depth means semantic depth=2; token length and read position include prefix.
110
+
111
+ Do not pass this class to train_halt.eval_fixed, which reads at position d.
112
+ """
113
+ def __init__(self, items, prefix_entities, e0, r0, device, batch_size):
114
+ if any(item["d"] != 2 or len(item["rels"]) != 2 for item in items):
115
+ raise ValueError("Prefix evaluation requires d2 queries")
116
+ prefix_entities = np.asarray(prefix_entities, dtype=np.int64)
117
+ if prefix_entities.ndim != 2 or len(prefix_entities) != len(items):
118
+ raise ValueError("One prefix row is required per query")
119
+ self.items, self.n, self.batch_size = items, len(items), batch_size
120
+ self.offset = int(prefix_entities.shape[1])
121
+ content = np.asarray([[e0 + item["x"]] + [r0 + r for r in item["rels"]]
122
+ for item in items], dtype=np.int64)
123
+ tokens = np.concatenate([e0 + prefix_entities, content], axis=1)
124
+ self.tokens = torch.tensor(tokens, dtype=torch.long, device=device)
125
+ self.gold = torch.tensor([e0 + item["gold"] for item in items], device=device)
126
+ self.last_position = self.offset + 2
127
+ if self.tokens.shape[1] != self.last_position + 1:
128
+ raise AssertionError("Wrong last read position for prefixed queries")
129
+
130
+ def batches(self):
131
+ for start in range(0, self.n, self.batch_size):
132
+ stop = min(start + self.batch_size, self.n)
133
+ yield 2, np.arange(start, stop), self.tokens[start:stop], self.gold[start:stop]
134
+
135
+
136
+ @torch.inference_mode()
137
+ def eval_prefix(model, evaluation, budgets=(2, 4)):
138
+ result = {budget: np.zeros(evaluation.n, dtype=bool) for budget in budgets}
139
+ for semantic_depth, indices, tokens, gold in evaluation.batches():
140
+ if semantic_depth != 2:
141
+ raise AssertionError("Prefix length must not become semantic depth")
142
+ last = torch.full((len(indices),), tokens.shape[1] - 1, dtype=torch.long, device=tokens.device)
143
+ if tokens.shape[1] - 1 != evaluation.last_position:
144
+ raise AssertionError("Prefix readout position mismatch")
145
+ state = model.embed_state(tokens)
146
+ for step in range(1, max(budgets) + 1):
147
+ hidden, state = model.step_state(state)
148
+ if step in result:
149
+ result[step][indices] = (model.logits(model.read(hidden, last)).argmax(-1) == gold).cpu().numpy()
150
+ return result
151
+
152
+
153
+ def paired_comparison(baseline_scores, current_scores):
154
+ result = {}
155
+ for budget, baseline in baseline_scores.items():
156
+ baseline = np.asarray(baseline, dtype=bool)
157
+ current = np.asarray(current_scores[budget], dtype=bool)
158
+ if current.shape != baseline.shape:
159
+ raise AssertionError("Paired content changed across prefix offsets")
160
+ denominator = int(baseline.sum())
161
+ result[str(budget)] = dict(
162
+ accuracy_difference_vs_offset0=float(current.mean() - baseline.mean()),
163
+ both_correct_count=int((current & baseline).sum()),
164
+ lost_correct_count=int((baseline & ~current).sum()),
165
+ gained_correct_count=int((~baseline & current).sum()),
166
+ retention_denominator_correct_offset0=denominator,
167
+ retention_given_correct_offset0=float(current[baseline].mean()) if denominator else None)
168
+ return result
169
+
170
+
171
+ def parse_args():
172
+ parser = argparse.ArgumentParser(description=__doc__)
173
+ parser.add_argument("--run-dir", required=True)
174
+ parser.add_argument("--data-dir", required=True)
175
+ parser.add_argument("--out", required=True)
176
+ parser.add_argument("--device", default="cuda")
177
+ parser.add_argument("--n", type=int, default=1000, help="Per deep category and per prefix offset")
178
+ parser.add_argument("--seed", type=int, default=20260923)
179
+ parser.add_argument("--batch-size", type=int, default=64)
180
+ parser.add_argument("--deep-depth", type=int, default=256)
181
+ parser.add_argument("--prefix-offsets", type=int, nargs="+", default=[0, 32, 128, 256])
182
+ parser.add_argument("--prefix-category", choices=CATEGORIES, default="all_seen")
183
+ parser.add_argument("--skip-deep", action="store_true")
184
+ parser.add_argument("--skip-prefix", action="store_true")
185
+ args = parser.parse_args()
186
+ if args.n < 1 or args.batch_size < 1 or args.deep_depth < 2:
187
+ parser.error("n/batch-size must be positive; deep-depth must be >=2")
188
+ if min(args.prefix_offsets) < 0 or 0 not in args.prefix_offsets:
189
+ parser.error("prefix-offsets must be nonnegative and include paired baseline 0")
190
+ if args.skip_deep and args.skip_prefix:
191
+ parser.error("At least one evaluation must be enabled")
192
+ args.prefix_offsets = sorted(set(args.prefix_offsets))
193
+ return args
194
+
195
+
196
+ def main():
197
+ args = parse_args()
198
+ start = time.monotonic()
199
+ run_dir, data_dir = Path(args.run_dir), Path(args.data_dir)
200
+ meta, vocab = read_json(data_dir / "meta.json"), read_json(data_dir / "vocab.json")
201
+ device = torch.device(args.device)
202
+ torch.backends.cuda.matmul.allow_tf32 = False
203
+ torch.backends.cudnn.allow_tf32 = False
204
+ model, manifest, update, e0, r0 = load_model(run_dir, vocab, meta, device)
205
+ lengths = ([args.deep_depth + 1] if not args.skip_deep else [])
206
+ lengths += ([max(args.prefix_offsets) + 3] if not args.skip_prefix else [])
207
+ if model.pos == "rope" and max(lengths) > model.rope_cos.shape[0]:
208
+ raise ValueError("Evaluation input exceeds checkpoint RoPE capacity")
209
+ result = dict(
210
+ extrapolation_evaluation=True, status="running", checkpoint="last.pt",
211
+ checkpoint_update=update, run_dir=str(run_dir.resolve()), config=vars(args),
212
+ training_seed=manifest.get("seed"), evaluation_seed=args.seed,
213
+ precision="strict_fp32_no_autocast_tf32_disabled", device=str(device),
214
+ torch_version=str(torch.__version__),
215
+ source_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
216
+ meta_sha256=hashlib.sha256((data_dir / "meta.json").read_bytes()).hexdigest(),
217
+ training_condition={key: manifest.get(key) for key in ("anchor", "consist_loss", "offset_max", "model")},
218
+ design_notes=[
219
+ "All labels are generated and used for evaluation only; no model parameters or training data are changed.",
220
+ "all_seen follows uniformly selected train-graph successors. one_new has exactly one unseen adjacency occurrence; that edge comes from meta.test_pool when available.",
221
+ "random uses iid uniform relation IDs without rejection based on novelty count; its actual unseen counts are reported.",
222
+ "Prefix tests use the same d2 query content at every offset. Independent prefix RNG generates one longest prefix; smaller prefixes use its suffix so the nearest distractors remain paired.",
223
+ "Prefixed d2 is read at offset+2 after exactly 2 or 4 loops. Prefix length never changes semantic depth or loop budget.",
224
+ "Long-chain scores use the ordinary model state API, no SNAP/reset/teacher intermediate states; all budgets share one rollout per batch.",
225
+ "New evaluation sets are fixed by this script and seed; they extend existing tests and do not replace the trainer's primary final evaluation."],
226
+ completed_stages=[])
227
+
228
+ def persist(active_stage):
229
+ result["active_stage"] = active_stage
230
+ result["elapsed_seconds"] = time.monotonic() - start
231
+ save_json(args.out, result)
232
+
233
+ persist("initializing")
234
+ # Cheap position tests run first so a late queue deadline cannot erase them.
235
+ if not args.skip_prefix:
236
+ items = make_chain_items(meta, 2, args.n, args.seed, args.prefix_category, "prefix-query")
237
+ prefix_rng = np.random.default_rng(named_seed(args.seed, "prefix-entities-independent"))
238
+ longest = prefix_rng.integers(0, int(meta["num_entities"]),
239
+ size=(args.n, max(args.prefix_offsets)), dtype=np.int64)
240
+ baseline_scores = None
241
+ result["prefix_d2"] = {}
242
+ for offset in args.prefix_offsets:
243
+ persist(f"prefix_offset_{offset}")
244
+ prefixes = longest[:, -offset:] if offset else longest[:, :0]
245
+ evaluation = PrefixEvalSet(items, prefixes, e0, r0, device, args.batch_size)
246
+ scores = eval_prefix(model, evaluation)
247
+ if offset == 0:
248
+ baseline_scores = scores
249
+ cell = scored_result(items, scores)
250
+ cell.update(offset=offset, semantic_depth=2, sequence_length=offset + 3,
251
+ last_read_position=offset + 2,
252
+ prefix_entity_ids_sha256=hashlib.sha256(prefixes.astype("<i8").tobytes()).hexdigest(),
253
+ paired_vs_offset0=paired_comparison(baseline_scores, scores))
254
+ result["prefix_d2"][str(offset)] = cell
255
+ result["completed_stages"].append(f"prefix_offset_{offset}")
256
+ persist(f"prefix_offset_{offset}_complete")
257
+ print(f"prefix offset={offset}, n={args.n}: {cell['accuracy']}", flush=True)
258
+
259
+ if not args.skip_deep:
260
+ result["deep"] = {}
261
+ for category in CATEGORIES:
262
+ persist(f"deep_d{args.deep_depth}_{category}")
263
+ items = make_chain_items(meta, args.deep_depth, args.n, args.seed, category, "depth-extrapolation")
264
+ evaluation = SmallEvalSet(items, e0, r0, device, args.batch_size)
265
+ scores = eval_fixed(model, evaluation, ["d", "d+2", "2d"])
266
+ cell = scored_result(items, scores)
267
+ cell["depth"] = args.deep_depth
268
+ cell["loop_budgets"] = {"d": args.deep_depth, "d+2": args.deep_depth + 2, "2d": 2 * args.deep_depth}
269
+ at_d = np.asarray(scores["d"], dtype=bool)
270
+ denominator = int(at_d.sum())
271
+ cell["retention_denominator_correct_R_d"] = denominator
272
+ cell["retention_given_correct_R_d"] = {
273
+ budget: float(np.asarray(scores[budget])[at_d].mean()) if denominator else None
274
+ for budget in ("d+2", "2d")}
275
+ result["deep"][category] = cell
276
+ result["completed_stages"].append(f"deep_d{args.deep_depth}_{category}")
277
+ persist(f"deep_d{args.deep_depth}_{category}_complete")
278
+ print(f"deep d={args.deep_depth} {category}, n={args.n}: {cell['accuracy']}", flush=True)
279
+ if set(result["deep"]) != set(CATEGORIES) or any(
280
+ cell["audit"]["n"] != args.n for cell in result["deep"].values()):
281
+ raise AssertionError("Final deep evaluation category counts are incomplete")
282
+
283
+ result["status"] = "complete"
284
+ persist("complete")
285
+ print(f"extrapolation complete in {result['elapsed_seconds']:.1f}s: {args.out}", flush=True)
286
+
287
+
288
+ if __name__ == "__main__":
289
+ main()
for Minegishi/reproducibility/frozen_source_sha256.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model/loop_gpt.py": "2245bd84895cb9cb587bf381e502d37cb52db6bd2ff382d892a4b31ffd20078b",
3
+ "model/__init__.py": "2f2b8681adf78b17209129795e950a954e3e9b030e97bdcc9446aa5aa9999247",
4
+ "model/gpt2.py": "f97c63c37d88d34698ac239507033334679d70023671f8542bcb6508b364537f",
5
+ "model/latent_executor.py": "26d33ae9d6ddb49b2c7e467ff9ede4aa1c43517336c53a2f58377e155af006e6",
6
+ "train_halt.py": "813e8f785c04341a85be4a7a58d1c6663ae3a3d4948f9ad9d063a8a71ef4eb0b",
7
+ "objectives.py": "f20f2507e995da5bccffa411edbfc8322b484f2f2d509d05ce374da282eab798",
8
+ "streams.py": "4bd4a83b21172b32079041117504020a5f8921da54340da28398b42539688c75",
9
+ "train_chain.py": "a9fc664c3dd03c0cd11e0e32a0deca0ecc60ef1b9a37a16d23d696f8a6f380b9",
10
+ "recipe.py": "3aae5bfe4189ec4323a59a935a8bb32c1a85a44e2c2d7f6cdd46915a5051cb41",
11
+ "eval_checkpoints.py": "6871b82c230566b824c0258cd40a0b5b8e6f72ae2578209ab5178a64fbc7871f",
12
+ "probes.py": "d32d3a7f5fa232c3cd7b23c28681e6cb10c869bdc45e303c95f7496522f8b10e",
13
+ "extrapolation.py": "7aa2779121f317384cc9e926ba36740033448f001e9d24853f7c4ee074da1962"
14
+ }
for Minegishi/reproducibility/load_checkpoint.py ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load the supplied custom LoopGPT evaluation snapshots with explicit RoPE."""
2
+ import json
3
+ from pathlib import Path
4
+ import torch
5
+ from model.loop_gpt import LoopGPT
6
+
7
+
8
+ def load_checkpoint(path, device='cpu'):
9
+ path = Path(path)
10
+ checkpoint = torch.load(path, map_location='cpu', weights_only=True)
11
+ if checkpoint.get('rope_base') != 10000 or checkpoint.get('max_pos') != 4097:
12
+ raise ValueError('Expected the original RoPE=10000, max_pos=4097 snapshots')
13
+ vocab = json.loads((Path(__file__).parent/'data/chain_loop_k8/vocab.json').read_text())
14
+ model = LoopGPT(len(vocab), d=768, n_head=12, n_layer=4,
15
+ pos='rope', rope_base=checkpoint['rope_base'], max_pos=checkpoint['max_pos'],
16
+ attn_window=0, reembed=0, ent_range=(1,501), rel_range=(501,521))
17
+ model.load_state_dict(checkpoint['model'], strict=True)
18
+ model = model.to(device).float().eval()
19
+ torch.backends.cuda.matmul.allow_tf32 = False
20
+ torch.backends.cudnn.allow_tf32 = False
21
+ metadata = {k:v for k,v in checkpoint.items() if k != 'model'}
22
+ return model, vocab, metadata
23
+
24
+
25
+ @torch.inference_mode()
26
+ def predict(model, tokens, loops):
27
+ if tokens.ndim != 2 or tokens.dtype != torch.long or loops < 1:
28
+ raise ValueError('Expected a batch of integer token sequences and loops >= 1')
29
+ if tokens.shape[1] > model.rope_cos.shape[0]:
30
+ raise ValueError('Input exceeds the saved RoPE cache capacity')
31
+ hidden = model.embed(tokens)
32
+ for _ in range(loops):
33
+ hidden = model.step(hidden)
34
+ last = torch.full((tokens.shape[0],), tokens.shape[1]-1, dtype=torch.long, device=tokens.device)
35
+ return model.logits(model.read(hidden, last)).argmax(-1)
for Minegishi/reproducibility/model/__init__.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ from .gpt2 import GPT2LikeEncoder
2
+ from .latent_executor import LatentExecutor
3
+ from .loop_gpt import LoopGPT
4
+
5
+ __all__ = ["GPT2LikeEncoder", "LatentExecutor", "LoopGPT"]
for Minegishi/reproducibility/model/gpt2.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ GPT-2-like Transformer model with Rotary Position Embedding (RoPE).
3
+ """
4
+
5
+ import math
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+
11
+
12
+ def rotate_half(x: torch.Tensor) -> torch.Tensor:
13
+ """Rotate half the hidden dims of the input."""
14
+ d = x.size(-1)
15
+ x1 = x[..., : d // 2]
16
+ x2 = x[..., d // 2:]
17
+ return torch.cat([-x2, x1], dim=-1)
18
+
19
+
20
+ def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
21
+ """Apply rotary position embedding to input tensor."""
22
+ return (x * cos) + (rotate_half(x) * sin)
23
+
24
+
25
+ class RotaryEmbedding(nn.Module):
26
+ """
27
+ Rotary Position Embedding (RoPE).
28
+ Precomputes cos/sin values for efficient position encoding.
29
+ """
30
+
31
+ def __init__(self, head_dim: int, max_len: int = 1024, base: float = 10000.0):
32
+ super().__init__()
33
+ assert head_dim % 2 == 0, "RoPE requires even head_dim"
34
+ self.head_dim = head_dim
35
+ self.max_len = max_len
36
+ self.base = base
37
+
38
+ inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim))
39
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
40
+ self._build_cache(max_len)
41
+
42
+ def _build_cache(self, max_len: int):
43
+ """Build cos/sin cache up to max_len."""
44
+ t = torch.arange(max_len, dtype=torch.float32)
45
+ freqs = torch.einsum("l,d->ld", t, self.inv_freq)
46
+ emb = torch.cat([freqs, freqs], dim=-1)
47
+ cos = emb.cos()[None, None, :, :]
48
+ sin = emb.sin()[None, None, :, :]
49
+ self.register_buffer("cos_cached", cos, persistent=False)
50
+ self.register_buffer("sin_cached", sin, persistent=False)
51
+ self.max_len = max_len
52
+
53
+ def forward(self, seq_len: int, device=None, dtype=None):
54
+ """Get cos/sin embeddings for given sequence length."""
55
+ if seq_len > self.max_len:
56
+ self._build_cache(seq_len)
57
+ cos = self.cos_cached[:, :, :seq_len, :]
58
+ sin = self.sin_cached[:, :, :seq_len, :]
59
+ if device is not None:
60
+ cos = cos.to(device)
61
+ sin = sin.to(device)
62
+ if dtype is not None:
63
+ cos = cos.to(dtype=dtype)
64
+ sin = sin.to(dtype=dtype)
65
+ return cos, sin
66
+
67
+
68
+ class CausalSelfAttentionRoPE(nn.Module):
69
+ """Causal Self-Attention with Rotary Position Embedding."""
70
+
71
+ def __init__(
72
+ self,
73
+ d_model: int,
74
+ n_head: int,
75
+ dropout: float,
76
+ max_len: int,
77
+ rope_base: float = 10000.0,
78
+ use_rope: bool = True,
79
+ ):
80
+ super().__init__()
81
+ assert d_model % n_head == 0, "d_model must be divisible by n_head"
82
+ self.use_rope = use_rope
83
+ self.d_model = d_model
84
+ self.n_head = n_head
85
+ self.head_dim = d_model // n_head
86
+ assert self.head_dim % 2 == 0, "head_dim must be even for RoPE"
87
+
88
+ self.qkv = nn.Linear(d_model, 3 * d_model)
89
+ self.out = nn.Linear(d_model, d_model)
90
+ self.attn_drop = nn.Dropout(dropout)
91
+ self.resid_drop = nn.Dropout(dropout)
92
+
93
+ self.rope = RotaryEmbedding(self.head_dim, max_len=max_len, base=rope_base)
94
+
95
+ causal = torch.triu(torch.ones(max_len, max_len), diagonal=1).bool()
96
+ self.register_buffer("causal_mask", causal[None, None, :, :], persistent=False)
97
+
98
+ def forward(self, x: torch.Tensor, pad_mask: torch.Tensor | None = None, pos_ids: torch.Tensor | None = None):
99
+ """
100
+ Args:
101
+ x: Input tensor of shape (B, L, C)
102
+ pad_mask: Padding mask of shape (B, L), True for PAD positions
103
+ pos_ids: optional explicit RoPE positions of shape (B, L) (e.g. reset to 0 at every call of the local
104
+ executor with history); the causal mask still follows the token order. Default: 0..L-1.
105
+ """
106
+ B, L, C = x.shape
107
+
108
+ qkv = self.qkv(x)
109
+ q, k, v = qkv.split(C, dim=-1)
110
+
111
+ q = q.view(B, L, self.n_head, self.head_dim).transpose(1, 2)
112
+ k = k.view(B, L, self.n_head, self.head_dim).transpose(1, 2)
113
+ v = v.view(B, L, self.n_head, self.head_dim).transpose(1, 2)
114
+
115
+ if self.use_rope:
116
+ if pos_ids is None:
117
+ cos, sin = self.rope(seq_len=L, device=x.device, dtype=x.dtype)
118
+ else:
119
+ need = int(pos_ids.max()) + 1
120
+ if need > self.rope.max_len:
121
+ self.rope._build_cache(need)
122
+ cos = self.rope.cos_cached[0, 0].to(x.device, x.dtype)[pos_ids][:, None] # (B, 1, L, hd)
123
+ sin = self.rope.sin_cached[0, 0].to(x.device, x.dtype)[pos_ids][:, None]
124
+ q = apply_rope(q, cos, sin)
125
+ k = apply_rope(k, cos, sin)
126
+
127
+ att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim)
128
+ att = att.masked_fill(self.causal_mask[:, :, :L, :L], float("-inf"))
129
+
130
+ if pad_mask is not None:
131
+ att = att.masked_fill(pad_mask[:, None, None, :], float("-inf"))
132
+
133
+ att = F.softmax(att, dim=-1)
134
+ att = self.attn_drop(att)
135
+
136
+ y = att @ v
137
+ y = y.transpose(1, 2).contiguous().view(B, L, C)
138
+ y = self.resid_drop(self.out(y))
139
+ return y
140
+
141
+
142
+ class GPT2BlockRoPE(nn.Module):
143
+ """GPT-2-like Transformer block with RoPE (Pre-LN variant)."""
144
+
145
+ def __init__(
146
+ self,
147
+ d_model: int,
148
+ n_head: int,
149
+ dropout: float,
150
+ max_len: int,
151
+ rope_base: float = 10000.0,
152
+ use_rope: bool = True,
153
+ ):
154
+ super().__init__()
155
+ self.ln_1 = nn.LayerNorm(d_model)
156
+ self.attn = CausalSelfAttentionRoPE(d_model, n_head, dropout, max_len, rope_base=rope_base, use_rope=use_rope)
157
+ self.ln_2 = nn.LayerNorm(d_model)
158
+
159
+ self.mlp = nn.Sequential(
160
+ nn.Linear(d_model, 4 * d_model),
161
+ nn.GELU(),
162
+ nn.Linear(4 * d_model, d_model),
163
+ nn.Dropout(dropout),
164
+ )
165
+
166
+ def forward(self, x: torch.Tensor, pad_mask: torch.Tensor | None = None, pos_ids: torch.Tensor | None = None):
167
+ x = x + self.attn(self.ln_1(x), pad_mask=pad_mask, pos_ids=pos_ids)
168
+ x = x + self.mlp(self.ln_2(x))
169
+ return x
170
+
171
+
172
+ class GPT2LikeEncoder(nn.Module):
173
+ """
174
+ GPT-2-like Transformer encoder with Rotary Position Embedding.
175
+ Uses pre-LN architecture and weight tying between embedding and output layers.
176
+ """
177
+
178
+ def __init__(
179
+ self,
180
+ vocab_size: int,
181
+ d_model: int = 768,
182
+ n_layer: int = 8,
183
+ n_head: int = 12,
184
+ dropout: float = 0.0,
185
+ max_len: int = 1024,
186
+ rope_base: float = 100.0,
187
+ use_rope: bool = True,
188
+ zero_init_out: bool = False,
189
+ ):
190
+ """
191
+ Args:
192
+ vocab_size: Size of the vocabulary
193
+ d_model: Model dimension
194
+ n_layer: Number of transformer layers
195
+ n_head: Number of attention heads
196
+ dropout: Dropout rate
197
+ max_len: Maximum sequence length
198
+ rope_base: Base for rotary position embedding
199
+ zero_init_out: zero-initialise the attention / MLP output projections of every block, so that the block
200
+ starts as the identity map (stabilises the recurrent-depth model that re-applies the
201
+ same blocks n_loops times; docs/experiments_loop_next.md §6)
202
+ """
203
+ super().__init__()
204
+ self.tok_emb = nn.Embedding(vocab_size, d_model)
205
+ self.drop = nn.Dropout(dropout)
206
+
207
+ self.blocks = nn.ModuleList([
208
+ GPT2BlockRoPE(d_model, n_head, dropout, max_len, rope_base=rope_base, use_rope=use_rope)
209
+ for _ in range(n_layer)
210
+ ])
211
+
212
+ self.ln_f = nn.LayerNorm(d_model)
213
+ self.head = nn.Linear(d_model, vocab_size, bias=False)
214
+ # self.head.weight = self.tok_emb.weight # Weight tying
215
+
216
+ self.max_len = max_len
217
+ self.zero_init_out = zero_init_out
218
+ self.apply(self._init)
219
+ if zero_init_out:
220
+ for blk in self.blocks:
221
+ nn.init.zeros_(blk.attn.out.weight)
222
+ nn.init.zeros_(blk.mlp[2].weight)
223
+
224
+ def _init(self, m):
225
+ """Initialize weights."""
226
+ if isinstance(m, (nn.Linear, nn.Embedding)):
227
+ nn.init.normal_(m.weight, mean=0.0, std=0.02)
228
+ if isinstance(m, nn.Linear) and m.bias is not None:
229
+ nn.init.zeros_(m.bias)
230
+ if isinstance(m, nn.LayerNorm):
231
+ nn.init.ones_(m.weight)
232
+ nn.init.zeros_(m.bias)
233
+
234
+ def forward(self, input_ids: torch.Tensor, pad_mask: torch.Tensor = None, pos_ids: torch.Tensor = None,
235
+ n_loops: int = 1):
236
+ """
237
+ Args:
238
+ input_ids: Input token IDs of shape (B, L)
239
+ pad_mask: Padding mask of shape (B, L), True for PAD positions
240
+ pos_ids: optional explicit RoPE positions (B, L); default 0..L-1
241
+ n_loops: apply the whole block stack this many times with shared weights (recurrent-depth model,
242
+ h^(t+1) = B_theta(h^(t)); positions are unchanged across loops, no loop embedding)
243
+
244
+ Returns:
245
+ Logits of shape (B, L, V)
246
+ """
247
+ B, L = input_ids.shape
248
+ if L > self.max_len:
249
+ raise ValueError(f"seq_len {L} exceeds max_len {self.max_len}")
250
+
251
+ x = self.drop(self.tok_emb(input_ids))
252
+
253
+ for _ in range(n_loops):
254
+ for blk in self.blocks:
255
+ x = blk(x, pad_mask=pad_mask, pos_ids=pos_ids)
256
+
257
+ x = self.ln_f(x)
258
+ return self.head(x)
for Minegishi/reproducibility/model/latent_executor.py ADDED
@@ -0,0 +1,213 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Fixed problem memory + autonomously updated latent workspace (docs/experiments_latent_batch2_20runs.md §4).
3
+
4
+ packet --Encoder (2 layers, bidirectional)--> M (computed ONCE per solve, never overwritten)
5
+ 4 learned latent slots Z0; Z_t = Core(Z_{t-1}, M) for t = 1..T (no loop index, no pointer, no gold state)
6
+ Z_T --AnswerDecoder (1 layer, causal, cross-attends to Z_T only)--> answer tokens
7
+
8
+ mode 'shared' the same 2-block core is applied T times (conditions A-D)
9
+ mode 'single' the core is applied once (E; same parameters as A)
10
+ mode 'untied' 8 cores with separate parameters, loop t uses core t, T <= 8 (F)
11
+
12
+ RoPE (base 100) is applied to Q / K after their projections, V is not rotated. Memory tokens carry their logical
13
+ position ids 0..261, latent slots 262..265, answer prefix 266.. . The supervised read head (conditions C / D) is the
14
+ memory cross-attention of the SECOND core block, latent slot 0, head 0: its actual soft attention probabilities are
15
+ returned for the loss, and its weighted values enter the residual stream as usual (no oracle read module, no mask).
16
+ The entity auxiliary head (LayerNorm + linear -> E classes on slot 0 after every loop) is instantiated in every
17
+ condition so parameters and init are identical; it only receives gradients in B / D.
18
+ """
19
+ import hashlib
20
+ import math
21
+ from typing import List, Optional, Tuple
22
+
23
+ import torch
24
+ import torch.nn as nn
25
+ import torch.nn.functional as F
26
+
27
+
28
+ def rotate_half(x):
29
+ d = x.size(-1); return torch.cat([-x[..., d // 2:], x[..., :d // 2]], -1)
30
+
31
+
32
+ class Rope(nn.Module):
33
+ def __init__(self, head_dim: int, base: float = 100.0, max_pos: int = 288):
34
+ super().__init__()
35
+ inv = 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim))
36
+ f = torch.einsum("l,d->ld", torch.arange(max_pos).float(), inv); emb = torch.cat([f, f], -1)
37
+ self.register_buffer("cos", emb.cos(), persistent=False); self.register_buffer("sin", emb.sin(), persistent=False)
38
+
39
+ def forward(self, pos: torch.Tensor):
40
+ """pos (B, L) or (L,) -> (cos, sin) broadcastable to (B, H, L, hd)"""
41
+ if pos.dim() == 1:
42
+ return self.cos[pos][None, None], self.sin[pos][None, None]
43
+ return self.cos[pos][:, None], self.sin[pos][:, None]
44
+
45
+
46
+ class Attn(nn.Module):
47
+ def __init__(self, d: int, n_head: int, manual: bool = False):
48
+ super().__init__()
49
+ self.h = n_head; self.hd = d // n_head; self.manual = manual
50
+ self.q = nn.Linear(d, d); self.k = nn.Linear(d, d); self.v = nn.Linear(d, d); self.o = nn.Linear(d, d)
51
+
52
+ def forward(self, x, y, rope_q, rope_k, key_pad=None, causal=False):
53
+ """x (B, Lq, C) queries, y (B, Lk, C) keys/values. Returns (out, probs or None); probs (B, H, Lq, Lk) only if manual."""
54
+ B, Lq, C = x.shape; Lk = y.shape[1]
55
+ sp = lambda t, L: t.view(B, L, self.h, self.hd).transpose(1, 2)
56
+ q, k, v = sp(self.q(x), Lq), sp(self.k(y), Lk), sp(self.v(y), Lk)
57
+ q = q * rope_q[0] + rotate_half(q) * rope_q[1]; k = k * rope_k[0] + rotate_half(k) * rope_k[1]
58
+ mask = None
59
+ if key_pad is not None:
60
+ mask = ~key_pad[:, None, None, :]
61
+ if causal:
62
+ cm = torch.ones(Lq, Lk, dtype=torch.bool, device=x.device).tril()[None, None]
63
+ mask = cm if mask is None else (mask & cm)
64
+ probs = None
65
+ if self.manual:
66
+ att = (q @ k.transpose(-2, -1)) / math.sqrt(self.hd)
67
+ if mask is not None:
68
+ att = att.masked_fill(~mask, float("-inf"))
69
+ probs = att.softmax(-1); out = probs @ v
70
+ else:
71
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)
72
+ return self.o(out.transpose(1, 2).reshape(B, Lq, C)), probs
73
+
74
+
75
+ def ffn(d):
76
+ return nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
77
+
78
+
79
+ class EncLayer(nn.Module):
80
+ def __init__(self, d, h):
81
+ super().__init__(); self.ln1 = nn.LayerNorm(d); self.att = Attn(d, h); self.ln2 = nn.LayerNorm(d); self.ff = ffn(d)
82
+
83
+ def forward(self, x, rope, pad):
84
+ n = self.ln1(x); x = x + self.att(n, n, rope, rope, key_pad=pad)[0]
85
+ return x + self.ff(self.ln2(x))
86
+
87
+
88
+ class CoreBlock(nn.Module):
89
+ """latent self-attention -> cross-attention to the full memory -> FFN (pre-LN residual branches)."""
90
+
91
+ def __init__(self, d, h, manual_cross):
92
+ super().__init__()
93
+ self.ln1 = nn.LayerNorm(d); self.self_att = Attn(d, h); self.ln2 = nn.LayerNorm(d); self.cross = Attn(d, h, manual=manual_cross)
94
+ self.ln3 = nn.LayerNorm(d); self.ff = ffn(d)
95
+
96
+ def forward(self, z, M, rope_z, rope_m, pad):
97
+ n = self.ln1(z); z = z + self.self_att(n, n, rope_z, rope_z)[0]
98
+ c, probs = self.cross(self.ln2(z), M, rope_z, rope_m, key_pad=pad); z = z + c
99
+ return z + self.ff(self.ln3(z)), probs
100
+
101
+
102
+ class Core(nn.Module):
103
+ def __init__(self, d, h, n_blocks=2):
104
+ super().__init__()
105
+ # the LAST block's cross-attention is the (potentially) supervised read: always the manual path, in every condition
106
+ self.blocks = nn.ModuleList([CoreBlock(d, h, manual_cross=(i == n_blocks - 1)) for i in range(n_blocks)])
107
+
108
+ def forward(self, z, M, rope_z, rope_m, pad):
109
+ probs = None
110
+ for blk in self.blocks:
111
+ z, p = blk(z, M, rope_z, rope_m, pad)
112
+ probs = p if p is not None else probs
113
+ return z, probs
114
+
115
+
116
+ class LatentExecutor(nn.Module):
117
+ def __init__(self, vocab_size: int, n_entities: int, d: int = 256, n_head: int = 2, n_latent: int = 4, enc_layers: int = 2,
118
+ core_blocks: int = 2, mode: str = "shared", n_untied: int = 8, rope_base: float = 100.0,
119
+ latent_pos0: int = 262, ans_pos0: int = 266):
120
+ super().__init__()
121
+ assert mode in ("shared", "single", "untied")
122
+ self.mode = mode; self.d = d; self.n_latent = n_latent; self.ans_pos0 = ans_pos0
123
+ self.tok_emb = nn.Embedding(vocab_size, d)
124
+ self.rope = Rope(d // n_head, rope_base)
125
+ self.enc = nn.ModuleList([EncLayer(d, n_head) for _ in range(enc_layers)]); self.enc_ln = nn.LayerNorm(d)
126
+ self.latents = nn.Parameter(torch.zeros(n_latent, d))
127
+ self.cores = nn.ModuleList([Core(d, n_head, core_blocks) for _ in range(n_untied if mode == "untied" else 1)])
128
+ self.ent_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, n_entities)) # auxiliary entity read-out of slot 0
129
+ self.z_ln = nn.LayerNorm(d)
130
+ self.dec_ln1 = nn.LayerNorm(d); self.dec_self = Attn(d, n_head); self.dec_ln2 = nn.LayerNorm(d); self.dec_cross = Attn(d, n_head)
131
+ self.dec_ln3 = nn.LayerNorm(d); self.dec_ff = ffn(d); self.dec_lnf = nn.LayerNorm(d); self.head = nn.Linear(d, vocab_size, bias=False)
132
+ self.register_buffer("latent_pos", torch.arange(latent_pos0, latent_pos0 + n_latent), persistent=False)
133
+
134
+ # ------------------------------------------------------------------------------------------------------------ init
135
+ def seeded_init(self, seed: int):
136
+ """Every tensor is drawn from its own generator keyed by (seed, parameter name): identically named tensors are
137
+ identical across conditions (A-E exactly; F's embedding / encoder / cores.0 / decoder / heads equal A's)."""
138
+ def gen(name):
139
+ return torch.Generator().manual_seed(int(hashlib.sha256(f"{seed}:{name}".encode()).hexdigest()[:12], 16))
140
+ with torch.no_grad():
141
+ for name, m in self.named_modules():
142
+ if isinstance(m, (nn.Linear, nn.Embedding)):
143
+ m.weight.copy_(torch.empty(m.weight.shape).normal_(0.0, 0.02, generator=gen(name + ".weight")))
144
+ if isinstance(m, nn.Linear) and m.bias is not None:
145
+ m.bias.zero_()
146
+ elif isinstance(m, nn.LayerNorm):
147
+ m.weight.fill_(1.0); m.bias.zero_()
148
+ self.latents.copy_(torch.empty(self.latents.shape).normal_(0.0, 0.02, generator=gen("latents")))
149
+ for core in self.cores: # residual branches of the core start as the identity map
150
+ for blk in core.blocks:
151
+ blk.self_att.o.weight.zero_(); blk.cross.o.weight.zero_(); blk.ff[2].weight.zero_()
152
+
153
+ # --------------------------------------------------------------------------------------------------------- forward
154
+ def encode(self, tok, pos, pad):
155
+ x = self.tok_emb(tok); rope = self.rope(pos)
156
+ for layer in self.enc:
157
+ x = layer(x, rope, pad)
158
+ return self.enc_ln(x), rope
159
+
160
+ def max_T(self) -> Optional[int]:
161
+ return {"shared": None, "single": 1, "untied": len(self.cores)}[self.mode]
162
+
163
+ def think(self, M, rope_m, pad, T: int, collect: bool = False, decode_at: Optional[List[int]] = None):
164
+ """Z_t = Core(Z_{t-1}, M), t = 1..T. collect -> per-loop slot-0 states and read-head probabilities (slot 0, head 0 of
165
+ the last block's cross-attention). decode_at -> also return the latent state after those loop counts (evaluation
166
+ of several budgets in one pass; identical to separate runs because the loop is deterministic)."""
167
+ mt = self.max_T(); assert mt is None or T <= mt, f"mode {self.mode} supports T <= {mt}"
168
+ B = M.shape[0]; z = self.latents[None].expand(B, -1, -1); rope_z = self.rope(self.latent_pos)
169
+ states, reads, snaps = [], [], {}
170
+ for t in range(T):
171
+ core = self.cores[t] if self.mode == "untied" else self.cores[0]
172
+ z, probs = core(z, M, rope_z, rope_m, pad)
173
+ if collect:
174
+ states.append(z[:, 0]); reads.append(probs[:, 0, 0, :])
175
+ if decode_at and (t + 1) in decode_at:
176
+ snaps[t + 1] = z
177
+ return z, states, reads, snaps
178
+
179
+ def decode(self, z, ans_in):
180
+ """z: final latent state (B, n_latent, C) -- the decoder never sees M. ans_in: (B, La) prefix starting with ANSWER."""
181
+ La = ans_in.shape[1]; pos = torch.arange(self.ans_pos0, self.ans_pos0 + La, device=ans_in.device)
182
+ rope_a = self.rope(pos); rope_z = self.rope(self.latent_pos); zk = self.z_ln(z)
183
+ y = self.tok_emb(ans_in); n = self.dec_ln1(y)
184
+ y = y + self.dec_self(n, n, rope_a, rope_a, causal=True)[0]
185
+ y = y + self.dec_cross(self.dec_ln2(y), zk, rope_a, rope_z)[0]
186
+ y = y + self.dec_ff(self.dec_ln3(y))
187
+ return self.head(self.dec_lnf(y))
188
+
189
+ def forward(self, tok, pos, pad, T: int, ans_in=None, collect: bool = False):
190
+ """The ONLY inputs are the complete packet (tok, pos, pad), the budget T and (teacher-forced) answer prefix for the
191
+ decoder. Gold entities / read addresses are never arguments of this function."""
192
+ M, rope_m = self.encode(tok, pos, pad)
193
+ z, states, reads, _ = self.think(M, rope_m, pad, T, collect=collect)
194
+ logits = self.decode(z, ans_in) if ans_in is not None else None
195
+ return dict(logits=logits, z=z, states=states, reads=reads, M=M)
196
+
197
+ @torch.no_grad()
198
+ def generate(self, z, answer_id: int, end_id: int, max_new: int = 4):
199
+ """greedy over the full vocabulary, no format forcing; returns (B, max_new) (tokens after END_ANSWER are padding = -1)."""
200
+ B = z.shape[0]; seq = torch.full((B, 1), answer_id, dtype=torch.long, device=z.device)
201
+ done = torch.zeros(B, dtype=torch.bool, device=z.device); out = []
202
+ for _ in range(max_new):
203
+ nxt = self.decode(z, seq)[:, -1].argmax(-1)
204
+ out.append(torch.where(done, torch.full_like(nxt, -1), nxt)); done = done | (nxt == end_id)
205
+ seq = torch.cat([seq, nxt[:, None]], 1)
206
+ return torch.stack(out, 1)
207
+
208
+ def param_groups_count(self):
209
+ cnt = lambda ms: sum(p.numel() for m in ms for p in m.parameters())
210
+ return dict(embedding=self.tok_emb.weight.numel(), encoder=cnt([self.enc, self.enc_ln]), latents=self.latents.numel(),
211
+ core_total=cnt([self.cores]), core_single=cnt([self.cores[0]]), n_cores=len(self.cores),
212
+ decoder=cnt([self.z_ln, self.dec_ln1, self.dec_self, self.dec_ln2, self.dec_cross, self.dec_ln3, self.dec_ff, self.dec_lnf]),
213
+ output_head=self.head.weight.numel(), aux_entity_head=cnt([self.ent_head]), total=sum(p.numel() for p in self.parameters()))
for Minegishi/reproducibility/model/loop_gpt.py ADDED
@@ -0,0 +1,216 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Full-sequence weight-shared loop Transformer aligned with "Loop, Think, & Generalize" v2
3
+ (docs/experiments_paper_aligned_loop_2026-09-19.md §1) + a per-sample stop head for the learned-halting arm
4
+ (docs/experiments_learned_halting_2026-09-19.md §1).
5
+
6
+ H_0 = Embed([x, r_1, ..., r_d]) one embedding pass, no positional encoding (NoPE), no input re-injection
7
+ H_t = F_theta(H_{t-1}) the SAME 4-block causal GPT-2-style stack, applied once per loop
8
+ z_t = LN_f(H_t[last valid input position])
9
+ a_t = softmax(W_tied z_t) LM head tied to the input embedding (same storage)
10
+ p_t = sigmoid(w_stop . z_t + b_stop) stop head (instantiated in every arm, only trained in arm H)
11
+
12
+ GPT-2 block: x + attn(LN(x)); x + mlp(LN(x)), GELU(tanh approximation = gelu_new), LayerNorm eps 1e-5, every dropout 0.
13
+ Init: N(0, 0.02) for embeddings / linear weights, zero biases, and ZERO weights for the two residual output projections of every
14
+ block (attention c_proj and MLP c_proj), so F_theta starts as the identity. Every tensor is drawn from its own generator keyed by
15
+ (seed, parameter name): the backbone init is identical across arms and independent of the stop head.
16
+ Right padding + causal attention: a valid position never attends to a later (pad) position, so no padding mask is needed
17
+ (checked in scripts/test_loop_gpt.py).
18
+ """
19
+ import hashlib
20
+
21
+ import torch
22
+ import torch.nn as nn
23
+ import torch.nn.functional as F
24
+
25
+
26
+ def _rot(x):
27
+ d = x.size(-1); return torch.cat([-x[..., d // 2:], x[..., :d // 2]], -1)
28
+
29
+
30
+ class Block(nn.Module):
31
+ def __init__(self, d, n_head):
32
+ super().__init__()
33
+ self.n_head = n_head; self.rope = None # (cos, sin) set by LoopGPT when pos == "rope"
34
+ self.window = 0 # > 0: local causal attention, a token sees itself and the previous `window` tokens
35
+ self.ln_1 = nn.LayerNorm(d, eps=1e-5); self.c_attn = nn.Linear(d, 3 * d); self.attn_proj = nn.Linear(d, d)
36
+ self.ln_2 = nn.LayerNorm(d, eps=1e-5); self.c_fc = nn.Linear(d, 4 * d); self.mlp_proj = nn.Linear(4 * d, d)
37
+
38
+ def forward(self, x, kv=None, strict_prev=False):
39
+ """kv: optional separate source for keys / values (same LayerNorm and projections); strict_prev: position i attends ONLY to i - 1
40
+ (position 0 to itself). Both are used by the stateless re-embedding modes, where a position's own stream is always its token embedding."""
41
+ B, L, C = x.shape
42
+ q, k, v = self.c_attn(self.ln_1(x)).split(C, dim=-1)
43
+ if kv is not None:
44
+ _, k, v = self.c_attn(self.ln_1(kv)).split(C, dim=-1)
45
+ sp = lambda t: t.view(B, L, self.n_head, C // self.n_head).transpose(1, 2)
46
+ q, k, v = sp(q), sp(k), sp(v)
47
+ if self.rope is not None: # relative positions: rotate Q / K after their projections, V untouched
48
+ cos, sin = self.rope[0][:L].to(q.dtype), self.rope[1][:L].to(q.dtype)
49
+ q = q * cos + _rot(q) * sin; k = k * cos + _rot(k) * sin
50
+ if strict_prev:
51
+ i = torch.arange(L, device=x.device); m = (i[None, :] == (i[:, None] - 1).clamp(min=0))
52
+ a = F.scaled_dot_product_attention(q, k, v, attn_mask=m)
53
+ elif self.window > 0:
54
+ i = torch.arange(L, device=x.device); m = (i[None, :] <= i[:, None]) & (i[None, :] >= i[:, None] - self.window)
55
+ a = F.scaled_dot_product_attention(q, k, v, attn_mask=m)
56
+ else:
57
+ a = F.scaled_dot_product_attention(q, k, v, is_causal=True)
58
+ x = x + self.attn_proj(a.transpose(1, 2).reshape(B, L, C))
59
+ return x + self.mlp_proj(F.gelu(self.c_fc(self.ln_2(x)), approximate="tanh"))
60
+
61
+
62
+ class LoopGPT(nn.Module):
63
+ def __init__(self, vocab_size: int, d: int = 768, n_head: int = 12, n_layer: int = 4, pos: str = "nope", rope_base: float = 100.0, max_pos: int = 512, attn_window: int = 0,
64
+ reembed: int = 0, ent_range=(0, 0), rel_range=(0, 0), untied: bool = False, untied_group: int = 1,
65
+ nr_init: bool = False, ready_state: bool = False):
66
+ super().__init__()
67
+ assert pos in ("nope", "rope")
68
+ self.d = d; self.pos = pos; self.attn_window = attn_window
69
+ # untied = the VANILLA (non-looped) control: the n_layer blocks are n_layer DIFFERENT layers applied once each, in order (step t uses
70
+ # block t). Everything else (window, re-embedding between layers, tied LM head read after any layer) is unchanged. A network of L
71
+ # layers has no layer L + 1: further "loops" are no-ops, so one hop per layer can never go beyond depth L.
72
+ self.untied = bool(untied)
73
+ # untied_group = layers per STEP of the vanilla control (= blocks per loop of the tied model it is compared with): step t applies layers t*g .. (t+1)*g-1,
74
+ # so a control with n_layer = g * L has L steps = the effective depth of the tied model run for L loops; the anchor / consistency losses treat one step as one loop.
75
+ self.untied_group = int(untied_group) if self.untied else 1
76
+ assert n_layer % self.untied_group == 0 and (self.untied_group == 1 or not reembed), "untied_group must divide n_layer; re-embedding between steps is defined for group 1"
77
+ # LEARNED HALTING ON THE STATELESS LOOP. nr_init: before anything has been computed a relation position holds a learned 'not ready' vector instead of its own
78
+ # token embedding. ready_state: the stop head is read at EVERY relation position and the state handed on is rho * E + (1 - rho) * nr (rho = the head's stop
79
+ # probability = 'my entity is ready'), so 'not ready' travels along the chain exactly like an entity does: a position that reads nr answers nr. One function
80
+ # f(token, previous state) -> (entity, ready), shared by all positions and loops; the halting probability of a query is rho at its last position.
81
+ self.nr_init = bool(nr_init or ready_state); self.ready_state = bool(ready_state)
82
+ if self.nr_init:
83
+ assert reembed in (2, 3) and not untied, "nr_init / ready_state are defined for the tied stateless modes 2 / 3"
84
+ self.nr = nn.Parameter(torch.zeros(d))
85
+ self.reembed = reembed; self.ent_range = tuple(ent_range); self.rel_range = tuple(rel_range)
86
+ self.wte = nn.Embedding(vocab_size, d)
87
+ self.blocks = nn.ModuleList([Block(d, n_head) for _ in range(n_layer)])
88
+ for blk in self.blocks:
89
+ blk.window = attn_window
90
+ self.ln_f = nn.LayerNorm(d, eps=1e-5)
91
+ self.stop = nn.Linear(d, 1) # 769 parameters; unused (grad None) outside arm H
92
+ if pos == "rope": # parameter-free: the backbone init is identical to the NoPE model's
93
+ hd = d // n_head; inv = 1.0 / (rope_base ** (torch.arange(0, hd, 2).float() / hd))
94
+ f = torch.einsum("l,d->ld", torch.arange(max_pos).float(), inv); emb = torch.cat([f, f], -1)
95
+ self.register_buffer("rope_cos", emb.cos(), persistent=False); self.register_buffer("rope_sin", emb.sin(), persistent=False)
96
+
97
+ def seeded_init(self, seed: int):
98
+ def gen(name):
99
+ return torch.Generator().manual_seed(int(hashlib.sha256(f"{seed}:{name}".encode()).hexdigest()[:12], 16))
100
+ with torch.no_grad():
101
+ for name, m in self.named_modules():
102
+ if isinstance(m, (nn.Linear, nn.Embedding)):
103
+ m.weight.copy_(torch.empty(m.weight.shape).normal_(0.0, 0.02, generator=gen(name + ".weight")))
104
+ if isinstance(m, nn.Linear):
105
+ m.bias.zero_()
106
+ elif isinstance(m, nn.LayerNorm):
107
+ m.weight.fill_(1.0); m.bias.zero_()
108
+ if self.nr_init:
109
+ self.nr.copy_(torch.empty(self.nr.shape).normal_(0.0, 0.02, generator=gen("nr")))
110
+ for blk in self.blocks: # residual output projections: scale 0
111
+ blk.attn_proj.weight.zero_(); blk.mlp_proj.weight.zero_()
112
+
113
+ def embed(self, tok):
114
+ return self.wte(tok)
115
+
116
+ def step(self, h):
117
+ """one loop = one pass through the shared stack"""
118
+ if self.pos == "rope":
119
+ for blk in self.blocks:
120
+ blk.rope = (self.rope_cos, self.rope_sin)
121
+ for blk in self.blocks:
122
+ h = blk(h)
123
+ return h
124
+
125
+ # ---- explicit loop state (used by arm D and by fixed-budget evaluation). Without re-embedding it is just the hidden states.
126
+ def embed_state(self, tok):
127
+ base = self.wte(tok); rel = ((tok >= self.rel_range[0]) & (tok < self.rel_range[1])).unsqueeze(-1).to(base.dtype)
128
+ h = base
129
+ if self.nr_init: # nothing computed yet: relation positions hold 'not ready'
130
+ h = base * (1 - rel) + rel * self.nr if self.reembed == 3 else base + rel * self.nr
131
+ return dict(h=h, base=base, rel=rel)
132
+
133
+ def step_state(self, st):
134
+ """-> (out, next_state). out = F_theta(h) is what the read-out sees. With reembed=True the state handed to the NEXT loop is rebuilt as
135
+ token embedding + (at relation positions) the embedding-space expectation of the position's own entity prediction:
136
+ h_next[p] = wte[tok_p] + sum_e softmax(logits_entities(out[p]))_e * wte[e]
137
+ so a computed entity is handed on in the SAME format as a raw entity token, nothing accumulates across loops, and every position's
138
+ state is a readable distribution over entities (a latent, in-place chain of thought). No label and no loop index is involved."""
139
+ if self.untied:
140
+ return self._step_untied(st)
141
+ if not self.reembed:
142
+ out = self.step(st["h"]); return out, dict(st, h=out)
143
+ W = self.wte.weight[self.ent_range[0]:self.ent_range[1]]
144
+ if self.reembed == 1:
145
+ out = self.step(st["h"])
146
+ else:
147
+ # STATELESS positions (modes 2 / 3): a position's own stream is ALWAYS its token embedding; the only thing that changes from loop to
148
+ # loop is what it can read from the previous position: out_p(t) = f(token_p, state_{p-1}(t-1)). Its output cannot depend on its own
149
+ # earlier outputs, so there is no private clock or accumulated garbage, and a finished prefix stays finished by construction.
150
+ assert len(self.blocks) == 1, "stateless re-embedding is defined for one block per loop"
151
+ out = self.blocks[0](st["base"], kv=st["h"], strict_prev=True)
152
+ if self.reembed in (4, 5):
153
+ return out, dict(st, h=self._raw_state(st, out))
154
+ pe = torch.softmax((self.ln_f(out) @ W.t()).float(), -1).to(out.dtype); E = pe @ W
155
+ if self.ready_state:
156
+ rho = torch.sigmoid(self.stop(self.ln_f(out)).float()).to(out.dtype); E = rho * E + (1 - rho) * self.nr
157
+ if self.reembed == 3: # mode 3: the next position sees ONLY the entity (the format of a raw entity token)
158
+ return out, dict(st, h=st["base"] * (1 - st["rel"]) + st["rel"] * E)
159
+ return out, dict(st, h=st["base"] + st["rel"] * E) # modes 1 / 2: token embedding + expected entity embedding
160
+
161
+ def _raw_state(self, st, out):
162
+ """Modes 4 / 5 = the BOTTLENECK ablation of modes 3 / 2: same stateless structure (own stream = token embedding, strict previous-position read),
163
+ but what a relation position hands on is its hidden vector itself instead of the embedding-space expectation of its entity prediction
164
+ (no LM-head read-out, no softmax, no projection back onto entity embeddings). Mode 4 <-> mode 3: only the computed update (out - token embedding);
165
+ mode 5 <-> mode 2: token embedding + update (= out)."""
166
+ upd = out - st["base"] if self.reembed == 4 else out
167
+ return st["base"] * (1 - st["rel"]) + st["rel"] * upd
168
+
169
+ def step_group(self, h, t):
170
+ """VANILLA control, plain residual stream: step t = layers t*g .. (t+1)*g - 1 in order (g = untied_group). No layer left -> h unchanged."""
171
+ blocks = self.blocks[t * self.untied_group:(t + 1) * self.untied_group]
172
+ if self.pos == "rope":
173
+ for blk in blocks:
174
+ blk.rope = (self.rope_cos, self.rope_sin)
175
+ for blk in blocks:
176
+ h = blk(h)
177
+ return h
178
+
179
+ @property
180
+ def n_groups(self):
181
+ return len(self.blocks) // self.untied_group if self.untied else None
182
+
183
+ def _step_untied(self, st):
184
+ t = st.get("t", 0)
185
+ if t * self.untied_group >= len(self.blocks): # no layer left
186
+ return st["out"], st
187
+ if self.reembed in (2, 3, 4, 5):
188
+ blk = self.blocks[t]
189
+ if self.pos == "rope":
190
+ blk.rope = (self.rope_cos, self.rope_sin)
191
+ out = blk(st["base"], kv=st["h"], strict_prev=True)
192
+ else:
193
+ out = self.step_group(st["h"], t)
194
+ nxt = dict(st, t=t + 1, out=out)
195
+ if not self.reembed:
196
+ nxt["h"] = out; return out, nxt
197
+ if self.reembed in (4, 5):
198
+ nxt["h"] = self._raw_state(st, out); return out, nxt
199
+ W = self.wte.weight[self.ent_range[0]:self.ent_range[1]]
200
+ pe = torch.softmax((self.ln_f(out) @ W.t()).float(), -1).to(out.dtype); E = pe @ W
201
+ nxt["h"] = st["base"] * (1 - st["rel"]) + st["rel"] * E if self.reembed == 3 else st["base"] + st["rel"] * E
202
+ return out, nxt
203
+
204
+ def read(self, h, last):
205
+ """z_t at the last valid input position. last: (B,) index of the last real token."""
206
+ # gather (backward = scatter_add) instead of advanced indexing, whose backward is a slow sort-based index_put
207
+ return self.ln_f(torch.gather(h, 1, last[:, None, None].expand(-1, 1, h.shape[-1])).squeeze(1))
208
+
209
+ def logits(self, z):
210
+ return z @ self.wte.weight.t() # tied head: the SAME parameter storage as the input embedding
211
+
212
+ def stop_logit(self, z):
213
+ return self.stop(z).squeeze(-1)
214
+
215
+ def backbone_names(self):
216
+ return [n for n, _ in self.named_parameters() if not n.startswith("stop.") and n != "nr"]
for Minegishi/reproducibility/objectives.py ADDED
@@ -0,0 +1,467 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Matched shallow supervision and separable consistency objectives.
2
+
3
+ The shallow trajectory ALWAYS runs four loops. Every d1 row is supervised at
4
+ loops 1, 2, 3 and every d2 row at loops 2, 3, 4, independently of loss switches.
5
+ Anchor targets and normalisation reproduce Miyabi ``loss_anchor(extra=2,
6
+ n_loops=2)``. Its four disjoint numerators are execute, wait, hold, and root;
7
+ root includes the initial entity and any random-entity prefix. Each numerator
8
+ uses the SAME full valid-target squared-norm denominator in each loop.
9
+
10
+ CS constructs detached ideal states and excludes the next frontier. Its terms
11
+ all use the SAME number of full valid positions. In particular, deleting wait
12
+ does not increase the coefficient of hold. The four main experimental modes
13
+ are root_only, wait_only (= root + wait), hold_only (= root + hold), and full.
14
+ This keeps the initial-entity constraint fixed in the wait x hold factorial.
15
+ ``off`` is available as a separate control, and differs from root_only.
16
+ The additional ``wait_no_root`` mode selects only the wait numerator while
17
+ keeping the same full nonfrontier-position denominator. It removes the root
18
+ CS term, not the root or hold targets in the separate shallow anchor loss.
19
+
20
+ All returned statistics are detached tensors; selected loss tensors retain
21
+ their graphs. There are no tensor-to-host conversions in the loss functions,
22
+ so they can be used inside the existing CUDA-graph training path.
23
+ """
24
+
25
+ import hashlib
26
+ import math
27
+ from typing import Dict, Optional, Tuple
28
+
29
+ import numpy as np
30
+
31
+ import torch
32
+ import torch.nn.functional as F
33
+
34
+
35
+ Tensor = torch.Tensor
36
+ SHALLOW_LOOPS = 4
37
+ SHALLOW_EXTRA = 2
38
+ ANCHOR_MODES = ("off", "full", "execute_only", "wait_only", "hold_only", "root_only")
39
+ SIMPLE_CS_MODES = ("t0_wait_denoise", "relation_identity")
40
+ CS_MODES = ("off", "root_only", "wait_only", "hold_only", "full", "wait_no_root") + SIMPLE_CS_MODES
41
+ CS_COMPONENTS = {
42
+ "off": (),
43
+ "root_only": ("root",),
44
+ "wait_only": ("root", "wait"),
45
+ "hold_only": ("root", "hold"),
46
+ "full": ("root", "wait", "hold"),
47
+ "wait_no_root": ("wait",),
48
+ "t0_wait_denoise": ("wait",),
49
+ "relation_identity": ("relation_identity",),
50
+ }
51
+
52
+
53
+ def validate_simple_cs_config(mode: str, depth: int, rows: int, noise_scale: float,
54
+ noise_kind: str = "embedding", noise_scope: str = "wait") -> None:
55
+ """Host-side validation before capture; the new batch freezes CS depth 32."""
56
+ if not math.isfinite(noise_scale) or noise_scale < 0:
57
+ raise ValueError("cs_noise_scale must be finite and nonnegative")
58
+ if mode in SIMPLE_CS_MODES and (depth != 32 or rows != 32):
59
+ raise ValueError("simple-CS experiment requires consist_depth=32 and consist_rows=32")
60
+ if mode == "t0_wait_denoise" and not 0 < noise_scale < 1:
61
+ raise ValueError("t0_wait_denoise requires 0 < cs_noise_scale < 1")
62
+ if mode != "t0_wait_denoise" and noise_scale != 0:
63
+ raise ValueError("cs_noise_scale is only meaningful for t0_wait_denoise")
64
+ if noise_kind not in ("embedding", "nearest", "isotropic") or noise_scope not in ("all", "wait"):
65
+ raise ValueError("unknown CS noise kind or scope")
66
+ if noise_kind == "embedding" and noise_scope != "wait":
67
+ raise ValueError("legacy embedding-norm perturbations support wait scope only")
68
+ if mode != "t0_wait_denoise" and (noise_kind != "embedding" or noise_scope != "wait"):
69
+ raise ValueError("CS noise kind/scope are only meaningful for t0_wait_denoise")
70
+
71
+
72
+ def simple_cs_metadata(mode: str, noise_scale: float, noise_kind: str = "embedding",
73
+ noise_scope: str = "wait") -> Dict:
74
+ """Explicit manifest semantics; coefficients match old expected wait mass."""
75
+ if mode == "t0_wait_denoise":
76
+ result = dict(denominator="rows_times_depth", internal_weight=0.5,
77
+ tfire=0, token_layout="entity_then_32_relations",
78
+ constrained_positions="2..32", input_noise_positions="0..32" if noise_scope == "all" else "2..32",
79
+ noise_relative_norm=noise_scale,
80
+ noise_kind=noise_kind, noise_scope=noise_scope,
81
+ noise_distribution="independent Gaussian directions, normalized per token",
82
+ noise_rng="independent torch.Generator; fill outside CUDA graph; checkpoint state",
83
+ target="clean detached input embeddings", input_detached=True,
84
+ note="root and frontier are neither noised nor directly constrained; shallow full anchor unchanged")
85
+ if noise_kind in ("nearest", "isotropic"):
86
+ result.update(noise_radius="uniform_fraction_of_nearest_embedding_distance",
87
+ noise_alpha_distribution="Uniform[0, noise_scale)",
88
+ nearest_candidates="all_entities_and_relations_excluding_self_and_special_tokens",
89
+ nearest_metric="Euclidean L2 of current detached embeddings; direct cdist (no TF32 distance matmul)",
90
+ noise_distribution="toward nearest embedding" if noise_kind == "nearest" else "normalized Gaussian direction with identical nearest-distance radius",
91
+ note="noise mask may include root/frontier; loss mask always only waiting positions 2..32; shallow full anchor unchanged")
92
+ return result
93
+ if mode == "relation_identity":
94
+ return dict(denominator="all_33_positions", internal_weight=31 / 64,
95
+ tfire=None, token_layout="33_relations_no_entity",
96
+ constrained_positions="0..32", input_noise_positions="none",
97
+ noise_relative_norm=0.0, target="clean detached input embeddings",
98
+ input_detached=True,
99
+ note="one-step identity on all-relation sequences; no entity or execution frontier")
100
+ return dict(denominator="all_nonfrontier_positions",
101
+ note="label-free one-step consistency on constructed states; frontier excluded")
102
+
103
+
104
+ def sample_simple_cs_tokens(rng, mode: str, rows: int, depth: int,
105
+ entity_start: int, entities: int,
106
+ relation_start: int, relations: int):
107
+ """Sample token IDs outside the graph, using the saved named CS stream.
108
+
109
+ No graph-state or gold intermediate entity labels are consulted. The
110
+ all-relation mode never samples an entity. ``tfire`` and ``ent_ids`` are
111
+ compatibility buffers only; the new loss branches do not use them.
112
+ """
113
+ if mode == "t0_wait_denoise":
114
+ tok = np.empty((rows, depth + 1), dtype=np.int64)
115
+ tok[:, 0] = rng.integers(entity_start, entity_start + entities, rows)
116
+ tok[:, 1:] = rng.integers(relation_start, relation_start + relations, (rows, depth))
117
+ elif mode == "relation_identity":
118
+ tok = rng.integers(relation_start, relation_start + relations, (rows, depth + 1), dtype=np.int64)
119
+ else:
120
+ raise ValueError("not a simple-CS mode")
121
+ return tok, np.zeros(rows, dtype=np.int64), np.zeros_like(tok)
122
+
123
+
124
+ def make_cs_noise_generator(seed: int, device) -> torch.Generator:
125
+ """Independent stream: capture/warm-up and global dropout RNG cannot consume it."""
126
+ stream_seed = int.from_bytes(hashlib.sha256(f"{seed}:simple_cs_noise".encode()).digest()[:8], "little")
127
+ return torch.Generator(device=device).manual_seed(stream_seed)
128
+
129
+
130
+ def fill_cs_noise_(buffer: Tensor, generator: torch.Generator) -> None:
131
+ """Call once per actual training update, outside CUDA graph capture/replay."""
132
+ with torch.no_grad():
133
+ buffer.normal_(generator=generator)
134
+
135
+
136
+ def fill_cs_alpha_(buffer: Tensor, generator: torch.Generator) -> None:
137
+ """Unit-uniform radii; multiplied by the frozen scale inside the graph."""
138
+ with torch.no_grad():
139
+ buffer.uniform_(generator=generator)
140
+
141
+
142
+ def _mode(mode: str, allowed: Tuple[str, ...]) -> str:
143
+ mode = mode.replace("-", "_")
144
+ mode = {"execute": "execute_only", "wait": "wait_only", "hold": "hold_only", "root": "root_only"}.get(mode, mode)
145
+ if mode not in allowed:
146
+ raise ValueError("unknown loss mode {!r}; choose {}".format(mode, allowed))
147
+ return mode
148
+
149
+
150
+ def _positive_denominator(value: Tensor) -> Tensor:
151
+ # Exactly unchanged for ordinary nonzero embeddings; also defines zero/zero
152
+ # as zero error for degenerate all-zero fixtures rather than returning NaN.
153
+ return value.clamp_min(torch.finfo(value.dtype).tiny)
154
+
155
+
156
+ def matched_ce_mask(depth: Tensor) -> Tensor:
157
+ """Return (4, N) mask: d1 -> loops 1..3; d2 -> loops 2..4.
158
+
159
+ ``depth`` must contain semantic depths 1 or 2, NOT absolute read positions.
160
+ Input/data validation should happen before CUDA graph capture.
161
+ """
162
+ loops = torch.arange(1, SHALLOW_LOOPS + 1, device=depth.device)[:, None]
163
+ return (loops >= depth[None, :]) & (loops <= depth[None, :] + SHALLOW_EXTRA)
164
+
165
+
166
+ def shallow_masks(depth: Tensor, width: int, loop: int, off: Optional[Tensor] = None) -> Dict[str, Tensor]:
167
+ """Disjoint post-loop anchor masks, including left entity prefixes.
168
+
169
+ Execute means *first* completion in this loop. Hold means a relation
170
+ completed earlier. Root means the head entity and all preceding fillers.
171
+ Right padding belongs to no component.
172
+ """
173
+ if loop < 1:
174
+ raise ValueError("loop is one-based")
175
+ off = torch.zeros_like(depth) if off is None else off
176
+ relative = torch.arange(width, device=depth.device)[None, :] - off[:, None]
177
+ valid = relative <= depth[:, None]
178
+ return {
179
+ "valid": valid,
180
+ "root": valid & (relative <= 0),
181
+ "execute": valid & (relative == loop),
182
+ "wait": valid & (relative > loop),
183
+ "hold": valid & (relative >= 1) & (relative < loop),
184
+ }
185
+
186
+
187
+ def consistency_masks(tfire: Tensor, width: int) -> Dict[str, Tensor]:
188
+ """Masks for a canonical state AFTER tfire loops, before its next update.
189
+
190
+ Position zero is the root. Positions 1..tfire hold completed entities.
191
+ Position tfire+1 is the frontier and is excluded from every CS loss.
192
+ Remaining positions wait. tfire=0 and tfire=depth are both supported.
193
+ Synthetic rows have no padding, matching the original Miyabi objective.
194
+ """
195
+ pos = torch.arange(width, device=tfire.device)[None, :]
196
+ time = tfire[:, None]
197
+ frontier = pos == time + 1
198
+ return {
199
+ "keep": ~frontier,
200
+ "frontier": frontier,
201
+ "root": (pos == 0).expand(tfire.shape[0], -1),
202
+ "hold": (pos >= 1) & (pos <= time),
203
+ "wait": pos > time + 1,
204
+ }
205
+
206
+
207
+ def _select_anchor(terms: Dict[str, Tensor], mode: str) -> Tensor:
208
+ if mode == "off":
209
+ return terms["full"].new_zeros(())
210
+ if mode == "full":
211
+ return terms["full"]
212
+ return terms[mode.removesuffix("_only")]
213
+
214
+
215
+ def _select_consistency(terms: Dict[str, Tensor], mode: str) -> Tensor:
216
+ if mode == "off":
217
+ return terms["full"].new_zeros(())
218
+ if mode == "full":
219
+ return terms["full"]
220
+ if mode == "root_only":
221
+ return terms["root"]
222
+ if mode == "wait_no_root":
223
+ return terms["wait"]
224
+ return terms["root"] + terms[mode.removesuffix("_only")]
225
+
226
+
227
+ def loss_shallow(model, tok: Tensor, depth: Tensor, tgt: Tensor, pv: Tensor,
228
+ off: Optional[Tensor] = None, anchor_mode: str = "full") -> Tuple[Tensor, Tensor, Dict[str, Tensor]]:
229
+ """Return (per-row matched CE, selected anchor, detached statistics).
230
+
231
+ Model API: embed(tok), wte(ids), step(hidden), read(hidden, last), logits(z).
232
+ ``pv`` contains the correct entity token at each relation's first completion;
233
+ those labels enter detached anchor targets only, never the forward state.
234
+ Every mode computes the identical forward trajectory and answer CE.
235
+ """
236
+ mode = _mode(anchor_mode, ANCHOR_MODES)
237
+ if getattr(model, "untied", False):
238
+ raise ValueError("this matched protocol is for the tied loop model")
239
+ if tok.ndim != 2 or pv.shape != tok.shape:
240
+ raise ValueError("tok and pv must have the same (rows, width) shape")
241
+ if depth.shape != tok.shape[:1] or tgt.shape != depth.shape:
242
+ raise ValueError("depth and tgt must contain one value per row")
243
+ if off is not None and off.shape != depth.shape:
244
+ raise ValueError("off must contain one prefix length per row")
245
+ off = torch.zeros_like(depth) if off is None else off
246
+ emb = model.embed(tok)
247
+ e_tok = emb.detach()
248
+ e_pv = model.wte(pv).detach()
249
+ h = emb
250
+ last = off + depth
251
+ terms = {name: emb.new_zeros(()) for name in ("full", "root", "execute", "wait", "hold")}
252
+ ces = []
253
+ for loop in range(1, SHALLOW_LOOPS + 1):
254
+ h = model.step(h)
255
+ masks = shallow_masks(depth, tok.shape[1], loop, off)
256
+ fired = masks["execute"] | masks["hold"]
257
+ target = torch.where(fired[..., None], e_pv, e_tok)
258
+ error = (h - target).square().sum(-1)
259
+ denominator = _positive_denominator((target.square().sum(-1) * masks["valid"]).sum())
260
+ terms["full"] = terms["full"] + (error * masks["valid"]).sum() / denominator
261
+ for name in ("root", "execute", "wait", "hold"):
262
+ terms[name] = terms[name] + (error * masks[name]).sum() / denominator
263
+ logits = model.logits(model.read(h, last)).float()
264
+ ces.append(F.cross_entropy(logits, tgt, reduction="none"))
265
+ terms = {name: value / SHALLOW_LOOPS for name, value in terms.items()}
266
+ ces = torch.stack(ces)
267
+ ce_mask = matched_ce_mask(depth).to(ces.dtype)
268
+ ce_per_row = (ces * ce_mask).sum(0) / ce_mask.sum(0)
269
+ selected = _select_anchor(terms, mode)
270
+ stats = {"anchor_" + name: value.detach() for name, value in terms.items()}
271
+ stats.update(anchor_selected=selected.detach(), ce=ce_per_row.mean().detach(), ce_by_loop=ces.mean(1).detach())
272
+ return ce_per_row, selected, stats
273
+
274
+
275
+ def construct_consistency_state(model, tok: Tensor, tfire: Tensor, ent_ids: Tensor) -> Tensor:
276
+ """Construct a detached ideal state; it contains no deep-chain gold labels."""
277
+ if tok.ndim != 2 or ent_ids.shape != tok.shape or tfire.shape != tok.shape[:1]:
278
+ raise ValueError("expected tok/ent_ids (rows, width) and tfire (rows,)")
279
+ masks = consistency_masks(tfire, tok.shape[1])
280
+ return torch.where(masks["hold"][..., None], model.wte(ent_ids), model.embed(tok)).detach()
281
+
282
+
283
+ def consistency_terms(output: Tensor, state: Tensor, tfire: Tensor) -> Dict[str, Tensor]:
284
+ """Differentiable CS components with a detached target and shared denominator.
285
+
286
+ The bounded per-position error is ||output-state||^2 /
287
+ (||output||^2 + ||state||^2), in [0, 2]. Root, wait, and hold components
288
+ sum to full, without reweighting any remaining term in an ablation.
289
+ This helper is also suitable for single-step stability probes.
290
+ """
291
+ if output.shape != state.shape or output.ndim != 3 or tfire.shape != output.shape[:1]:
292
+ raise ValueError("expected output/state (rows, width, hidden) and tfire (rows,)")
293
+ state = state.detach()
294
+ masks = consistency_masks(tfire, state.shape[1])
295
+ error = (output - state).square().sum(-1)
296
+ scale = _positive_denominator(output.square().sum(-1) + state.square().sum(-1))
297
+ bounded = error / scale
298
+ denominator = masks["keep"].to(bounded.dtype).sum().clamp_min(1)
299
+ terms = {name: (bounded * masks[name]).sum() / denominator for name in ("root", "wait", "hold")}
300
+ terms["full"] = (bounded * masks["keep"]).sum() / denominator
301
+ return terms
302
+
303
+
304
+ def initial_wait_state(clean: Tensor, noise: Tensor, noise_scale: float) -> Tensor:
305
+ """Detached t=0 input with perturbations only at waiting positions >=2.
306
+
307
+ ``noise`` contains raw Gaussian directions supplied by the batch sampler.
308
+ Each nonzero direction is rescaled to ``noise_scale * ||clean_token||``;
309
+ a zero direction (used during graph capture) produces no perturbation.
310
+ Root and first-relation frontier remain bitwise identical to clean input.
311
+ """
312
+ if clean.ndim != 3 or noise.shape != clean.shape:
313
+ raise ValueError("noise must have the (rows, width, hidden) clean-state shape")
314
+ if clean.shape[1] < 3 or not math.isfinite(noise_scale) or noise_scale < 0:
315
+ raise ValueError("t0 wait noise requires width >= 3 and a finite nonnegative scale")
316
+ clean, noise = clean.detach(), noise.detach()
317
+ direction_norm = noise.norm(dim=-1, keepdim=True)
318
+ direction = noise / _positive_denominator(direction_norm)
319
+ wait = torch.arange(clean.shape[1], device=clean.device)[None, :, None] >= 2
320
+ delta = direction * clean.norm(dim=-1, keepdim=True) * noise_scale
321
+ return torch.where(wait, clean + delta, clean).detach()
322
+
323
+
324
+ def nearest_embedding_geometry(model, tok: Tensor, candidate_ids: Tensor,
325
+ relation_ids: Tensor) -> Tuple[Tensor, Tensor]:
326
+ """Nearest nonself entity/relation embedding, detached and recomputed each step.
327
+
328
+ Query the B roots and relation vocabulary once. Candidate IDs include
329
+ only entity/relation ranges. Direct Euclidean cdist avoids TF32 Gram
330
+ cancellation; duplicate embedding vectors for distinct IDs may give zero.
331
+ """
332
+ if candidate_ids.ndim != 1 or candidate_ids.numel() < 2 or relation_ids.ndim != 1:
333
+ raise ValueError("need at least two candidate token IDs and a relation-ID vector")
334
+ query_ids = torch.cat((tok[:, 0], relation_ids))
335
+ query = model.wte(query_ids).detach()
336
+ candidates = model.wte(candidate_ids).detach()
337
+ distances = torch.cdist(query.float(), candidates.float(), p=2, compute_mode="donot_use_mm_for_euclid_dist")
338
+ distances = distances.masked_fill(query_ids[:, None] == candidate_ids[None, :], float("inf"))
339
+ nearest_distance, nearest_index = distances.min(dim=1)
340
+ neighbor = candidates[nearest_index]
341
+ lookup = torch.cat((torch.arange(tok.shape[0], device=tok.device)[:, None],
342
+ tok[:, 1:] - relation_ids[0] + tok.shape[0]), dim=1)
343
+ return neighbor[lookup].detach(), nearest_distance[lookup][..., None].detach()
344
+
345
+
346
+ def nearest_radius_state(clean: Tensor, neighbor: Tensor, nearest_distance: Tensor,
347
+ noise: Tensor, alpha: Tensor, noise_scale: float,
348
+ noise_kind: str, noise_scope: str) -> Tensor:
349
+ """Matched radius: U[0,scale) times nearest distance in either direction.
350
+
351
+ alpha is a U[0,1) batch buffer. Nearest direction points toward the nearest
352
+ nonself embedding; isotropic direction is a normalized Gaussian draw.
353
+ """
354
+ if clean.ndim != 3 or neighbor.shape != clean.shape or noise.shape != clean.shape:
355
+ raise ValueError("clean, neighbor, and noise must have the same 3-D shape")
356
+ if alpha.shape != clean.shape[:2] + (1,) or nearest_distance.shape != alpha.shape:
357
+ raise ValueError("alpha and nearest_distance require one scalar per token")
358
+ if noise_kind not in ("nearest", "isotropic") or noise_scope not in ("all", "wait"):
359
+ raise ValueError("invalid nearest-radius noise configuration")
360
+ clean, neighbor, nearest_distance = clean.detach(), neighbor.detach(), nearest_distance.detach()
361
+ alpha, noise = alpha.detach(), noise.detach()
362
+ if noise_kind == "nearest":
363
+ delta = alpha * noise_scale * (neighbor - clean)
364
+ else:
365
+ direction = noise / _positive_denominator(noise.norm(dim=-1, keepdim=True))
366
+ delta = alpha * noise_scale * nearest_distance * direction
367
+ if noise_scope == "wait":
368
+ mask = torch.arange(clean.shape[1], device=clean.device)[None, :, None] >= 2
369
+ delta = torch.where(mask, delta, torch.zeros_like(delta))
370
+ return (clean + delta).detach()
371
+
372
+
373
+ def loss_simple_consistency(model, tok: Tensor, mode: str,
374
+ noise: Optional[Tensor] = None,
375
+ noise_scale: float = 0.0, *, noise_kind: str = "embedding",
376
+ noise_scope: str = "wait", alpha: Optional[Tensor] = None,
377
+ candidate_ids: Optional[Tensor] = None,
378
+ relation_ids: Optional[Tensor] = None) -> Tuple[Tensor, Dict[str, Tensor]]:
379
+ """Two initial-state objectives with no sampled execution time or labels.
380
+
381
+ Coefficients match the expected selected-position mass of the old D=32
382
+ pure-wait loss: E[#wait]/32 = 15.5/32 = 31/64. This is an expectation
383
+ match, not an assertion that the distributions or gradient norms match.
384
+ """
385
+ if tok.ndim != 2 or tok.shape[1] != 33:
386
+ raise ValueError("simple-CS loss requires exactly 33 token positions")
387
+ clean = model.embed(tok).detach()
388
+ if mode == "t0_wait_denoise":
389
+ if noise is None:
390
+ raise ValueError("t0_wait_denoise needs an explicit batch noise tensor")
391
+ if noise_kind == "embedding":
392
+ state = initial_wait_state(clean, noise, noise_scale)
393
+ else:
394
+ if alpha is None or candidate_ids is None or relation_ids is None:
395
+ raise ValueError("nearest-radius noise requires alpha and entity/relation candidate IDs")
396
+ neighbor, nearest_distance = nearest_embedding_geometry(model, tok, candidate_ids, relation_ids)
397
+ state = nearest_radius_state(clean, neighbor, nearest_distance, noise, alpha, noise_scale, noise_kind, noise_scope)
398
+ output = model.step(state)
399
+ fixed_t0 = torch.zeros(tok.shape[0], dtype=torch.long, device=tok.device)
400
+ terms = consistency_terms(output, clean, fixed_t0)
401
+ selected = 0.5 * terms["wait"]
402
+ stats = {"cs_" + name: value.detach() for name, value in terms.items()}
403
+ elif mode == "relation_identity":
404
+ if noise is not None or noise_scale != 0:
405
+ raise ValueError("relation_identity has no noise")
406
+ output = model.step(clean)
407
+ error = (output - clean).square().sum(-1)
408
+ denominator = _positive_denominator(output.square().sum(-1) + clean.square().sum(-1))
409
+ identity = (error / denominator).mean()
410
+ selected = (31 / 64) * identity
411
+ stats = dict(cs_identity=identity.detach())
412
+ else:
413
+ raise ValueError("not a simple-CS mode")
414
+ stats["cs_selected"] = selected.detach()
415
+ return selected, stats
416
+
417
+
418
+ def loss_consistency(model, tok: Tensor, tfire: Tensor, ent_ids: Tensor,
419
+ mode: str = "full", *, noise: Optional[Tensor] = None,
420
+ noise_scale: float = 0.0, noise_kind: str = "embedding",
421
+ noise_scope: str = "wait", alpha: Optional[Tensor] = None,
422
+ candidate_ids: Optional[Tensor] = None,
423
+ relation_ids: Optional[Tensor] = None) -> Tuple[Tensor, Dict[str, Tensor]]:
424
+ """Return (selected one-step CS loss, detached component statistics).
425
+
426
+ All modes, including off/root_only, use exactly the same constructed state
427
+ and one model step. The caller may separately disable the entire CS branch
428
+ before the preregistered CS warm-up boundary.
429
+ """
430
+ mode = _mode(mode, CS_MODES)
431
+ if getattr(model, "untied", False):
432
+ raise ValueError("this consistency protocol is for the tied loop model")
433
+ if mode in SIMPLE_CS_MODES:
434
+ return loss_simple_consistency(model, tok, mode, noise, noise_scale,
435
+ noise_kind=noise_kind, noise_scope=noise_scope, alpha=alpha,
436
+ candidate_ids=candidate_ids, relation_ids=relation_ids)
437
+ if noise is not None or noise_scale != 0:
438
+ raise ValueError("noise is supported only by t0_wait_denoise")
439
+ state = construct_consistency_state(model, tok, tfire, ent_ids)
440
+ terms = consistency_terms(model.step(state), state, tfire)
441
+ selected = _select_consistency(terms, mode)
442
+ stats = {"cs_" + name: value.detach() for name, value in terms.items()}
443
+ stats["cs_selected"] = selected.detach()
444
+ return selected, stats
445
+
446
+
447
+ def loss_anchor(model, tok, depth, tgt, pv, extra=2, off=None, n_loops=2,
448
+ k_delay=0, pad_id=None, *, anchor_mode="full"):
449
+ """Drop-in original-signature wrapper -> (CE per row, selected anchor).
450
+
451
+ The experiment deliberately supports only matched shallow four-loop
452
+ supervision. n_loops=None also means two semantic loops, never width-1.
453
+ Padding ID is accepted for source compatibility; delayed starts are not.
454
+ """
455
+ if extra != 2 or n_loops not in (None, 2) or k_delay != 0:
456
+ raise ValueError("matched protocol requires extra=2, n_loops=2, k_delay=0")
457
+ ce, anchor, _ = loss_shallow(model, tok, depth, tgt, pv, off=off, anchor_mode=anchor_mode)
458
+ return ce, anchor
459
+
460
+
461
+ def loss_consist(model, tok, tfire, ent_ids, ent_lo=None, ent_hi=None, *, mode="full"):
462
+ """Drop-in original-signature wrapper -> scalar; entity bounds are unused.
463
+
464
+ Existing code can retain all five arguments and supply ``mode=...``.
465
+ """
466
+ loss, _ = loss_consistency(model, tok, tfire, ent_ids, mode=mode)
467
+ return loss
for Minegishi/reproducibility/probes.py ADDED
@@ -0,0 +1,371 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Exploratory, read-only probes of final loop checkpoints; no SNAP or state reset.
2
+
3
+ Canonical-interface probes assume the entity/token-embedding interface used by the
4
+ anchor experiments. They are NOT architecture-neutral capability measurements.
5
+ Ordinary, unmodified free rollout is reported separately.
6
+ """
7
+ import argparse
8
+ import collections
9
+ import hashlib
10
+ import json
11
+ import os
12
+ from pathlib import Path
13
+ import time
14
+
15
+ import numpy as np
16
+ import torch
17
+
18
+ from model.loop_gpt import LoopGPT
19
+ from train_halt import EvalSet, eval_fixed
20
+
21
+
22
+ def read_json(path):
23
+ with open(path) as stream:
24
+ return json.load(stream)
25
+
26
+
27
+ def named_seed(seed, name):
28
+ return (int(seed) + int(hashlib.sha256(name.encode()).hexdigest()[:12], 16)) % (2**63 - 1)
29
+
30
+
31
+ def summary(values):
32
+ values = np.asarray(values, dtype=np.float64).reshape(-1)
33
+ good = np.isfinite(values)
34
+ finite = values[good]
35
+ return dict(n=int(values.size), nonfinite=int((~good).sum()),
36
+ mean=float(finite.mean()) if finite.size else None,
37
+ p50=float(np.median(finite)) if finite.size else None,
38
+ p95=float(np.percentile(finite, 95)) if finite.size else None,
39
+ maximum=float(finite.max()) if finite.size else None)
40
+
41
+
42
+ class Accumulator:
43
+ def __init__(self):
44
+ self.values = collections.defaultdict(list)
45
+
46
+ def add(self, name, values):
47
+ if torch.is_tensor(values):
48
+ values = values.detach().float().cpu().numpy()
49
+ self.values[name].append(np.asarray(values).reshape(-1))
50
+
51
+ def result(self):
52
+ return {k: summary(np.concatenate(v)) for k, v in self.values.items()}
53
+
54
+
55
+ def bounded_error(a, b):
56
+ # Identical to loss_consist on nonzero vectors; only protect a degenerate 0/0.
57
+ numerator = (a - b).square().sum(-1)
58
+ denominator = a.square().sum(-1) + b.square().sum(-1)
59
+ return numerator / denominator.clamp_min(torch.finfo(a.dtype).tiny)
60
+
61
+
62
+ def save_json(path, result):
63
+ path = Path(path)
64
+ path.parent.mkdir(parents=True, exist_ok=True)
65
+ temp = path.with_name(path.name + ".tmp")
66
+ with temp.open("w") as stream:
67
+ json.dump(result, stream, indent=2, allow_nan=False)
68
+ os.replace(temp, path)
69
+
70
+
71
+ def load_model(run_dir, vocab, meta, device):
72
+ manifest = read_json(run_dir / "manifest.json")
73
+ cfg = manifest["model"]
74
+ token_ids = {token: i for i, token in enumerate(vocab)}
75
+ e0, r0 = token_ids["<e_0>"], token_ids["<r_0>"]
76
+ ne, nr = int(meta["num_entities"]), int(meta["num_relations"])
77
+ if any(token_ids[f"<e_{i}>"] != e0 + i for i in range(ne)):
78
+ raise ValueError("Entity token IDs must be contiguous")
79
+ if any(token_ids[f"<r_{i}>"] != r0 + i for i in range(nr)):
80
+ raise ValueError("Relation token IDs must be contiguous")
81
+ if cfg.get("untied") or cfg.get("reembed", 0):
82
+ raise ValueError("Canonical probes currently support plain tied loops only; "
83
+ "do not silently substitute a different state transition")
84
+ model = LoopGPT(len(vocab), cfg["d_model"], cfg["n_head"], cfg["n_layer"],
85
+ pos=cfg.get("position", "nope"),
86
+ rope_base=cfg.get("rope_base", 100.0), max_pos=cfg.get("max_pos", 512),
87
+ attn_window=cfg.get("attn_window", 0),
88
+ ent_range=(e0, e0 + ne), rel_range=(r0, r0 + nr))
89
+ # Local training checkpoints include optimizer and RNG state. They are trusted
90
+ # local artifacts, not arbitrary downloaded pickle files.
91
+ checkpoint = torch.load(run_dir / "last.pt", map_location="cpu", weights_only=False)
92
+ model.load_state_dict(checkpoint["model"], strict=True)
93
+ model.to(device).eval()
94
+ model.requires_grad_(False)
95
+ return model, manifest, checkpoint.get("update"), e0, r0
96
+
97
+
98
+ def canonical_batch(depth, n, seed, ne, nr, e0, r0, rel_map, device):
99
+ rng = np.random.default_rng(named_seed(seed, f"canonical-d{depth}"))
100
+ cuts = sorted({0, 1, 2, depth // 4, depth // 2, depth - 2, depth - 1})
101
+ cuts = [k for k in cuts if 0 <= k < depth]
102
+ tfire = np.resize(np.asarray(cuts, dtype=np.int64), n)
103
+ rng.shuffle(tfire)
104
+ entities = rng.integers(0, ne, size=(n, depth + 1))
105
+ relations = rng.integers(0, nr, size=(n, depth + 1))
106
+ positions = np.arange(depth + 1)[None, :]
107
+ ids = np.where(positions <= tfire[:, None], e0 + entities, r0 + relations)
108
+ boundary = tfire + 1
109
+ predecessor = entities[np.arange(n), tfire]
110
+ relation = relations[np.arange(n), boundary]
111
+ gold = e0 + rel_map[relation, predecessor]
112
+ return (torch.tensor(ids, device=device), torch.tensor(tfire, device=device),
113
+ torch.tensor(gold, device=device), cuts)
114
+
115
+
116
+ def masks_for(tfire, length):
117
+ positions = torch.arange(length, device=tfire.device)[None, :]
118
+ return dict(root=(positions == 0).expand(len(tfire), -1),
119
+ hold=(positions >= 1) & (positions <= tfire[:, None]),
120
+ wait=positions > tfire[:, None] + 1,
121
+ frontier=positions == tfire[:, None] + 1)
122
+
123
+
124
+ def add_canonical(acc, source, output, tfire, gold, model, e0, ne):
125
+ masks = masks_for(tfire, source.shape[1])
126
+ bounded = bounded_error(output, source)
127
+ relative = (output - source).norm(dim=-1) / source.norm(dim=-1).clamp_min(1e-12)
128
+ rows = torch.arange(len(tfire), device=source.device)
129
+ logits = model.logits(model.ln_f(output[rows, tfire + 1]))
130
+ acc.add("frontier_accuracy_full_vocab", logits.argmax(-1) == gold)
131
+ acc.add("frontier_accuracy_entity_only", logits[:, e0:e0 + ne].argmax(-1) + e0 == gold)
132
+ acc.add("keep_bounded", bounded[~masks["frontier"]])
133
+ for region in ("root", "hold", "wait"):
134
+ acc.add(f"{region}_bounded", bounded[masks[region]])
135
+ acc.add(f"{region}_relative_l2", relative[masks[region]])
136
+
137
+
138
+ def finalize_canonical(acc):
139
+ result = acc.result()
140
+ count = result["keep_bounded"]["n"]
141
+ result["bounded_contributions_full_keep_denominator"] = {
142
+ region: (result[f"{region}_bounded"]["mean"] or 0.0)
143
+ * result[f"{region}_bounded"]["n"] / max(count, 1)
144
+ for region in ("root", "hold", "wait")}
145
+ return result
146
+
147
+
148
+ def add_perturbation(acc, baseline, perturbed, initial_delta_norm, masks):
149
+ difference = perturbed - baseline
150
+ delta_norm = difference.flatten(1).norm(dim=-1)
151
+ acc.add("relative_frobenius_per_example",
152
+ delta_norm / baseline.flatten(1).norm(dim=-1).clamp_min(1e-12))
153
+ valid = initial_delta_norm > 0
154
+ acc.add("gain_vs_actual_initial_delta_per_example", delta_norm[valid] / initial_delta_norm[valid])
155
+ paired = bounded_error(perturbed, baseline)
156
+ for region, mask in masks.items():
157
+ acc.add(f"initial_{region}_paired_bounded", paired[mask])
158
+
159
+
160
+ @torch.inference_mode()
161
+ def canonical_and_perturbation(model, args, meta, e0, r0, device, progress):
162
+ ne, nr = int(meta["num_entities"]), int(meta["num_relations"])
163
+ rel_map = np.asarray(meta["rel_map"], dtype=np.int64)
164
+ if rel_map.shape != (nr, ne):
165
+ raise ValueError(f"Unexpected relation map shape {rel_map.shape}")
166
+ all_results = {}
167
+ for depth in args.canonical_depths:
168
+ ids, cuts, gold, cut_values = canonical_batch(
169
+ depth, args.n, args.seed, ne, nr, e0, r0, rel_map, device)
170
+ overall = Accumulator()
171
+ per_cut = {k: Accumulator() for k in cut_values}
172
+ perturb = collections.defaultdict(Accumulator)
173
+ for start in range(0, args.n, args.batch_size):
174
+ tok, tfire, targets = ids[start:start + args.batch_size], cuts[start:start + args.batch_size], gold[start:start + args.batch_size]
175
+ source = model.embed(tok).detach()
176
+ masks = masks_for(tfire, source.shape[1])
177
+ baseline = source
178
+ branches = []
179
+ for region in args.perturb_regions:
180
+ generator = torch.Generator(device=device).manual_seed(named_seed(
181
+ args.seed, f"noise-d{depth}-batch{start}-{region}"))
182
+ noise = torch.randn(source.shape, device=device, dtype=source.dtype, generator=generator)
183
+ noise = noise / noise.norm(dim=-1, keepdim=True).clamp_min(1e-12)
184
+ noise = noise * source.norm(dim=-1, keepdim=True) * masks[region][..., None]
185
+ for epsilon in args.epsilons:
186
+ h = source + epsilon * noise
187
+ delta = (h - source).flatten(1).norm(dim=-1)
188
+ branches.append(dict(region=region, epsilon=epsilon, h=h, delta=delta))
189
+ for step in range(1, max(args.perturb_steps) + 1):
190
+ baseline = model.step(baseline)
191
+ if step == 1:
192
+ add_canonical(overall, source, baseline, tfire, targets, model, e0, ne)
193
+ for cut in cut_values:
194
+ chosen = tfire == cut
195
+ if chosen.any():
196
+ add_canonical(per_cut[cut], source[chosen], baseline[chosen], tfire[chosen], targets[chosen], model, e0, ne)
197
+ for branch in branches:
198
+ # The epsilon=0 reference is reused, not independently rerun.
199
+ branch["h"] = baseline if branch["epsilon"] == 0 else model.step(branch["h"])
200
+ if step in args.perturb_steps:
201
+ key = f"{branch['region']}|eps={branch['epsilon']:g}|step={step}"
202
+ add_perturbation(perturb[key], baseline, branch["h"], branch["delta"], masks)
203
+ all_results[str(depth)] = dict(
204
+ n=args.n, tfire_counts={str(k): int((cuts == k).sum().item()) for k in cut_values},
205
+ canonical=finalize_canonical(overall),
206
+ canonical_by_tfire={str(k): finalize_canonical(a) for k, a in per_cut.items()
207
+ if a.values},
208
+ perturbation={key: acc.result() for key, acc in perturb.items()})
209
+ progress(all_results)
210
+ print(f"canonical/perturbation d{depth}, n={args.n}: "
211
+ f"frontier={all_results[str(depth)]['canonical']['frontier_accuracy_full_vocab']['mean']:.4f}", flush=True)
212
+ return all_results
213
+
214
+
215
+ class SmallEvalSet(EvalSet):
216
+ def __init__(self, items, e0, r0, device, batch_size):
217
+ super().__init__(items, e0, r0, device)
218
+ self.batch_size = batch_size
219
+
220
+ def batches(self, tokens_per_batch=60000):
221
+ for depth, indices, tokens, gold in self.groups:
222
+ for start in range(0, len(indices), self.batch_size):
223
+ stop = start + self.batch_size
224
+ yield depth, indices[start:stop], tokens[start:stop], gold[start:stop]
225
+
226
+
227
+ @torch.inference_mode()
228
+ def free_rollout(model, args, e0, r0, device, progress):
229
+ tests = read_json(Path(args.data_dir) / "tests.json")
230
+ cells = collections.defaultdict(list)
231
+ for item in tests:
232
+ if item["d"] in args.rollout_depths:
233
+ key = (item["d"], item["cat"])
234
+ if len(cells[key]) < args.n:
235
+ cells[key].append(dict(item, gold=item["states"][-1]))
236
+ results = {}
237
+ for (depth, category), items in sorted(cells.items()):
238
+ evaluation = SmallEvalSet(items, e0, r0, device, args.batch_size)
239
+ # All round readouts are obtained in ONE unmodified rollout, using the
240
+ # trainer's exact state API/readout. No oracle information enters F.
241
+ correct = eval_fixed(model, evaluation, list(range(1, 2 * depth + 1)))
242
+ trajectory = np.stack([correct[t] for t in range(1, 2 * depth + 1)])
243
+ ever = trajectory.any(axis=0)
244
+ first = np.where(ever, trajectory.argmax(axis=0) + 1, 0)
245
+ at_d = correct[depth]
246
+ denominator = int(at_d.sum())
247
+ after_first = np.arange(1, 2 * depth + 1)[:, None] >= first[None, :]
248
+ forgot = ((~trajectory) & after_first).any(axis=0) & ever
249
+ key = f"{category}|d{depth}"
250
+ results[key] = dict(
251
+ n=len(items), test_ids=[int(item["id"]) for item in items],
252
+ accuracy_R_d=float(at_d.mean()),
253
+ accuracy_R_d_plus_2=float(correct[depth + 2].mean()),
254
+ accuracy_R_2d=float(correct[2 * depth].mean()),
255
+ retention_denominator_correct_at_R_d=denominator,
256
+ retention_d_plus_2_given_correct_d=float(correct[depth + 2][at_d].mean()) if denominator else None,
257
+ retention_2d_given_correct_d=float(correct[2 * depth][at_d].mean()) if denominator else None,
258
+ uninterrupted_retention_d_to_2d_given_correct_d=float(trajectory[depth - 1:, at_d].all(axis=0).mean()) if denominator else None,
259
+ first_correct_round_retrospective=summary(first[ever]),
260
+ never_correct_fraction=float((~ever).mean()),
261
+ later_incorrect_given_ever_correct=float(forgot[ever].mean()) if ever.any() else None,
262
+ accuracy_by_round=[float(value) for value in trajectory.mean(axis=1)],
263
+ per_example_first_correct_round_zero_if_never=first.tolist())
264
+ progress(results)
265
+ print(f"free rollout {key}, n={len(items)}: "
266
+ f"R=d {at_d.mean():.4f}, R=2d {correct[2 * depth].mean():.4f}", flush=True)
267
+ return results
268
+
269
+
270
+ def monitor_crosscheck(run_dir, update, results):
271
+ path = run_dir / "metrics.jsonl"
272
+ if not path.exists():
273
+ return dict(status="no_metrics_jsonl")
274
+ matching = None
275
+ with path.open() as stream:
276
+ for line in stream:
277
+ try:
278
+ record = json.loads(line)
279
+ except json.JSONDecodeError:
280
+ continue
281
+ if record.get("update", record.get("step")) == update and record.get("monitor_budgets"):
282
+ matching = record
283
+ if matching is None:
284
+ return dict(status="no_same_checkpoint_monitor_budgets")
285
+ comparisons = {}
286
+ for key, cell in results.items():
287
+ if cell["n"] != 100:
288
+ continue
289
+ for budget, metric in (("d", "accuracy_R_d"), ("d+2", "accuracy_R_d_plus_2"), ("2d", "accuracy_R_2d")):
290
+ recorded = matching["monitor_budgets"].get(budget, {}).get(key)
291
+ if recorded is not None:
292
+ comparisons[f"{key}|{budget}"] = dict(probe=cell[metric], monitor=recorded,
293
+ difference=cell[metric] - recorded)
294
+ return dict(status="compared_where_available", note="Assumes trainer monitor_per_cell=100; compare test IDs/config if values differ.", cells=comparisons)
295
+
296
+
297
+ def parse_args():
298
+ parser = argparse.ArgumentParser(description=__doc__)
299
+ parser.add_argument("--run-dir", required=True)
300
+ parser.add_argument("--data-dir", required=True)
301
+ parser.add_argument("--out", required=True)
302
+ parser.add_argument("--device", default="cuda")
303
+ parser.add_argument("--seed", type=int, default=20260923, help="Fixed probe seed, shared across model seeds")
304
+ parser.add_argument("--n", type=int, default=100, help="Total canonical samples per depth; free samples per depth/category")
305
+ parser.add_argument("--batch-size", type=int, default=25)
306
+ parser.add_argument("--canonical-depths", type=int, nargs="+", default=[32, 128, 256])
307
+ parser.add_argument("--epsilons", type=float, nargs="+", default=[0.0, 1e-3, 1e-2])
308
+ parser.add_argument("--perturb-steps", type=int, nargs="+", default=[1, 2, 4, 8])
309
+ parser.add_argument("--perturb-regions", choices=["wait", "hold", "root"], nargs="+", default=["wait", "hold"])
310
+ parser.add_argument("--rollout-depths", type=int, nargs="+", default=[16, 32, 64, 128])
311
+ parser.add_argument("--skip-rollout", action="store_true")
312
+ args = parser.parse_args()
313
+ if args.n < 1 or args.batch_size < 1 or min(args.canonical_depths + args.rollout_depths) < 2:
314
+ parser.error("n/batch-size must be positive and depths must be >=2")
315
+ if min(args.epsilons) < 0 or min(args.perturb_steps) < 1:
316
+ parser.error("epsilons must be nonnegative and perturb steps must be positive")
317
+ return args
318
+
319
+
320
+ def main():
321
+ args = parse_args()
322
+ start = time.monotonic()
323
+ run_dir, data_dir = Path(args.run_dir), Path(args.data_dir)
324
+ meta, vocab = read_json(data_dir / "meta.json"), read_json(data_dir / "vocab.json")
325
+ device = torch.device(args.device)
326
+ # Match the trainer's strict fp32 evaluation even if its training used TF32.
327
+ torch.backends.cuda.matmul.allow_tf32 = False
328
+ torch.backends.cudnn.allow_tf32 = False
329
+ model, manifest, update, e0, r0 = load_model(run_dir, vocab, meta, device)
330
+ maximum = max(args.canonical_depths + ([] if args.skip_rollout else args.rollout_depths))
331
+ if model.pos == "rope" and maximum + 1 > model.rope_cos.shape[0]:
332
+ raise ValueError("Requested sequence length exceeds the checkpoint's RoPE capacity")
333
+ result = dict(
334
+ exploratory_probe=True, status="running", run_dir=str(run_dir.resolve()),
335
+ checkpoint="last.pt", checkpoint_update=update, training_seed=manifest.get("seed"),
336
+ probe_seed=args.seed, config=vars(args), torch_version=str(torch.__version__),
337
+ dtype=str(next(model.parameters()).dtype), device=str(device),
338
+ evaluation_precision="strict_fp32_no_autocast_tf32_disabled",
339
+ training_condition={key: manifest.get(key) for key in ("anchor", "consist_loss", "offset_max", "model")},
340
+ probe_source_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
341
+ meta_sha256=hashlib.sha256((data_dir / "meta.json").read_bytes()).hexdigest(),
342
+ interface_limitations=[
343
+ "Canonical states assume the anchor's token/entity embedding interface. A low canonical score does not establish low capability for a model trained without that interface.",
344
+ "Independent prefix entities form synthetic local states, not necessarily reachable trajectories of real deep inputs.",
345
+ "tfire specifies constructed inputs and reporting masks only. It is never passed to the model, and free rollout has no SNAP, reset, or intermediate answer injection.",
346
+ "Perturbation compares against the same unperturbed evolving batch, not the static input. Region names denote INITIAL positions; some waiting positions naturally execute later.",
347
+ "epsilon is relative L2 norm PER PERTURBED TOKEN. Gain excludes examples with an empty perturbation region; epsilon=0 reuses the reference and has undefined gain.",
348
+ "First-correct round is a retrospective diagnostic using the answer, not a realizable halting rule or reported task success criterion.",
349
+ "Small n, selected depths, and canonical probes are exploratory, not a proof of arbitrary-depth generalization."],
350
+ bounded_error_formula="sum((h-S)^2)/(sum(h^2)+sum(S^2)); full contributions use all nonfrontier tokens as denominator")
351
+ def progress_canonical(values):
352
+ result["canonical_and_perturbation"] = values
353
+ result["elapsed_seconds"] = time.monotonic() - start
354
+ save_json(args.out, result)
355
+ def progress_rollout(values):
356
+ result["free_rollout"] = values
357
+ result["elapsed_seconds"] = time.monotonic() - start
358
+ save_json(args.out, result)
359
+ save_json(args.out, result)
360
+ canonical_and_perturbation(model, args, meta, e0, r0, device, progress_canonical)
361
+ if not args.skip_rollout:
362
+ rollout = free_rollout(model, args, e0, r0, device, progress_rollout)
363
+ result["same_checkpoint_monitor_crosscheck"] = monitor_crosscheck(run_dir, update, rollout)
364
+ result["status"] = "complete"
365
+ result["elapsed_seconds"] = time.monotonic() - start
366
+ save_json(args.out, result)
367
+ print(f"probes complete in {result['elapsed_seconds']:.1f}s: {args.out}", flush=True)
368
+
369
+
370
+ if __name__ == "__main__":
371
+ main()
for Minegishi/reproducibility/recipe.py ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """One source of truth for the two authorized RoPE=10000 training commands."""
2
+ import sys
3
+ from pathlib import Path
4
+
5
+ def command(code, cfg, run, gate=False):
6
+ code=Path(code)
7
+ assert cfg['seed']==7 and cfg['updates']==200000 and cfg['anchor_mode']=='full'
8
+ assert cfg['rope_base']==10000 and cfg['max_pos']==4097
9
+ assert cfg['cs_mode'] in ('wait_no_root','t0_wait_denoise')
10
+ assert cfg['cs_noise_scale']==(0.2 if cfg['cs_mode']=='t0_wait_denoise' else 0)
11
+ assert cfg['cs_noise_scope']=='wait' and cfg['cs_noise_kind']==('isotropic' if cfg['cs_mode']=='t0_wait_denoise' else 'embedding')
12
+ return [sys.executable,'-u',str(code/'train_halt.py'),'--data_dir',str(code/'data/chain_loop_k8'),
13
+ '--atomic',str(code/'data/atomic_joint_2026-09-19/train_atomic.json'),'--save_dir',str(run),
14
+ '--arm','D','--seed','7','--updates','200000','--lr_decay','linear',
15
+ '--pos','rope','--rope_base','10000','--max_pos','4097','--anchor','1','--anchor_mode','full',
16
+ '--consist_w','1','--cs_mode',cfg['cs_mode'],'--consist_depth','32','--consist_rows','32',
17
+ '--consist_start','1' if gate else '20000','--offset_max','128','--offset_fill','entity',
18
+ '--cs_noise_scale',str(cfg['cs_noise_scale']),'--cs_noise_kind',cfg['cs_noise_kind'],'--cs_noise_scope','wait',
19
+ '--train_precision','tf32','--cuda_graph','1','--monitor_per_cell','1' if gate else '100',
20
+ '--eval_every','10000','--ckpt_every','100' if gate else '5000','--final_n','5000','--diagnostic_n','100']
for Minegishi/reproducibility/streams.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Unmodified Stream/load_json extracted from the downloaded Miyabi train_chain.py."""
2
+ import hashlib,json,os
3
+ import numpy as np
4
+
5
+ def load_json(d, name):
6
+ return json.load(open(os.path.join(d, f"{name}.json")))
7
+
8
+ class Stream:
9
+ def __init__(self, n, seed, name):
10
+ self.n = n; self.rng = np.random.default_rng([seed, int(hashlib.sha256(name.encode()).hexdigest()[:8], 16)])
11
+ self.order = None; self.cursor = 0; self.cycles = 0
12
+
13
+ def take(self, k):
14
+ out = []
15
+ while k > 0:
16
+ if self.order is None or self.cursor >= self.n:
17
+ self.order = self.rng.permutation(self.n); self.cursor = 0; self.cycles += 1
18
+ chunk = self.order[self.cursor:self.cursor + k]; self.cursor += len(chunk); k -= len(chunk); out.append(chunk)
19
+ return np.concatenate(out)
20
+
21
+ def state(self):
22
+ return dict(order=None if self.order is None else self.order.tolist(), cursor=self.cursor, cycles=self.cycles,
23
+ rng=self.rng.bit_generator.state)
24
+
25
+ def load(self, st):
26
+ self.order = None if st["order"] is None else np.array(st["order"]); self.cursor = st["cursor"]; self.cycles = st["cycles"]
27
+ self.rng.bit_generator.state = st["rng"]
for Minegishi/reproducibility/train_chain.py ADDED
@@ -0,0 +1,435 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Training + autonomous evaluation for the strict w2/d2 chain task (docs/shallow_training_deep_composition_plan_2026-09-16.md).
3
+
4
+ python train_chain.py --data_dir ../data/chain_L --save_dir ../runs/chain_B_s1 --protocol B --seed 1 [--resume]
5
+
6
+ Schedule (plan §4.2): phase 1 = `--phase1_updates` updates of 256 w2 rows; phase 2 = `--phase2_updates` updates of
7
+ 128 w2 + 128 d2 rows. Loss per update = sum over rows of sum_t w_t CE_t / 256 with the protocol weights of
8
+ data/chain_format.py (final 1, bridge 1, END 1, copies share 1; w2 sub-questions halved). Streams are frozen per
9
+ (seed, stream name) so every protocol sees the same question ids at the same update.
10
+
11
+ Every `--eval_every` updates: teacher-forced CE per token role (train subsets, validation d2), FREE greedy rollouts
12
+ (no gold history) on the exhaustive atomic set, validation d2 (model selection), validation w2 (second question
13
+ conditioned on the model's own first answer) and a small test subset; metrics.jsonl + model checkpoint
14
+ ckpt_UUUUUU.pt; last.pt holds the full state (optimizer, streams, RNG) for exact resume; best_by_val.pt tracks
15
+ validation d2 final-answer accuracy (ties: lower validation loss). At the end (or with --eval_only) the full test
16
+ matrix is rolled out for the final and best checkpoints -> final_eval_<tag>.json, predictions_<tag>.jsonl.
17
+ """
18
+ import argparse
19
+ import collections
20
+ import hashlib
21
+ import json
22
+ import os
23
+ import time
24
+
25
+ import numpy as np
26
+ import torch
27
+ import torch.nn.functional as F
28
+
29
+ from data.chain_format import (ROLE_NAMES, ChainVocab, completion_length, parse_completion, prompt_ids, render_chain,
30
+ render_d2, render_w2, score_chain)
31
+ from data.local_format import extend_vocab
32
+ from model import GPT2LikeEncoder
33
+
34
+
35
+ # ----------------------------------------------------------------------------------------------------------------- data
36
+ def load_json(d, name):
37
+ return json.load(open(os.path.join(d, f"{name}.json")))
38
+
39
+
40
+ class Rendered:
41
+ """Padded tensors for rendered rows: inputs = ids[:-1], targets = ids[1:], weights/roles aligned with targets."""
42
+
43
+ def __init__(self, rows, pad, L, device):
44
+ n = len(rows)
45
+ inp = torch.full((n, L), pad, dtype=torch.long); tgt = torch.full((n, L), -100, dtype=torch.long)
46
+ w = torch.zeros((n, L)); roles = torch.full((n, L), -1, dtype=torch.long); pm = torch.ones((n, L), dtype=torch.bool)
47
+ for i, (ids, ww, rr) in enumerate(rows):
48
+ k = len(ids) - 1
49
+ inp[i, :k] = torch.tensor(ids[:-1]); tgt[i, :k] = torch.tensor(ids[1:])
50
+ w[i, :k] = torch.tensor(ww[1:]); roles[i, :k] = torch.tensor(rr[1:]); pm[i, :k] = False
51
+ self.input_ids, self.target_ids, self.weights, self.roles, self.pad_mask = (t.to(device) for t in (inp, tgt, w, roles, pm))
52
+ self.n = n
53
+ self.n_tokens = int((~pm).sum()); self.n_supervised = int((w > 0).sum())
54
+
55
+
56
+ class Stream:
57
+ def __init__(self, n, seed, name):
58
+ self.n = n; self.rng = np.random.default_rng([seed, int(hashlib.sha256(name.encode()).hexdigest()[:8], 16)])
59
+ self.order = None; self.cursor = 0; self.cycles = 0
60
+
61
+ def take(self, k):
62
+ out = []
63
+ while k > 0:
64
+ if self.order is None or self.cursor >= self.n:
65
+ self.order = self.rng.permutation(self.n); self.cursor = 0; self.cycles += 1
66
+ chunk = self.order[self.cursor:self.cursor + k]; self.cursor += len(chunk); k -= len(chunk); out.append(chunk)
67
+ return np.concatenate(out)
68
+
69
+ def state(self):
70
+ return dict(order=None if self.order is None else self.order.tolist(), cursor=self.cursor, cycles=self.cycles,
71
+ rng=self.rng.bit_generator.state)
72
+
73
+ def load(self, st):
74
+ self.order = None if st["order"] is None else np.array(st["order"]); self.cursor = st["cursor"]; self.cycles = st["cycles"]
75
+ self.rng.bit_generator.state = st["rng"]
76
+
77
+
78
+ # --------------------------------------------------------------------------------------------------------------- rollout
79
+ ROLLOUT_BATCH = 1000
80
+
81
+
82
+ @torch.no_grad()
83
+ def rollout(model, prompts, max_new, device, batch=None, n_loops=1):
84
+ """Greedy free generation for prompts of EQUAL length. Returns (N, max_new) generated ids.
85
+ n_loops: recurrent-depth budget (the same for every generated token; docs/experiments_loop_next.md §6)."""
86
+ batch = batch or ROLLOUT_BATCH
87
+ outs = []
88
+ for s in range(0, len(prompts), batch):
89
+ x = torch.tensor(prompts[s:s + batch], dtype=torch.long, device=device)
90
+ for _ in range(max_new):
91
+ logits = model(x, n_loops=n_loops)[:, -1, :]
92
+ x = torch.cat([x, logits.argmax(-1, keepdim=True)], 1)
93
+ outs.append(x[:, -max_new:].cpu())
94
+ return torch.cat(outs).numpy()
95
+
96
+
97
+ def eval_chains(model, v, protocol, rel, chains, device, keep_predictions=False, loops=None):
98
+ """chains: list of dict(x, rels, states, ...). Groups by depth. Returns per-chain scores (and parsed outputs).
99
+ loops: None (plain model) or a callable depth -> recurrent-depth budget T."""
100
+ results = []
101
+ by_d = collections.defaultdict(list)
102
+ for c in chains:
103
+ by_d[len(c["rels"])].append(c)
104
+ for d, cs in by_d.items():
105
+ prompts = [prompt_ids(v, c["x"], c["rels"]) for c in cs]
106
+ gen = rollout(model, prompts, completion_length(protocol, d) + 2, device, n_loops=loops(d) if loops else 1)
107
+ for c, g in zip(cs, gen):
108
+ parsed = parse_completion(v, protocol, g.tolist(), d)
109
+ sc = score_chain(v, protocol, rel, c["x"], c["rels"], c["states"], parsed)
110
+ rec = dict(sc)
111
+ if keep_predictions:
112
+ rec.update(id=c.get("id"), cat=c.get("cat"), d=d, pred_states=parsed["states"], pred_rels=parsed["rels"],
113
+ ended=parsed["ended"], valid_format=parsed["valid_format"])
114
+ results.append((c, rec))
115
+ return results
116
+
117
+
118
+ def summarize(results):
119
+ n = len(results)
120
+ if n == 0:
121
+ return {}
122
+ fe = collections.Counter(r["first_error_type"] for _, r in results if r["first_error_type"])
123
+ stop = collections.Counter(r["stop"] for _, r in results)
124
+ own_c = sum(r["own_update_correct"] for _, r in results); own_t = sum(r["own_update_total"] for _, r in results)
125
+ return dict(n=n, final_acc=sum(r["final_correct"] for _, r in results) / n, traj_acc=sum(r["traj_correct"] for _, r in results) / n,
126
+ own_update_acc=(own_c / own_t) if own_t else None, stop={k: v / n for k, v in stop.items()},
127
+ first_error={k: v / n for k, v in fe.items()})
128
+
129
+
130
+ @torch.no_grad()
131
+ def eval_w2_val(model, v, protocol, rel, val_w2, device, n_loops=1):
132
+ """Second question answered with the model's OWN first output as history (plan §5)."""
133
+ correct = 0; both = 0
134
+ for q in val_w2:
135
+ (x, r), (z, s) = q["q1"], q["q2"]
136
+ p1 = prompt_ids(v, x, [r]); g1 = rollout(model, [p1], completion_length(protocol, 1) + 2, device, n_loops=n_loops)[0].tolist()
137
+ end1 = g1.index(v.END) + 1 if v.END in g1 else len(g1)
138
+ hist = p1 + g1[:end1] + prompt_ids(v, z, [s])
139
+ g2 = rollout(model, [hist], completion_length(protocol, 1) + 2, device, n_loops=n_loops)[0].tolist()
140
+ s1 = score_chain(v, protocol, rel, x, [r], [rel[r][x]], parse_completion(v, protocol, g1, 1))
141
+ s2 = score_chain(v, protocol, rel, z, [s], [rel[s][z]], parse_completion(v, protocol, g2, 1))
142
+ correct += int(s2["final_correct"]); both += int(s1["final_correct"] and s2["final_correct"])
143
+ return dict(n=len(val_w2), second_acc=correct / len(val_w2), both_acc=both / len(val_w2))
144
+
145
+
146
+ @torch.no_grad()
147
+ def teacher_forced(model, R: Rendered, device, batch=1024, n_loops=1):
148
+ """Mean CE per role and the weighted loss per row."""
149
+ sums = collections.defaultdict(float); cnts = collections.defaultdict(int); wloss = 0.0
150
+ for s in range(0, R.n, batch):
151
+ sl = slice(s, s + batch)
152
+ logits = model(R.input_ids[sl], pad_mask=R.pad_mask[sl], n_loops=n_loops)
153
+ ce = F.cross_entropy(logits.transpose(1, 2), R.target_ids[sl], reduction="none", ignore_index=-100)
154
+ wloss += float((ce * R.weights[sl]).sum())
155
+ for k, name in enumerate(ROLE_NAMES):
156
+ m = R.roles[sl] == k
157
+ if k == 0:
158
+ continue
159
+ sums[name] += float(ce[m].sum()); cnts[name] += int(m.sum())
160
+ out = {f"ce/{k}": sums[k] / cnts[k] for k in sums if cnts[k]}
161
+ out["loss_per_row"] = wloss / R.n
162
+ return out
163
+
164
+
165
+ # ------------------------------------------------------------------------------------------------------------------ main
166
+ def main():
167
+ ap = argparse.ArgumentParser()
168
+ ap.add_argument("--data_dir", required=True); ap.add_argument("--save_dir", required=True)
169
+ ap.add_argument("--protocol", choices=["A", "B", "C", "D"], required=True)
170
+ ap.add_argument("--no_rope", action="store_true", help="NoPE control: no positional encoding at all (plan §6.1)")
171
+ ap.add_argument("--seed", type=int, default=1)
172
+ ap.add_argument("--d_model", type=int, default=256); ap.add_argument("--n_layer", type=int, default=2); ap.add_argument("--n_head", type=int, default=2)
173
+ ap.add_argument("--max_len", type=int, default=512); ap.add_argument("--rope_base", type=float, default=100.0)
174
+ ap.add_argument("--lr", type=float, default=3e-4); ap.add_argument("--weight_decay", type=float, default=0.1); ap.add_argument("--warmup", type=int, default=100)
175
+ ap.add_argument("--rows_per_update", type=int, default=256)
176
+ ap.add_argument("--d3_per_update", type=int, default=0, help="CONTROL: rows from train_d3.json per phase-2 update (taken out of the d2 share)")
177
+ ap.add_argument("--phase1_updates", type=int, default=50000); ap.add_argument("--phase2_updates", type=int, default=200000)
178
+ ap.add_argument("--eval_every", type=int, default=5000); ap.add_argument("--save_every", type=int, default=5000)
179
+ ap.add_argument("--small_test_per_cell", type=int, default=500); ap.add_argument("--small_test_depths", default="2,3,4,8")
180
+ ap.add_argument("--resume", action="store_true")
181
+ ap.add_argument("--eval_only", default=None, help="checkpoint to evaluate on the full test matrix (no training)")
182
+ ap.add_argument("--no_final_eval", action="store_true")
183
+ ap.add_argument("--rollout_batch", type=int, default=1000, help="chains per rollout batch (protocol D at d=32 needs ~250 on a shared GPU)")
184
+ ap.add_argument("--final_max_per_cell", type=int, default=0, help="limit chains per (cat, d) cell in the final eval (0 = all; smoke tests)")
185
+ ap.add_argument("--final_depths", default="", help="comma-separated depths for the final eval (default: all in the dataset)")
186
+ ap.add_argument("--final_extra_ckpts", default="", help="'auto' = also evaluate the two checkpoints before the last; or a comma-separated list")
187
+ ap.add_argument("--eval_tag", default=None, help="tag for --eval_only outputs (default 'evalonly')")
188
+ # recurrent-depth arm R (docs/experiments_loop_next.md §6): the same block stack is applied T times
189
+ ap.add_argument("--loop_T", default="", help="training loop budgets sampled per update, e.g. 2,4,8 (empty = plain model)")
190
+ ap.add_argument("--loop_cap", type=int, default=128, help="test rule T(d) = smallest power of two >= max(8, d), capped here")
191
+ ap.add_argument("--loop_fixed", type=int, default=8, help="fixed-budget control reported next to T(d) in the final eval")
192
+ ap.add_argument("--loop_matrix_at", default="50000,150000,250000", help="updates at which the full depth x T matrix is computed on the monitoring set")
193
+ ap.add_argument("--loop_matrix_T", default="1,2,4,8,16,32,64,128")
194
+ ap.add_argument("--zero_init_out", action="store_true", help="zero-init the attention / MLP output projections (looped model)")
195
+ ap.add_argument("--vocab_ext", action="store_true", help="use the frozen extended vocabulary shared with the local-executor arms")
196
+ ap.add_argument("--select_by", choices=["d2", "w2"], default="d2", help="checkpoint selection: validation d2 (original) or w2-only free execution (new arms)")
197
+ args = ap.parse_args()
198
+ global ROLLOUT_BATCH
199
+ ROLLOUT_BATCH = args.rollout_batch
200
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
201
+ os.makedirs(args.save_dir, exist_ok=True)
202
+ torch.manual_seed(args.seed); np.random.seed(args.seed)
203
+
204
+ meta = load_json(args.data_dir, "meta"); rel = meta["rel_map"]
205
+ base_vocab = load_json(args.data_dir, "vocab")
206
+ v = ChainVocab(extend_vocab(base_vocab) if args.vocab_ext else base_vocab)
207
+ loop_T = [int(x) for x in args.loop_T.split(",")] if args.loop_T else []
208
+ looped = bool(loop_T)
209
+
210
+ def T_of_d(d):
211
+ """pre-registered budget rule (plan §6): smallest power of two >= max(8, d), capped at --loop_cap; 1 for the plain model."""
212
+ if not looped:
213
+ return 1
214
+ t = 8
215
+ while t < d:
216
+ t *= 2
217
+ return min(args.loop_cap, t)
218
+ train_w2 = load_json(args.data_dir, "train_w2"); train_d2 = load_json(args.data_dir, "train_d2")
219
+ val_d2 = load_json(args.data_dir, "val_d2"); val_w2 = load_json(args.data_dir, "val_w2"); tests = load_json(args.data_dir, "tests")
220
+ P = args.protocol
221
+ rw2 = [render_w2(v, P, rel, q) for q in train_w2]; rd2 = [render_d2(v, P, rel, q) for q in train_d2]
222
+ train_d3 = load_json(args.data_dir, "train_d3") if (args.d3_per_update and os.path.exists(os.path.join(args.data_dir, "train_d3.json"))) else []
223
+ if args.d3_per_update and not train_d3:
224
+ raise SystemExit("--d3_per_update > 0 but the dataset has no train_d3.json (generate with --train_d3 1)")
225
+ rd3 = [render_chain(v, P, rel, q) for q in train_d3]
226
+ L = max([max(len(r[0]) for r in rw2), max(len(r[0]) for r in rd2)] + ([max(len(r[0]) for r in rd3)] if rd3 else [])) - 1
227
+ W2 = Rendered(rw2, v.pad, L, device); D2 = Rendered(rd2, v.pad, L, device)
228
+ D3 = Rendered(rd3, v.pad, L, device) if rd3 else None
229
+ VAL = Rendered([render_d2(v, P, rel, q) for q in val_d2], v.pad, L, device)
230
+ VALW2 = Rendered([render_w2(v, P, rel, q) for q in val_w2], v.pad, L, device)
231
+ SUB_W2 = Rendered(rw2[:1000], v.pad, L, device); SUB_D2 = Rendered(rd2[:1000], v.pad, L, device)
232
+ atomic = [dict(x=x, rels=[r], states=[rel[r][x]]) for r in range(v.R) for x in range(v.E)]
233
+ small_depths = {int(d) for d in args.small_test_depths.split(",")}
234
+ small_tests = []
235
+ per_cell = collections.Counter()
236
+ for t in tests:
237
+ key = (t["cat"], t["d"])
238
+ if t["d"] in small_depths and per_cell[key] < args.small_test_per_cell:
239
+ small_tests.append(t); per_cell[key] += 1
240
+ val_chains = [dict(x=q["x"], rels=[q["r"], q["s"]], states=[rel[q["r"]][q["x"]], rel[q["s"]][rel[q["r"]][q["x"]]]]) for q in val_d2]
241
+
242
+ model = GPT2LikeEncoder(len(v.vocab), d_model=args.d_model, n_layer=args.n_layer, n_head=args.n_head, dropout=0.0,
243
+ max_len=args.max_len, rope_base=args.rope_base, use_rope=not args.no_rope, zero_init_out=args.zero_init_out).to(device)
244
+ n_params = sum(p.numel() for p in model.parameters())
245
+ T1 = T_of_d(1)
246
+
247
+ def cells_of(res):
248
+ cells = collections.defaultdict(list)
249
+ for c, r in res:
250
+ cells[f"{c['cat']}|d{c['d']}"].append((c, r))
251
+ return {k: summarize(vv) for k, vv in sorted(cells.items())}
252
+
253
+ def full_eval(tag, ckpt_path=None):
254
+ if ckpt_path:
255
+ model.load_state_dict(torch.load(ckpt_path, map_location=device, weights_only=False)["model"])
256
+ model.eval()
257
+ pool = tests
258
+ final_depths = {int(x) for x in args.final_depths.split(",")} if args.final_depths else None
259
+ if args.final_max_per_cell or final_depths:
260
+ cnt = collections.Counter(); pool = []
261
+ for t in tests:
262
+ if final_depths and t["d"] not in final_depths:
263
+ continue
264
+ if not args.final_max_per_cell or cnt[(t["cat"], t["d"])] < args.final_max_per_cell:
265
+ pool.append(t); cnt[(t["cat"], t["d"])] += 1
266
+ t_eval = time.time()
267
+ res = eval_chains(model, v, P, rel, pool, device, keep_predictions=True, loops=T_of_d)
268
+ cells = collections.defaultdict(list)
269
+ for c, r in res:
270
+ cells[(c["cat"], c["d"])].append((c, r))
271
+ for flag in ("repeated_relation", "revisits_entity"):
272
+ if c.get(flag):
273
+ cells[(f"{c['cat']}|{flag}", c["d"])].append((c, r))
274
+ summary = {f"{k[0]}|d{k[1]}": summarize(vv) for k, vv in sorted(cells.items(), key=lambda kv: (str(kv[0][0]), kv[0][1]))}
275
+ summary["atomic"] = summarize(eval_chains(model, v, P, rel, atomic, device, loops=T_of_d))
276
+ if looped:
277
+ summary["loop_rule"] = dict(rule="min(cap, smallest power of two >= max(8, d))", cap=args.loop_cap, T_by_depth={d: T_of_d(d) for d in sorted({t["d"] for t in pool})})
278
+ summary[f"fixed_T{args.loop_fixed}"] = cells_of(eval_chains(model, v, P, rel, pool, device, loops=lambda d: args.loop_fixed))
279
+ # plan §6.1: EXTERNAL step-wise calling -- the program feeds the model one-step prompts Q s_t r_(t+1) ANS and
280
+ # chains the predicted entities itself (identifies whether the atomic update is reliable; not autonomous)
281
+ stepwise = collections.defaultdict(lambda: [0, 0])
282
+ for d in sorted({t["d"] for t in pool}):
283
+ cs = [t for t in pool if t["d"] == d]
284
+ cur = [c["x"] for c in cs]; ok = [True] * len(cs)
285
+ for step in range(d):
286
+ prompts = [prompt_ids(v, s_, [c["rels"][step]]) for s_, c in zip(cur, cs)]
287
+ gen = rollout(model, prompts, completion_length(P, 1) + 1, device, n_loops=T1)
288
+ nxt = []
289
+ for i, g in enumerate(gen):
290
+ pr = parse_completion(v, P, g.tolist(), 1)
291
+ st_ = pr["states"][-1] if pr["states"] else -1
292
+ ok[i] = ok[i] and (st_ == cs[i]["states"][step]); nxt.append(st_ if st_ >= 0 else 0)
293
+ cur = nxt
294
+ for c, o in zip(cs, ok):
295
+ stepwise[f"{c['cat']}|d{d}"][0] += int(o); stepwise[f"{c['cat']}|d{d}"][1] += 1
296
+ summary["stepwise_external"] = {k: dict(n=n, final_acc=c / n) for k, (c, n) in sorted(stepwise.items())}
297
+ summary["val_d2"] = summarize(eval_chains(model, v, P, rel, val_chains, device, loops=T_of_d))
298
+ summary["val_w2"] = eval_w2_val(model, v, P, rel, val_w2, device, n_loops=T1)
299
+ summary["checkpoint"] = ckpt_path; summary["eval_seconds"] = time.time() - t_eval
300
+ json.dump(summary, open(os.path.join(args.save_dir, f"final_eval_{tag}.json"), "w"), indent=1)
301
+ with open(os.path.join(args.save_dir, f"predictions_{tag}.jsonl"), "w") as f:
302
+ for c, r in res:
303
+ f.write(json.dumps(dict(r, x=c["x"], rels=c["rels"], gold_states=c["states"], new_adjacencies=c.get("new_adjacencies"))) + "\n")
304
+ print(f"[final eval {tag}] " + " ".join(f"{k}:{s['final_acc']:.3f}/{s['traj_acc']:.3f}" for k, s in summary.items() if isinstance(s, dict) and "final_acc" in s and "|" in k and "repeated" not in k and "revisits" not in k))
305
+ return summary
306
+
307
+ if args.eval_only:
308
+ full_eval(args.eval_tag or "evalonly", args.eval_only); return
309
+
310
+ def loop_matrix(u):
311
+ """depth x T matrix on the monitoring set (plan §6: fixed checkpoints, fixed 500 chains per cell, every T)."""
312
+ out = {}
313
+ for T in [int(x) for x in args.loop_matrix_T.split(",")]:
314
+ out[f"T{T}"] = cells_of(eval_chains(model, v, P, rel, small_tests, device, loops=lambda d, T=T: T))
315
+ json.dump(dict(update=u, matrix=out), open(os.path.join(args.save_dir, f"loop_matrix_{u:06d}.json"), "w"), indent=1)
316
+
317
+ opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
318
+ s_w2 = Stream(W2.n, args.seed, "w2"); s_d2 = Stream(D2.n, args.seed, "d2"); s_d3 = Stream(D3.n, args.seed, "d3") if D3 else None
319
+ loop_rng = np.random.default_rng([args.seed, int(hashlib.sha256(b"loopT").hexdigest()[:8], 16)])
320
+ matrix_at = {int(x) for x in args.loop_matrix_at.split(",")} if looped else set()
321
+ total = args.phase1_updates + args.phase2_updates
322
+ start = 1; best = (-1.0, float("inf"))
323
+ metrics_path = os.path.join(args.save_dir, "metrics.jsonl"); log_path = os.path.join(args.save_dir, "train_log.jsonl")
324
+ last_path = os.path.join(args.save_dir, "last.pt")
325
+ if args.resume and os.path.exists(last_path):
326
+ st = torch.load(last_path, map_location=device, weights_only=False)
327
+ model.load_state_dict(st["model"]); opt.load_state_dict(st["opt"]); s_w2.load(st["s_w2"]); s_d2.load(st["s_d2"])
328
+ if s_d3 and st.get("s_d3"): s_d3.load(st["s_d3"])
329
+ if st.get("loop_rng"): loop_rng.bit_generator.state = st["loop_rng"]
330
+ torch.set_rng_state(st["torch_rng"]); start = st["update"] + 1; best = tuple(st["best"])
331
+ print(f"resumed from update {st['update']}")
332
+ else:
333
+ open(metrics_path, "w").close(); open(log_path, "w").close()
334
+ manifest = dict(protocol=P, seed=args.seed, n_params=n_params, rows_per_update=args.rows_per_update,
335
+ phase1_updates=args.phase1_updates, phase2_updates=args.phase2_updates, L=L,
336
+ w2=dict(rows=W2.n, tokens=W2.n_tokens, supervised_positions=W2.n_supervised),
337
+ d2=dict(rows=D2.n, tokens=D2.n_tokens, supervised_positions=D2.n_supervised),
338
+ d3=(dict(rows=D3.n, tokens=D3.n_tokens, supervised_positions=D3.n_supervised, per_update=args.d3_per_update) if D3 else None),
339
+ data_hashes=meta.get("hashes"), model=dict(d_model=args.d_model, n_layer=args.n_layer, n_head=args.n_head, max_len=args.max_len, rope_base=args.rope_base, use_rope=not args.no_rope),
340
+ optimizer=dict(lr=args.lr, weight_decay=args.weight_decay, warmup=args.warmup), vocab_size=len(v.vocab),
341
+ vocab_ext=args.vocab_ext, select_by=args.select_by,
342
+ loop=(dict(train_T=loop_T, cap=args.loop_cap, fixed_control=args.loop_fixed, zero_init_out=args.zero_init_out,
343
+ rule="min(cap, smallest power of two >= max(8, d))", matrix_at=sorted(matrix_at), matrix_T=args.loop_matrix_T) if looped else None))
344
+ json.dump(manifest, open(os.path.join(args.save_dir, "manifest.json"), "w"), indent=1)
345
+ print(f"protocol {P} | params {n_params:,} | L {L} | w2 rows {W2.n} d2 rows {D2.n} | device {device}")
346
+
347
+ run_ce = collections.defaultdict(float); run_cnt = collections.defaultdict(int); run_loss = 0.0; run_n = 0; tok_seen = 0
348
+ t0 = time.time()
349
+ for u in range(start, total + 1):
350
+ model.train()
351
+ phase = 1 if u <= args.phase1_updates else 2
352
+ if phase == 1:
353
+ idx_w2 = s_w2.take(args.rows_per_update); idx_d2 = None
354
+ else:
355
+ n_d3 = args.d3_per_update if D3 else 0
356
+ idx_w2 = s_w2.take(args.rows_per_update // 2); idx_d2 = s_d2.take(args.rows_per_update // 2 - n_d3)
357
+ idx_d3 = s_d3.take(n_d3) if n_d3 else None
358
+ parts = [(W2, torch.from_numpy(idx_w2).to(device), "w2")]
359
+ if idx_d2 is not None:
360
+ parts.append((D2, torch.from_numpy(idx_d2).to(device), "d2"))
361
+ if phase == 2 and D3 is not None and args.d3_per_update:
362
+ parts.append((D3, torch.from_numpy(idx_d3).to(device), "d3"))
363
+ inp = torch.cat([R.input_ids[i] for R, i, _ in parts]); tgt = torch.cat([R.target_ids[i] for R, i, _ in parts])
364
+ w = torch.cat([R.weights[i] for R, i, _ in parts]); roles = torch.cat([R.roles[i] for R, i, _ in parts]); pm = torch.cat([R.pad_mask[i] for R, i, _ in parts])
365
+ for g in opt.param_groups:
366
+ g["lr"] = args.lr * min(1.0, u / max(1, args.warmup))
367
+ T_u = int(loop_rng.choice(loop_T)) if looped else 1
368
+ logits = model(inp, pad_mask=pm, n_loops=T_u)
369
+ ce = F.cross_entropy(logits.transpose(1, 2), tgt, reduction="none", ignore_index=-100)
370
+ loss = (ce * w).sum() / args.rows_per_update
371
+ opt.zero_grad(set_to_none=True); loss.backward(); opt.step()
372
+ with torch.no_grad():
373
+ run_loss += float(loss); run_n += 1; tok_seen += int((~pm).sum())
374
+ off = 0
375
+ for R, i, name in parts:
376
+ k = len(i); sl = slice(off, off + k); off += k
377
+ for rk, rname in enumerate(ROLE_NAMES):
378
+ if rk == 0:
379
+ continue
380
+ m = roles[sl] == rk
381
+ if m.any():
382
+ run_ce[f"{name}/{rname}"] += float(ce[sl][m].sum()); run_cnt[f"{name}/{rname}"] += int(m.sum())
383
+ if u % 100 == 0 or u == 1:
384
+ with open(log_path, "a") as f:
385
+ f.write(json.dumps(dict(update=u, phase=phase, loss=run_loss / max(1, run_n), lr=opt.param_groups[0]["lr"], tokens_seen=tok_seen,
386
+ elapsed=time.time() - t0, T=T_u, **{f"ce/{k}": run_ce[k] / run_cnt[k] for k in run_ce if run_cnt[k]})) + "\n")
387
+ run_ce.clear(); run_cnt.clear(); run_loss = 0.0; run_n = 0
388
+ if u % args.eval_every == 0 or u == total or u == args.phase1_updates:
389
+ model.eval()
390
+ t_eval = time.time()
391
+ rec = dict(update=u, phase=phase, elapsed=time.time() - t0, tokens_seen=tok_seen,
392
+ tf_val=teacher_forced(model, VAL, device, n_loops=T1), tf_train_d2=teacher_forced(model, SUB_D2, device, n_loops=T1),
393
+ tf_train_w2=teacher_forced(model, SUB_W2, device, n_loops=T1), tf_val_w2=teacher_forced(model, VALW2, device, n_loops=T1))
394
+ rec["atomic"] = summarize(eval_chains(model, v, P, rel, atomic, device, loops=T_of_d))
395
+ rec["val_d2"] = summarize(eval_chains(model, v, P, rel, val_chains, device, loops=T_of_d))
396
+ rec["val_w2"] = eval_w2_val(model, v, P, rel, val_w2, device, n_loops=T1)
397
+ res = eval_chains(model, v, P, rel, small_tests, device, loops=T_of_d)
398
+ cells = collections.defaultdict(list)
399
+ for c, r in res:
400
+ cells[f"{c['cat']}|d{c['d']}"].append((c, r))
401
+ rec["test_small"] = {k: summarize(vv) for k, vv in sorted(cells.items())}
402
+ if u in matrix_at:
403
+ loop_matrix(u)
404
+ rec["eval_seconds"] = time.time() - t_eval
405
+ with open(metrics_path, "a") as f:
406
+ f.write(json.dumps(rec) + "\n")
407
+ print(f"u {u} ph{phase} | val_d2 {rec['val_d2']['final_acc']:.3f}/{rec['val_d2']['traj_acc']:.3f} | atomic {rec['atomic']['final_acc']:.3f} | "
408
+ + " ".join(f"{k}:{s['final_acc']:.2f}" for k, s in rec["test_small"].items()) + f" | eval {rec['eval_seconds']:.0f}s | {time.time() - t0:.0f}s")
409
+ if args.select_by == "w2":
410
+ score = (rec["val_w2"]["both_acc"], -rec["tf_val_w2"]["loss_per_row"]); best_name = "best_by_valw2.pt"
411
+ else:
412
+ score = (rec["val_d2"]["final_acc"], -rec["tf_val"]["loss_per_row"]); best_name = "best_by_val.pt"
413
+ if score > (best[0], -best[1]):
414
+ best = (score[0], -score[1]); torch.save(dict(model=model.state_dict(), update=u), os.path.join(args.save_dir, best_name))
415
+ if u % args.save_every == 0 or u == total:
416
+ torch.save(dict(model=model.state_dict(), update=u), os.path.join(args.save_dir, f"ckpt_{u:06d}.pt"))
417
+ torch.save(dict(model=model.state_dict(), opt=opt.state_dict(), update=u, s_w2=s_w2.state(), s_d2=s_d2.state(),
418
+ s_d3=(s_d3.state() if s_d3 else None), loop_rng=loop_rng.bit_generator.state, torch_rng=torch.get_rng_state(),
419
+ best=list(best), tokens_seen=tok_seen), last_path)
420
+ if not args.no_final_eval:
421
+ full_eval("final")
422
+ extra = [] if args.final_extra_ckpts == "" else args.final_extra_ckpts.split(",")
423
+ if args.final_extra_ckpts == "auto":
424
+ extra = [f"ckpt_{total - k * args.save_every:06d}.pt" for k in (1, 2) if total - k * args.save_every > 0]
425
+ for ck in extra:
426
+ if os.path.exists(os.path.join(args.save_dir, ck)):
427
+ full_eval(ck.replace(".pt", ""), os.path.join(args.save_dir, ck))
428
+ best_name = "best_by_valw2.pt" if args.select_by == "w2" else "best_by_val.pt"
429
+ if os.path.exists(os.path.join(args.save_dir, best_name)):
430
+ full_eval("best", os.path.join(args.save_dir, best_name))
431
+ print("done")
432
+
433
+
434
+ if __name__ == "__main__":
435
+ main()
for Minegishi/reproducibility/train_halt.py ADDED
@@ -0,0 +1,996 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Sixth batch: externally set loop count (D) vs learned halting (H) on the paper-aligned full-sequence loop Transformer
3
+ (docs/experiments_learned_halting_2026-09-19.md, docs/experiments_paper_aligned_loop_2026-09-19.md; log docs/experiments_halting_log.md).
4
+
5
+ python train_halt.py --data_dir ../data/chain_loop --atomic ../data/atomic_joint_2026-09-19/train_atomic.json \
6
+ --save_dir ../runs/halt_H_s1 --arm H --seed 1 [--resume] [--cuda_graph 1]
7
+
8
+ Arms D the program sets R = d (number of relations in the input): atomic queries are read after 1 loop, d2 after 2, a depth-d
9
+ test chain after d loops. loss = CE at loop d
10
+ H the model decides CONTINUE / STOP every loop with a stop head on the read-out state; trained with the expected final-answer
11
+ loss under its own halting distribution, unrolled --b_train loops, remaining mass collected at the last unrolled loop
12
+ (logged as truncation mass, never credited as an active STOP), plus --lam * E[T] (main arm: lam = 0).
13
+ No halting labels, no intermediate entities, no loop index. Inference: STOP at the first loop with p_t >= --stop_thr
14
+ (pre-registered 0.5), global cap --b_eval loops for every depth; hitting the cap = timeout = failure.
15
+ P paper-style baseline (extra, not part of the D / H comparison): R ~ clip(Poisson(4), 2, 8) per batch, CE at loop R.
16
+ Data one update = 128 single-chain queries = 32 atomic + 16 w2 sources rendered as 32 independent atomic queries + 64 d2; the
17
+ mean of the 128 final-answer CEs. Inputs are compact [x, r] / [x, r1, r2] (no wrappers, no END); the label is never an input.
18
+ Training depth <= 2; deeper chains exist only in evaluation.
19
+ Optim AdamW lr 1e-4, wd 0.01, 2000 warm-up updates then constant, no label smoothing, global grad-norm clip 1.0, fp32.
20
+ """
21
+ import argparse, math
22
+ import collections
23
+ import hashlib
24
+ import json
25
+ import os
26
+ import time
27
+ import signal
28
+ import sys
29
+ from eval_checkpoints import save_eval_checkpoint
30
+ from objectives import (CS_COMPONENTS, CS_MODES, SIMPLE_CS_MODES, loss_shallow, loss_consistency,
31
+ validate_simple_cs_config, simple_cs_metadata, sample_simple_cs_tokens,
32
+ make_cs_noise_generator, fill_cs_noise_, fill_cs_alpha_)
33
+
34
+ import numpy as np
35
+ import torch
36
+ import torch.nn.functional as F
37
+
38
+ from model.loop_gpt import LoopGPT
39
+ from streams import Stream, load_json
40
+
41
+ N_AT, N_W2SRC, N_D2 = 32, 16, 64
42
+ N_ROWS = N_AT + 2 * N_W2SRC + N_D2
43
+ GRID = [1, 2, 4, 8, 16, 32, 64, 128, 256]
44
+
45
+
46
+ def _m(x):
47
+ return x.mean() if x.numel() else x.new_zeros(())
48
+
49
+
50
+ def named_rng(seed, name):
51
+ return np.random.default_rng([seed, int(hashlib.sha256(name.encode()).hexdigest()[:8], 16)])
52
+
53
+
54
+ def arm_ok(args):
55
+ return args.arm == "D" and not args.joint and not args.prefix_sup
56
+
57
+
58
+ # --------------------------------------------------------------------------------------------------------------- losses
59
+ def loss_fn(model, arm, tok, last, tgt, R=None, b_train=8, lam=0.0):
60
+ """-> (loss, stats dict of 0-d / 1-d tensors). Row layout: [32 atomic | 32 w2-derived atomic | 64 d2]."""
61
+ h = model.embed(tok); n1 = N_AT + 2 * N_W2SRC; st = {}
62
+ if arm == "D":
63
+ if model.reembed:
64
+ o1, s1 = model.step_state(model.embed_state(tok)); o2, _ = model.step_state(s1)
65
+ z1 = model.read(o1, last); z2 = model.read(o2, last)[n1:]
66
+ else:
67
+ h = model.step(h); z1 = model.read(h, last)
68
+ h2 = model.step(h[n1:]); z2 = model.read(h2, last[n1:])
69
+ ce = F.cross_entropy(torch.cat([model.logits(z1[:n1]), model.logits(z2)]).float(), tgt, reduction="none")
70
+ loss = ce.mean()
71
+ elif arm == "P":
72
+ for _ in range(R):
73
+ h = model.step(h)
74
+ ce = F.cross_entropy(model.logits(model.read(h, last)).float(), tgt, reduction="none"); loss = ce.mean()
75
+ else:
76
+ ces, ps = [], []; stt = model.embed_state(tok) # state API: identical to model.step for the plain loop, required by the re-embedding loops
77
+ for _ in range(b_train):
78
+ h, stt = model.step_state(stt); z = model.read(h, last)
79
+ ces.append(F.cross_entropy(model.logits(z).float(), tgt, reduction="none")); ps.append(model.stop_logit(z).float())
80
+ ce_t = torch.stack(ces); sl = torch.stack(ps); p = torch.sigmoid(sl) # (B_train, N)
81
+ # survival prod_{j<=t}(1 - p_j) in log space: log(1 - sigmoid(s)) = logsigmoid(-s) (cumprod's backward syncs with the host,
82
+ # which a captured CUDA graph does not allow; the log form is also the numerically safer one)
83
+ surv = torch.exp(torch.cumsum(F.logsigmoid(-sl), 0)); prev = torch.cat([torch.ones_like(surv[:1]), surv[:-1]])
84
+ q = torch.cat([(p * prev)[:-1], prev[-1:]]) # last loop collects ALL remaining mass
85
+ steps = torch.arange(1, b_train + 1, device=tok.device, dtype=q.dtype)[:, None]
86
+ ET = (q * steps).sum(0); ce = (q * ce_t).sum(0); loss = (ce + lam * ET).mean()
87
+ st.update(ce_by_loop_atomic=ce_t[:, :n1].mean(1), ce_by_loop_d2=ce_t[:, n1:].mean(1), q_atomic=q[:, :n1].mean(1), q_d2=q[:, n1:].mean(1),
88
+ tail_atomic=surv[-1, :n1].mean(), tail_d2=surv[-1, n1:].mean(), ET_atomic=ET[:n1].mean(), ET_d2=ET[n1:].mean())
89
+ st.update(ce_atomic=_m(ce[:N_AT]), ce_w2=_m(ce[N_AT:n1]), ce_d2=_m(ce[n1:]))
90
+ return loss, st
91
+
92
+
93
+ def loss_fn_untied(model, tok, last, tgt, read, L, depth=None):
94
+ """Arm D on the VANILLA control (L different layers, one pass; `last` equals the depth of the row, 1 or 2). read =
95
+ 'at_d' : the answer of a depth-p row is read after layer p (the analogue of R = d; layers > 2 never receive gradient);
96
+ 'final': every answer is read after the last layer (the ordinary fixed-depth transformer);
97
+ 'all' : the answer of a depth-p row is supervised after EVERY layer l >= p (final-answer labels only, depth <= 2 rows only; per-row mean
98
+ over its layers) -- the only one of the three in which every layer is trained."""
99
+ n1 = N_AT + 2 * N_W2SRC; stt = model.embed_state(tok); ces = [] # L = steps: layers of the vanilla control, or loops of the tied model (--unroll)
100
+ depth = last if depth is None else depth # with a prefix (--offset_max) the read position `last` = offset + depth
101
+ for _ in range(2 if read == "at_d" else L):
102
+ o, stt = model.step_state(stt); ces.append(F.cross_entropy(model.logits(model.read(o, last)).float(), tgt, reduction="none"))
103
+ ces = torch.stack(ces) # (layers run, N)
104
+ if read == "final":
105
+ ce = ces[-1]
106
+ else:
107
+ layer = torch.arange(1, ces.shape[0] + 1, device=tok.device)[:, None]
108
+ m = ((layer == depth[None, :]) if read == "at_d" else (layer >= depth[None, :])).to(ces.dtype); ce = (ces * m).sum(0) / m.sum(0)
109
+ st = dict(ce_atomic=_m(ce[:N_AT]), ce_w2=_m(ce[N_AT:n1]), ce_d2=_m(ce[n1:]), ce_last_layer_d2=_m(ces[-1, n1:]))
110
+ return ce.mean(), st
111
+
112
+
113
+ def offset_sampler(dist, K):
114
+ """-> f(rng, n): prefix lengths in 0..K. 'uniform'; 'boltz<b>': p(k) ~ exp(b * k / K) (b > 0 puts more mass on LONG prefixes -- longer sequences are harder and get
115
+ more data); 'poisson<lam>': Poisson(lam) truncated to 0..K (peaked)."""
116
+ ks = np.arange(K + 1, dtype=np.float64)
117
+ if dist == "uniform":
118
+ w = np.ones(K + 1)
119
+ elif dist.startswith("boltz"):
120
+ w = np.exp(float(dist[5:]) * ks / max(1, K))
121
+ elif dist.startswith("poisson"):
122
+ lam = float(dist[7:]); w = np.exp(ks * np.log(lam) - lam - np.array([math.lgamma(k + 1) for k in ks]))
123
+ else:
124
+ raise ValueError(dist)
125
+ p = w / w.sum()
126
+ return lambda rng, n: rng.choice(K + 1, size=n, p=p)
127
+
128
+
129
+ def pack_rows_prefix(q, pack, pad_id, pk, ks, ents):
130
+ """Host-side packing WITH a random-entity prefix per sequence: sequence r starts with ks[r] random entity tokens (ents[r, :ks[r]]), then queries are laid end to end;
131
+ a query may start at any column < pack (so the width is pack + 2). q: (n, 6) numpy [t0, t1, t2 | PAD, depth, gold, gold1]. -> number of sequences used."""
132
+ n = q.shape[0]; W = pk["tok"].shape[1]; S = pk["seg_last"].shape[1]
133
+ tok = np.full((n, W), pad_id, dtype=np.int64); valid = np.zeros((n, W), dtype=bool); pv = np.zeros((n, W), dtype=np.int64); fire = np.zeros((n, W), dtype=np.int64)
134
+ sl = np.zeros((n, S), dtype=np.int64); sd = np.ones((n, S), dtype=np.int64); sg = np.zeros((n, S), dtype=np.int64); sv = np.zeros((n, S), dtype=bool)
135
+ r = -1; col = pack; slot = 0
136
+ for i in range(n):
137
+ d = int(q[i, 3])
138
+ if col >= pack:
139
+ r += 1; k = int(ks[r]); tok[r, :k] = ents[r, :k]; valid[r, :k] = True; col = k; slot = 0
140
+ tok[r, col:col + d + 1] = q[i, :d + 1]; valid[r, col:col + d + 1] = True; fire[r, col:col + d + 1] = np.arange(d + 1)
141
+ pv[r, col + 1] = q[i, 5]
142
+ if d == 2:
143
+ pv[r, col + 2] = q[i, 4]
144
+ sl[r, slot] = col + d; sd[r, slot] = d; sg[r, slot] = q[i, 4]; sv[r, slot] = True; slot += 1; col += d + 1
145
+ used = r + 1
146
+ for name, arr in (("tok", tok), ("valid", valid), ("pv", pv), ("fire", fire), ("seg_last", sl), ("seg_depth", sd), ("seg_gold", sg), ("seg_valid", sv)):
147
+ pk[name][:used].copy_(torch.from_numpy(arr[:used]))
148
+ return used
149
+
150
+
151
+ def pack_rows(q, pack, pad_id, pk):
152
+ """q (n, 6) = [t0, t1, t2 | PAD, depth, gold, gold1] in the order in which the queries are laid end to end; a new sequence starts every `pack` tokens (a query that starts
153
+ before the boundary is finished in the same sequence, hence width pack + 2). Fills the static buffers pk[...] in place, without host synchronisation."""
154
+ n = q.shape[0]; ln = q[:, 3] + 1; start = torch.cumsum(ln, 0) - ln; ri = start // pack; col = start % pack
155
+ slot = torch.arange(n, device=q.device) - torch.searchsorted(ri.contiguous(), ri.contiguous()); is2 = (q[:, 3] == 2); one = torch.ones_like(is2); z = torch.zeros_like(ri)
156
+ for b in pk.values():
157
+ b.zero_()
158
+ pk["tok"].fill_(pad_id); pk["seg_depth"].fill_(1)
159
+ # write order matters: the third column of an atomic query is the first column of the next query
160
+ pk["tok"].index_put_((ri, col + 2), q[:, 2]); pk["valid"].index_put_((ri, col + 2), is2); pk["fire"].index_put_((ri, col + 2), 2 * is2.long()); pk["pv"].index_put_((ri, col + 2), q[:, 4] * is2.long())
161
+ pk["tok"].index_put_((ri, col + 1), q[:, 1]); pk["valid"].index_put_((ri, col + 1), one); pk["fire"].index_put_((ri, col + 1), z + 1); pk["pv"].index_put_((ri, col + 1), q[:, 5])
162
+ pk["tok"].index_put_((ri, col), q[:, 0]); pk["valid"].index_put_((ri, col), one); pk["fire"].index_put_((ri, col), z); pk["pv"].index_put_((ri, col), z)
163
+ pk["seg_last"].index_put_((ri, slot), col + q[:, 3]); pk["seg_depth"].index_put_((ri, slot), q[:, 3]); pk["seg_gold"].index_put_((ri, slot), q[:, 4]); pk["seg_valid"].index_put_((ri, slot), one)
164
+
165
+
166
+ def loss_packed(model, tok, valid, pv, fire, seg_last, seg_depth, seg_gold, seg_valid, anchor_w, extra):
167
+ """WIDTH-n RENDERING (user proposal): the shallow queries of an update are PACKED into long causal sequences -- n independent d1 / d2 queries side by side,
168
+ <e_x1><r_1> <e_x2><r_2><s_2> <e_x3><r_3> ..., no separator -- and every query is read at its OWN last relation token (after loop = its depth). Only shallow labels are
169
+ used (atomic answers and d2 answers), yet relation tokens now sit at every position of a long sequence and have to read the token right before them among many earlier
170
+ entity-like states: what a deep chain looks like locally. tok / valid / pv / fire: (R, W); fire = loop at which a relation position produces its value (0 = never), pv = that
171
+ value's token id; seg_*: (R, S) per-query read position, depth, gold, validity. With anchor_w > 0 the anchored-state loss is applied to every position (see loss_anchor).
172
+ -> per-query CE (R, S) (0 where invalid), anchor loss"""
173
+ emb = model.embed(tok); e_tok = emb.detach(); e_pv = model.wte(pv).detach(); h = emb; ces = []; anc = emb.new_zeros(()); T = 2 + extra; C = emb.shape[-1]; v = valid.to(emb.dtype)
174
+ idx = seg_last[..., None].expand(-1, -1, C)
175
+ for t in range(1, T + 1):
176
+ h = model.step(h)
177
+ if anchor_w:
178
+ target = torch.where(((fire > 0) & (fire <= t))[..., None], e_pv, e_tok); anc = anc + (((h - target) ** 2).sum(-1) * v).sum() / ((target ** 2).sum(-1) * v).sum()
179
+ lg = model.logits(model.ln_f(h.gather(1, idx))).float(); ces.append(F.cross_entropy(lg.flatten(0, 1), seg_gold.flatten(), reduction="none").view_as(seg_gold))
180
+ ces = torch.stack(ces); loop = torch.arange(1, T + 1, device=tok.device)[:, None, None]; m = ((loop >= seg_depth[None]) & (loop <= seg_depth[None] + extra) & seg_valid[None]).to(ces.dtype)
181
+ return (ces * m).sum(0) / m.sum(0).clamp(min=1), anc / T
182
+
183
+
184
+ def stepper(model):
185
+ """h -> h': the loop map F of the tied model; for the VANILLA control (--untied, --untied_group g) the t-th call applies layer group t, and a call beyond the last
186
+ group leaves h unchanged (a network of L groups has no group L + 1). Used by loss_anchor so that the vanilla control is trained with exactly the loop's schedule."""
187
+ if not getattr(model, "untied", False):
188
+ return model.step
189
+ cnt = [0]
190
+
191
+ def f(h):
192
+ h = model.step_group(h, cnt[0]); cnt[0] += 1; return h
193
+ return f
194
+
195
+
196
+ def loss_consist(model, tok, tfire, ent_ids, ent_lo, ent_hi):
197
+ """ONE-STEP CONSISTENCY ON CONSTRUCTED STATES (label-free, one loop; handover section 7). A deep-chain state after loop t is CONSTRUCTED instead of computed: positions
198
+ 1..t hold entity embeddings (random entities -- the format constraint does not care which), the rest hold their token embeddings. One loop is applied; every position except
199
+ t + 1 (the one whose turn it is -- its hop is left to the labelled rows) must come back to where it was. Random chains tok (F, D + 1), per-row t = tfire (F,), random entity ids
200
+ ent_ids (F, D + 1). Bounded symmetric error. -> scalar"""
201
+ emb = model.embed(tok); Wr = model.wte(ent_ids); pos = torch.arange(tok.shape[1], device=tok.device)[None]
202
+ fired = ((pos >= 1) & (pos <= tfire[:, None]))[..., None]; S = torch.where(fired, Wr, emb).detach()
203
+ # VANILLA control: every layer group plays the role of F and receives the same constructed states F receives (per-map exposure matched); mean over groups
204
+ hs = [model.step_group(S, g) for g in range(model.n_groups)] if getattr(model, "untied", False) else [model.step(S)]
205
+ keep = (pos != (tfire[:, None] + 1)).to(emb.dtype)
206
+ return sum(((((h - S) ** 2).sum(-1) / ((h ** 2).sum(-1) + (S ** 2).sum(-1))) * keep).sum() / keep.sum() for h in hs) / len(hs)
207
+
208
+
209
+ def loss_format(model, tok, extra, ent_lo, ent_hi, bounded=True):
210
+ """LABEL-FREE FORMAT LOSS on random chains of ANY depth (handover_stability_2026-09-22.md, method A). tok (F, D + 1) = [random entity, D random relations]; no answer is
211
+ known or used. After every loop t = 1..D + extra, every position is pulled to a loop-independent target: positions that have not fired yet (p > t) and the head -> their own
212
+ token embedding; fired positions (1 <= p <= t) -> the NEAREST entity embedding (whichever it is), which also makes the fired state a fixed point. Relative squared error,
213
+ targets detached. This is the training-time version of the inference-time SNAP diagnostic: it pins the waiting / holding dynamics at loop counts and positions the shallow
214
+ labelled rows never reach, without any deep label and without changing the architecture. -> scalar"""
215
+ emb = model.embed(tok); e_tok = emb.detach(); W = model.wte.weight[ent_lo:ent_hi].detach(); Wn = (W ** 2).sum(-1); h = emb; F_, L = tok.shape; pos = torch.arange(L, device=tok.device)[None]
216
+ T = L - 1 + extra; tot = emb.new_zeros(())
217
+ for t in range(1, T + 1):
218
+ h = model.step(h); hd = h.detach()
219
+ near = W[((hd ** 2).sum(-1, keepdim=True) - 2 * hd @ W.t() + Wn).argmin(-1)]; fired = ((pos >= 1) & (pos <= t))[..., None]; target = torch.where(fired, near, e_tok)
220
+ if bounded: # symmetric relative error in [0, 2]: an exploding residual stream cannot blow the loss up (the unbounded form reached 1e4
221
+ tot = tot + (((h - target) ** 2).sum(-1) / ((h ** 2).sum(-1) + (target ** 2).sum(-1))).mean() # and stalled the shallow task, see log 13.7)
222
+ else:
223
+ tot = tot + ((h - target) ** 2).sum() / (target ** 2).sum()
224
+ return tot / T
225
+
226
+
227
+ def loss_anchor(model, tok, depth, tgt, pv, extra, off=None, n_loops=None, k_delay=0, pad_id=None):
228
+ """Arm D, STANDARD loop, LEARNING-OBJECTIVE remedy for the diagnosed blocker (log section 12: the states a hop reads drift with the number of loops run, so the hop only
229
+ works at the loop counts seen in training). ANCHORED STATES: between loops every position is pulled towards a loop-independent vector --
230
+ its own token embedding while it has not fired yet (loop t < position p), and for the head entity always,
231
+ the embedding of its prefix value v_p once it has fired (t >= p), for ALL later loops (so finished states are fixed points),
232
+ with a relative squared error; targets are detached. The hop at any depth then reads exactly what the hop at depth 1 / 2 reads: an entity embedding at its predecessor and
233
+ a relation embedding at itself. No architectural change (full causal attention, residual stream, R = d); labels used: the final answer and, for depth >= 2 rows, the prefix
234
+ values, which for a depth-2 row is just the atomic fact r1(x). tok (N, L) = [x, r1..rd, PAD..]; pv (N, L) = token id of v_p at relation position p.
235
+ -> per-row CE (N,) on the final answer at loops d .. d + extra, anchor loss (scalar)"""
236
+ N, L = tok.shape; emb = model.embed(tok); pos = torch.arange(L, device=tok.device)[None]; off = torch.zeros_like(depth) if off is None else off; stp = stepper(model)
237
+ # with --offset_max the chain sits at columns off .. off + depth behind `off` filler tokens (random ENTITY tokens: once states are anchored a finished prefix of a deep
238
+ # chain IS a run of entity embeddings, so a shallow chain behind random entities is what the last hops of a deep chain look like); fillers are anchored to themselves
239
+ rel_pos = pos - off[:, None]; valid = (rel_pos <= depth[:, None]).to(emb.dtype); last = depth + off
240
+ e_tok = emb.detach(); e_pv = model.wte(pv).detach(); h = emb; ces = []; anc = 0.0; T = (L - 1 if n_loops is None else n_loops) + extra
241
+ if k_delay:
242
+ # DELAYED START with anchoring: for k loops the head entity is hidden (PAD embedding), every position is anchored to what it holds, nothing can fire; then the entity is
243
+ # written in and the chain runs as usual. With shallow labels only, every relation position has now WAITED k extra loops (in a static context) before its hop.
244
+ head = (rel_pos == 0)[..., None]; h = torch.where(head, model.wte.weight[pad_id].expand_as(emb), emb); e_wait = h.detach()
245
+ for _ in range(k_delay): # bounded (symmetric) error during the wait: an untrained loop inflates the residual stream and the unbounded form blew up (log 13.8)
246
+ h = stp(h); anc = anc + ((((h - e_wait) ** 2).sum(-1) / ((h ** 2).sum(-1) + (e_wait ** 2).sum(-1))) * valid).sum() / valid.sum()
247
+ h = torch.where(head, emb, h)
248
+ for t in range(1, T + 1):
249
+ h = stp(h); fired = ((rel_pos >= 1) & (rel_pos <= t))[..., None]; target = torch.where(fired, e_pv, e_tok)
250
+ anc = anc + (((h - target) ** 2).sum(-1) * valid).sum() / ((target ** 2).sum(-1) * valid).sum()
251
+ ces.append(F.cross_entropy(model.logits(model.read(h, last)).float(), tgt, reduction="none"))
252
+ ces = torch.stack(ces); loop = torch.arange(1, T + 1, device=tok.device)[:, None]; m = ((loop >= depth[None, :]) & (loop <= depth[None, :] + extra)).to(ces.dtype)
253
+ return (ces * m).sum(0) / m.sum(0), anc / (T + k_delay)
254
+
255
+
256
+ def loss_delay_persist(model, tok, last, tgt, k_delay, persist, pad_id):
257
+ """Arm D, plain loop, two DATA-PROTOCOL augmentations (no architectural change, evaluation unchanged):
258
+ k_delay : the head entity ARRIVES LATE -- for the first k loops position 0 holds the PAD embedding, after loop k its state is overwritten by the entity's embedding; a row of
259
+ depth d is then read after loop k + d. Every relation position therefore waits k loops longer than its depth requires (waiting ages beyond the training depth
260
+ are seen with shallow chains).
261
+ persist : the answer of a depth-d row is supervised after EVERY loop d .. d + persist (per-row mean), so a finished answer has to stay (fixed point) and finished states are
262
+ seen at larger loop counts.
263
+ `last` = depth of the row. -> per-row CE (N,)"""
264
+ emb = model.embed(tok); h = emb
265
+ if k_delay:
266
+ h = emb.clone(); h[:, 0] = model.wte.weight[pad_id]
267
+ for _ in range(k_delay):
268
+ h = model.step(h)
269
+ h = torch.cat([emb[:, :1], h[:, 1:]], 1)
270
+ dmax = tok.shape[1] - 1; ces = []
271
+ for t in range(1, dmax + persist + 1):
272
+ h = model.step(h); ces.append(F.cross_entropy(model.logits(model.read(h, last)).float(), tgt, reduction="none"))
273
+ ces = torch.stack(ces); loop = torch.arange(1, ces.shape[0] + 1, device=tok.device)[:, None]; m = ((loop >= last[None, :]) & (loop <= last[None, :] + persist)).to(ces.dtype)
274
+ return (ces * m).sum(0) / m.sum(0)
275
+
276
+
277
+ def loss_deep(model, tok, last, tgt, rows_per_depth, depths):
278
+ """Arm D, DEEPER TRAINING ROWS (--deep): rows are sorted by depth (rows_per_depth rows for each depth in `depths`, ascending); a row of depth d is read after loop d
279
+ (R = d, the same protocol as the shallow rows), then leaves the batch. -> per-row CE (n_deep,)"""
280
+ stt = model.embed_state(tok); ces = []; M = rows_per_depth; nxt = 0
281
+ for t in range(1, depths[-1] + 1):
282
+ h, stt = model.step_state(stt)
283
+ if nxt < len(depths) and t == depths[nxt]:
284
+ ces.append(F.cross_entropy(model.logits(model.read(h[:M], last[:M])).float(), tgt[:M], reduction="none"))
285
+ stt = {k: (v[M:] if torch.is_tensor(v) else v) for k, v in stt.items()}; last = last[M:]; tgt = tgt[M:]; nxt += 1
286
+ return torch.cat(ces)
287
+
288
+
289
+ def loss_fn_prefix(model, tok, last, tgt, tgt1, n_loops):
290
+ """Arm D + PREFIX-ANSWER supervision. Every prefix of a shallow query is itself a shallow query whose answer is in the allowed data:
291
+ the prefix [x, r1] of a d2 row is the ATOMIC fact r1(x). All rows run n_loops loops; the prefix of depth p (read at relation position p)
292
+ is supervised at EVERY loop t >= p: position 1 -> r1(x) at loops 1..n, and for d2 rows position 2 -> r2(r1(x)) at loops 2..n. So a finished
293
+ position must hold its entity in the read-out format and KEEP it while later loops run. No label deeper than 2 hops, no 'not ready'
294
+ label, no loop index. Per-row loss = mean over its supervised (position, loop) cells; batch loss = mean over the 128 rows."""
295
+ stt = model.embed_state(tok); n1 = N_AT + 2 * N_W2SRC; one = torch.ones_like(last); two = 2 * one; is2 = (last == 2).float()
296
+ c1, c2 = [], []
297
+ for t in range(1, n_loops + 1):
298
+ h, stt = model.step_state(stt)
299
+ c1.append(F.cross_entropy(model.logits(model.read(h, one)).float(), tgt1, reduction="none"))
300
+ if t >= 2:
301
+ c2.append(F.cross_entropy(model.logits(model.read(h, two.clamp(max=tok.shape[1] - 1))).float(), tgt, reduction="none"))
302
+ c1 = torch.stack(c1); c2 = torch.stack(c2) # (n, N), (n - 1, N)
303
+ per_row = (c1.sum(0) + is2 * c2.sum(0)) / (n_loops + is2 * (n_loops - 1))
304
+ st = dict(ce_atomic=c1[0, :N_AT].mean(), ce_w2=c1[0, N_AT:n1].mean(), ce_d2=c2[0, n1:].mean(), ce_pos1_first=c1[0].mean(), ce_pos1_last=c1[-1].mean(),
305
+ ce_pos2_first=c2[0, n1:].mean(), ce_pos2_last=c2[-1, n1:].mean())
306
+ return per_row.mean(), st
307
+
308
+
309
+ def loss_fn_joint(model, tok, pos_a, pos_b, dep_a, dep_b, tgt_a, tgt_b):
310
+ """Arm D with JOINT rendering. tok (S, 6): [query a][query b] right-padded; every query is atomic (depth 1) or d2 (depth 2) and is read
311
+ at ITS OWN last token after ITS OWN depth in loops (R = d per query). Semantic depth stays <= 2; what changes is that a program no
312
+ longer starts at position 0 and that earlier, unrelated entity states are present in the context."""
313
+ h1 = model.step(model.embed(tok)); h2 = model.step(h1); ces = []
314
+ for pos, dep, tgt in ((pos_a, dep_a, tgt_a), (pos_b, dep_b, tgt_b)):
315
+ c1 = F.cross_entropy(model.logits(model.read(h1, pos)).float(), tgt, reduction="none")
316
+ c2 = F.cross_entropy(model.logits(model.read(h2, pos)).float(), tgt, reduction="none")
317
+ ces.append(torch.where(dep == 1, c1, c2))
318
+ ce = torch.cat(ces); dep = torch.cat([dep_a, dep_b]); d1 = (dep == 1).float(); d2 = 1.0 - d1
319
+ st = dict(ce_atomic=(ce * d1).sum() / d1.sum().clamp(min=1), ce_w2=(ce * d1).sum() / d1.sum().clamp(min=1), ce_d2=(ce * d2).sum() / d2.sum().clamp(min=1),
320
+ ce_second_slot=ces[1].mean())
321
+ return ce.mean(), st
322
+
323
+
324
+ # ----------------------------------------------------------------------------------------------------------- evaluation
325
+ class EvalSet:
326
+ def __init__(self, items, e0, r0, device):
327
+ """items: dict(x, rels, gold[, cat, d, id]); grouped by depth (equal length -> no padding inside a group)."""
328
+ self.items = items; self.n = len(items); self.groups = []
329
+ by_d = collections.defaultdict(list)
330
+ for i, it in enumerate(items):
331
+ by_d[len(it["rels"])].append(i)
332
+ for d, idx in sorted(by_d.items()):
333
+ tok = torch.tensor([[e0 + items[i]["x"]] + [r0 + r for r in items[i]["rels"]] for i in idx], device=device)
334
+ self.groups.append((d, np.array(idx), tok, torch.tensor([e0 + items[i]["gold"] for i in idx], device=device)))
335
+
336
+ def batches(self, tokens_per_batch=60000):
337
+ for d, idx, tok, gold in self.groups:
338
+ bs = max(32, tokens_per_batch // (d + 1))
339
+ for s in range(0, len(idx), bs):
340
+ yield d, idx[s:s + bs], tok[s:s + bs], gold[s:s + bs]
341
+
342
+
343
+ @torch.no_grad()
344
+ def eval_fixed(model, ES, budgets):
345
+ """budgets: list of ints (the same R for every depth) and / or the string 'd' (R = depth). -> {budget: correct (N,) bool}"""
346
+ def loops(b, d):
347
+ """budget spec: int = the same R for every depth; 'd' = depth; 'd+K' = depth + K; 'Kd' = K x depth"""
348
+ if isinstance(b, int):
349
+ return b
350
+ if b == "d":
351
+ return d
352
+ return d + int(b[2:]) if b.startswith("d+") else int(b[:-1]) * d
353
+ out = {b: np.zeros(ES.n, dtype=bool) for b in budgets}
354
+ for d, idx, tok, gold in ES.batches():
355
+ want = sorted({loops(b, d) for b in budgets}); last = torch.full((len(idx),), d, device=tok.device)
356
+ stt = model.embed_state(tok); snap = {}
357
+ for t in range(1, want[-1] + 1):
358
+ h, stt = model.step_state(stt)
359
+ if t in want:
360
+ snap[t] = (model.logits(model.read(h, last)).argmax(-1) == gold).cpu().numpy()
361
+ for b in budgets:
362
+ out[b][idx] = snap[loops(b, d)]
363
+ return out
364
+
365
+
366
+ @torch.no_grad()
367
+ def eval_halt(model, ES, b_eval, thr=0.5, sample_seed=None, rule="head"):
368
+ """Per-sample stopping with an active set that shrinks. rule 'head': stop head (p >= thr, or Bernoulli(p) if sample_seed);
369
+ rule 'kl': the paper's inference rule (KL(a_t || a_{t-1}) < 0.01 and entropy(a_t) < 3.0, at least 2 loops).
370
+ -> correct (N,), T (N,) loop at which the sample stopped (b_eval if it never did), stopped (N,) bool."""
371
+ correct = np.zeros(ES.n, dtype=bool); T = np.full(ES.n, b_eval, dtype=np.int32); stopped = np.zeros(ES.n, dtype=bool)
372
+ gen = torch.Generator(device=ES.groups[0][2].device).manual_seed(sample_seed) if sample_seed is not None else None
373
+ for d, idx, tok, gold in ES.batches():
374
+ stt = model.embed_state(tok); act = torch.arange(len(idx), device=tok.device); last = torch.full((len(idx),), d, device=tok.device); prev_lp = None
375
+ for t in range(1, b_eval + 1):
376
+ h, stt = model.step_state(stt); z = model.read(h, last[:len(act)]); lg = model.logits(z)
377
+ if rule == "head":
378
+ p = torch.sigmoid(model.stop_logit(z)); stop = (torch.rand(p.shape, device=p.device, generator=gen) < p) if gen is not None else (p >= thr)
379
+ else:
380
+ lp = F.log_softmax(lg.float(), -1); ent = -(lp.exp() * lp).sum(-1)
381
+ stop = torch.zeros(len(act), dtype=torch.bool, device=tok.device) if prev_lp is None else (((lp.exp() * (lp - prev_lp)).sum(-1) < 0.01) & (ent < 3.0))
382
+ if t == b_eval:
383
+ ok = (lg.argmax(-1) == gold[act]).cpu().numpy(); correct[idx[act.cpu().numpy()]] = ok # recorded, but counted as timeout
384
+ break
385
+ if bool(stop.any()):
386
+ a = act[stop].cpu().numpy(); correct[idx[a]] = (lg[stop].argmax(-1) == gold[act[stop]]).cpu().numpy(); T[idx[a]] = t; stopped[idx[a]] = True
387
+ keep = ~stop; act = act[keep]; stt = {k: (v[keep] if torch.is_tensor(v) else v) for k, v in stt.items()}
388
+ if rule == "kl":
389
+ lp = lp[keep]
390
+ if len(act) == 0:
391
+ break
392
+ if rule == "kl":
393
+ prev_lp = lp
394
+ return correct, T, stopped
395
+
396
+
397
+ def by_cell(items, **arrays):
398
+ agg = collections.defaultdict(list)
399
+ for i, it in enumerate(items):
400
+ agg[f"{it['cat']}|d{it['d']}"].append(i)
401
+ out = {}
402
+ for k, ix in sorted(agg.items()):
403
+ ix = np.array(ix); rec = dict(n=len(ix))
404
+ for name, arr in arrays.items():
405
+ rec[name] = float(np.mean(arr[ix]))
406
+ if "T" in arrays:
407
+ rec["T_p10_p50_p90"] = [float(x) for x in np.percentile(arrays["T"][ix], [10, 50, 90])]
408
+ out[k] = rec
409
+ return out
410
+
411
+
412
+ # ------------------------------------------------------------------------------------------------------------------ main
413
+ def main():
414
+ ap = argparse.ArgumentParser()
415
+ ap.add_argument("--data_dir", required=True); ap.add_argument("--atomic", required=True); ap.add_argument("--save_dir", required=True)
416
+ ap.add_argument("--arm", choices=["D", "H", "P"], required=True); ap.add_argument("--seed", type=int, default=1)
417
+ ap.add_argument("--d_model", type=int, default=768); ap.add_argument("--n_head", type=int, default=12); ap.add_argument("--n_layer", type=int, default=4)
418
+ ap.add_argument("--reembed", type=int, default=0, help="arm D only (1 = re-embed, own stream keeps its state; 2 = stateless positions, previous position seen as token + entity; 3 = stateless, previous position seen as entity only; 4 / 5 = bottleneck ablations of 3 / 2: the hidden vector is handed on directly, no vocabulary read-out): between loops every relation position is rebuilt as token embedding + the embedding-space "
419
+ "expectation of its own entity prediction (computed entities are handed on in the format of raw entity tokens)")
420
+ ap.add_argument("--prefix_sup", type=int, default=0, help="arm D only: supervise the prefix answer at every relation position (position 1 = the atomic fact) "
421
+ "at every loop >= its depth; all rows run 2 + --extra_loops loops")
422
+ ap.add_argument("--extra_loops", type=int, default=2)
423
+ ap.add_argument("--joint", type=int, default=0, help="arm D only: render the 128 queries of an update as 64 sequences of TWO independent queries "
424
+ "(the 16 w2 sources stay together; the other 96 queries are paired at random), each answered at its own last token")
425
+ ap.add_argument("--attn_window", type=int, default=0, help="0 = full causal attention (paper-aligned); W > 0 = local causal attention: every token sees itself "
426
+ "and the previous W tokens only (an architectural locality prior; no supervision, pointer or loop index is added)")
427
+ ap.add_argument("--pos", choices=["nope", "rope"], default="nope", help="nope = paper-aligned main setting; rope = relative positions (RoPE base 100), a labelled variant")
428
+ ap.add_argument("--lr", type=float, default=1e-4); ap.add_argument("--weight_decay", type=float, default=0.01); ap.add_argument("--warmup", type=int, default=2000)
429
+ ap.add_argument("--clip", type=float, default=1.0); ap.add_argument("--updates", type=int, default=1000000)
430
+ ap.add_argument("--lr_decay", choices=["none", "linear", "cosine"], default="none", help="linear: after the warm-up the lr falls linearly to 0 at --updates (the constant schedule "
431
+ "lets Adam blow up once the loss sits at the fp32 floor: loss 5e-8 for 100k updates, then a spike back to chance)")
432
+ ap.add_argument("--adam_eps", type=float, default=1e-8)
433
+ ap.add_argument("--mix", default="", help="queries per update as 'atomic,w2_sources,d2' (default 32,16,64 = 128 queries); e.g. 128,0,0 = atomic-only control")
434
+ ap.add_argument("--deep", type=int, default=0, help="arm D: ALSO train on chains of depth 3..DEEP (0 = off). The chains walk the TRAINED relation graph, heads are random, "
435
+ "chains of the test matrix are excluded; read-out after loop d like every other row")
436
+ ap.add_argument("--anchor", type=float, default=0.0, help="arm D, standard loop: weight of the ANCHORED-STATE loss (see loss_anchor); 0 = off")
437
+ ap.add_argument("--consist_w", type=float, default=0.0, help="weight of the one-step consistency loss on constructed deep states (loss_consist); 0 = off")
438
+ ap.add_argument("--consist_depth", type=int, default=32); ap.add_argument("--consist_rows", type=int, default=32); ap.add_argument("--consist_start", type=int, default=20000)
439
+ ap.add_argument("--format_w", type=float, default=0.0, help="weight of the label-free FORMAT loss on random chains (loss_format); 0 = off")
440
+ ap.add_argument("--format_start", type=int, default=0, help="first update at which the format loss is applied (let the shallow task and the anchors settle first)")
441
+ ap.add_argument("--format_bounded", type=int, default=1, help="1 = symmetric relative error (bounded), 0 = the original relative squared error")
442
+ ap.add_argument("--format_depth", type=int, default=32, help="depth of the random unlabelled chains of the format loss"); ap.add_argument("--format_rows", type=int, default=16, help="such chains per update")
443
+ ap.add_argument("--anchor_extra", type=int, default=2, help="with --anchor: loops run beyond the depth of a row (answer and anchors supervised there too)")
444
+ ap.add_argument("--delay_max", type=int, default=0, help="arm D, plain loop: random start delay k ~ U{0..DELAY_MAX} per update (the head entity arrives after k loops; read at k + d)")
445
+ ap.add_argument("--delay_start", type=int, default=0, help="first update at which --delay_max is active (k = 0 before)")
446
+ ap.add_argument("--persist", type=int, default=0, help="arm D, plain loop: supervise the answer after every loop d .. d + PERSIST (the answer has to stay)")
447
+ ap.add_argument("--pack", type=int, default=0, help="arm D, standard loop: WIDTH-n rendering -- pack the shallow queries of every update into causal sequences of about PACK tokens "
448
+ "(n independent queries side by side, each read at its own last relation token); 0 = one query per sequence. Evaluation is unchanged (single chains)")
449
+ ap.add_argument("--pack_prefix_max", type=int, default=0, help="with --pack: every packed sequence starts with k random ENTITY tokens, k ~ --offset_dist over 0..PACK_PREFIX_MAX "
450
+ "(host-side packing, dynamic number of sequences: the CUDA graph is switched off)")
451
+ ap.add_argument("--offset_dist", default="uniform", help="distribution of prefix lengths for --offset_max / --pack_prefix_max: uniform | boltz<b> (p ~ exp(b k / K), long prefixes favoured) | poisson<lam>")
452
+ ap.add_argument("--pack_what", choices=["all", "atomic"], default="all", help="all = d1 and d2 queries are packed together; atomic = only the 64 atomic queries are packed (pure w_n), d2 rows stay alone")
453
+ ap.add_argument("--offset_fill", choices=["pad", "entity"], default="pad", help="filler tokens of --offset_max: PAD, or random ENTITY tokens (use with --anchor)")
454
+ ap.add_argument("--offset_max", type=int, default=0, help="DATA-RENDERING remedy for position-calibrated addressing (arm D, plain protocol): every TRAINING row is shifted right by a random "
455
+ "k ~ U{0..OFFSET_MAX} (k filler tokens in front), so large absolute positions are seen with shallow chains; evaluation is unchanged (no fillers)")
456
+ ap.add_argument("--deep_list", default="", help="explicit list of extra training depths, e.g. 3,5,6,10,12,20,24,32 (overrides the range 3..DEEP; untrained depths inside the range test "
457
+ "interpolation, depths beyond it extrapolation)")
458
+ ap.add_argument("--deep_n", type=int, default=10000, help="number of distinct training chains per extra depth (the data budget)")
459
+ ap.add_argument("--deep_rows", type=int, default=16, help="rows per extra depth per update (added to the 128 shallow queries; the loss stays a per-query mean)")
460
+ ap.add_argument("--deep_labels", choices=["truth", "self"], default="truth", help="truth = gold answers; self = NO deep labels: the target is what the CURRENT model returns "
461
+ "when it is called one atomic step at a time by an external loop (its own shallow skill, composed outside), i.e. self-distillation of external iteration into the loop")
462
+ ap.add_argument("--deep_start", type=int, default=0, help="first update at which the deep rows enter the loss")
463
+ ap.add_argument("--deep_seed", type=int, default=20260921)
464
+ ap.add_argument("--nr_init", type=int, default=0, help="stateless loops (--reembed 2|3): relation positions start from a learned 'not ready' vector instead of their token embedding")
465
+ ap.add_argument("--ready_state", type=int, default=0, help="arm H on a stateless loop: the stop head is read at every relation position and the state handed on is "
466
+ "rho * entity expectation + (1 - rho) * 'not ready' (implies --nr_init); the query halts on rho at its last position")
467
+ ap.add_argument("--untied", type=int, default=0, help="arm D only: 1 = VANILLA control, the --n_layer blocks are different layers applied once each (no weight sharing, no loop)")
468
+ ap.add_argument("--unroll", type=int, default=0, help="arm D, TIED model: train with the same fixed-depth protocol as the vanilla control -- L loops of the shared block(s), "
469
+ "answer read according to --untied_read; the main test protocol is then R = L for final / all and R = d for at_d (matched loop-vs-vanilla comparison)")
470
+ ap.add_argument("--untied_group", type=int, default=1, help="with --untied: layers per step (= blocks per loop of the tied model under comparison); the control then has n_layer / untied_group steps, "
471
+ "and --anchor / --consist_w / --offset_max treat one step as one loop (every layer group receives the consistency states)")
472
+ ap.add_argument("--untied_read", choices=["at_d", "final", "all"], default="final", help="where the answer is read / supervised in the vanilla control (see loss_fn_untied)")
473
+ ap.add_argument("--eval_ckpt", default="", help="with --eval_only: evaluate this checkpoint file (e.g. best.pt) instead of last.pt"); ap.add_argument("--eval_tag", default="")
474
+ ap.add_argument("--b_train", type=int, default=8, help="H: loops unrolled in training (not tied to the depth of the query)")
475
+ ap.add_argument("--b_eval", type=int, default=256, help="H: global loop cap at inference, the same for every depth; reaching it is a timeout")
476
+ ap.add_argument("--lam", type=float, default=0.0, help="H: compute cost weight on E[T]"); ap.add_argument("--stop_thr", type=float, default=0.5)
477
+ ap.add_argument("--stop_bias", type=float, default=0.0, help="initial bias of the stop head")
478
+ ap.add_argument("--eval_every", type=int, default=10000); ap.add_argument("--monitor_per_cell", type=int, default=100)
479
+ ap.add_argument("--cuda_graph", type=int, default=0); ap.add_argument("--resume", action="store_true"); ap.add_argument("--eval_only", action="store_true")
480
+ ap.add_argument("--train_precision", choices=["fp32", "tf32", "bf16"], default="fp32",
481
+ help="matmul precision of the TRAINING step only (tf32 = TensorFloat-32 matmuls, bf16 = autocast); every evaluation runs in strict fp32")
482
+ ap.add_argument("--ckpt_every", type=int, default=0, help="also save last.pt every N updates (no evaluation) -- for chained short jobs, so a walltime kill loses at most N updates")
483
+ ap.add_argument("--stop_after", type=int, default=0); ap.add_argument("--smoke", type=int, default=0)
484
+ ap.add_argument("--cs_mode", choices=CS_MODES, default="full",
485
+ help="wait_only retains root CS; wait_no_root selects only wait CS, keeping the all-nonfrontier denominator")
486
+ ap.add_argument("--cs_noise_scale", type=float, default=0.0,
487
+ help="t0_wait_denoise: max nearest-distance fraction for nearest/isotropic (preregistered 0.2); fixed embedding-norm fraction for legacy embedding kind")
488
+ ap.add_argument("--cs_noise_kind", choices=["embedding", "nearest", "isotropic"], default="embedding",
489
+ help="nearest/isotropic: U[0, cs_noise_scale) times nearest nonself embedding distance")
490
+ ap.add_argument("--cs_noise_scope", choices=["all", "wait"], default="wait",
491
+ help="noise all input positions, or only waiting positions >=2; loss always only wait")
492
+ ap.add_argument("--anchor_mode", choices=["off", "full", "execute_only", "wait_only", "hold_only"], default="full")
493
+ ap.add_argument("--final_n", type=int, default=5000)
494
+ ap.add_argument("--diagnostic_n", type=int, default=100)
495
+ ap.add_argument("--rope_base", type=float, default=100.0)
496
+ ap.add_argument("--max_pos", type=int, default=512)
497
+ args = ap.parse_args()
498
+ assert math.isfinite(args.rope_base) and args.rope_base > 1 and args.max_pos >= args.offset_max + 3
499
+ validate_simple_cs_config(args.cs_mode, args.consist_depth, args.consist_rows, args.cs_noise_scale,
500
+ args.cs_noise_kind, args.cs_noise_scope)
501
+ if args.cs_mode in SIMPLE_CS_MODES:
502
+ assert args.consist_w > 0 and args.anchor_mode == "full", "new simple-CS modes keep full shallow anchor and active CS"
503
+ assert args.arm == "D" and args.anchor == 1 and args.anchor_extra == 2
504
+ assert not any([args.pack, args.deep, args.joint, args.untied, args.unroll, args.reembed, args.delay_max, args.persist, args.format_w])
505
+ stop_requested = [False]
506
+ in_final_eval = [False]
507
+ def request_stop(signum, frame):
508
+ stop_requested[0] = True
509
+ if in_final_eval[0]:
510
+ sys.exit(75)
511
+ signal.signal(signal.SIGTERM, request_stop)
512
+ signal.signal(signal.SIGUSR1, request_stop)
513
+ if args.mix:
514
+ global N_AT, N_W2SRC, N_D2, N_ROWS
515
+ N_AT, N_W2SRC, N_D2 = (int(v) for v in args.mix.split(",")); N_ROWS = N_AT + 2 * N_W2SRC + N_D2
516
+ assert arm_ok(args), "--mix is implemented for arm D without --joint / --prefix_sup"
517
+ arm = args.arm; device = torch.device("cuda" if torch.cuda.is_available() else "cpu"); os.makedirs(args.save_dir, exist_ok=True)
518
+ torch.manual_seed(args.seed); np.random.seed(args.seed)
519
+
520
+ meta = load_json(args.data_dir, "meta"); rel = meta["rel_map"]; vocab = load_json(args.data_dir, "vocab"); t2i = {t: i for i, t in enumerate(vocab)}
521
+ e0, r0 = t2i["<e_0>"], t2i["<r_0>"]; PAD = t2i["<pad>"]; E = meta["num_entities"]
522
+ atomic = json.load(open(args.atomic)); train_w2 = load_json(args.data_dir, "train_w2"); train_d2 = load_json(args.data_dir, "train_d2")
523
+ assert all(a["y"] == rel[a["r"]][a["x"]] for a in atomic) and len(atomic) == E * meta["num_relations"]
524
+ val_d2 = load_json(args.data_dir, "val_d2"); tests = load_json(args.data_dir, "tests")
525
+ # source tables on the device: [x, r, s or PAD, last index, gold] (w2 sources become two independent atomic queries)
526
+ T_at = torch.tensor([[e0 + a["x"], r0 + a["r"], PAD, 1, e0 + a["y"], e0 + a["y"]] for a in atomic], device=device)
527
+ T_w2 = torch.tensor([[[e0 + x, r0 + r, PAD, 1, e0 + rel[r][x], e0 + rel[r][x]] for (x, r) in (q["q1"], q["q2"])] for q in train_w2], device=device)
528
+ T_d2 = torch.tensor([[e0 + q["x"], r0 + q["r"], r0 + q["s"], 2, e0 + rel[q["s"]][rel[q["r"]][q["x"]]], e0 + rel[q["r"]][q["x"]]] for q in train_d2], device=device)
529
+
530
+ DEPTHS_X = list(range(3, args.deep + 1)) if args.deep >= 3 else []
531
+ if args.deep_list:
532
+ DEPTHS_X = sorted({int(v) for v in args.deep_list.split(",")}); assert DEPTHS_X[0] >= 3; args.deep = DEPTHS_X[-1]
533
+ assert not DEPTHS_X or (arm == "D" and not (args.joint or args.prefix_sup or args.untied or args.unroll)), "--deep is implemented for the plain arm D protocol (R = d)"
534
+ deep_pool = {}
535
+ if DEPTHS_X:
536
+ succ = collections.defaultdict(list)
537
+ for a, b in meta["train_edges"]:
538
+ succ[a].append(b)
539
+ heads_ok = sorted(succ); test_keys = {(t["x"], tuple(t["rels"])) for t in tests if t["d"] in DEPTHS_X}; relm = np.array(rel)
540
+ for d in DEPTHS_X:
541
+ rg = np.random.default_rng([args.deep_seed, d]); seen = set(); rows = []
542
+ while len(rows) < args.deep_n:
543
+ x = int(rg.integers(0, E)); rs = [int(rg.choice(heads_ok))]
544
+ while len(rs) < d:
545
+ rs.append(int(rg.choice(succ[rs[-1]])))
546
+ key = (x, tuple(rs))
547
+ if key in seen or key in test_keys:
548
+ continue
549
+ seen.add(key); y = x
550
+ for r in rs:
551
+ y = int(relm[r][y])
552
+ rows.append([e0 + x] + [r0 + r for r in rs] + [PAD] * (args.deep - d) + [d, e0 + y])
553
+ deep_pool[d] = torch.tensor(rows, device=device) # [x, r1..rd, PAD.., last index, gold]
554
+ model = LoopGPT(len(vocab), args.d_model, args.n_head, args.n_layer, pos=args.pos, rope_base=args.rope_base, max_pos=args.max_pos, attn_window=args.attn_window, reembed=int(args.reembed),
555
+ ent_range=(e0, e0 + E), rel_range=(r0, r0 + meta["num_relations"]), untied=bool(args.untied), untied_group=args.untied_group,
556
+ nr_init=bool(args.nr_init), ready_state=bool(args.ready_state)); model.seeded_init(args.seed)
557
+ assert not (args.untied or args.unroll) or (arm == "D" and not args.joint and not args.prefix_sup), "--untied / --unroll are implemented for arm D without --joint / --prefix_sup"
558
+ FIXED_L = (args.n_layer // args.untied_group) if args.untied else args.unroll # steps of the fixed-depth protocols; 0 = the ordinary loop protocol (R = d)
559
+ with torch.no_grad():
560
+ model.stop.bias.fill_(args.stop_bias)
561
+ model.to(device); assert model.wte.weight.data_ptr() == model.wte.weight.data_ptr()
562
+ hsh = lambda names: hashlib.sha256(b"".join(p.detach().cpu().numpy().tobytes() for n, p in sorted(model.named_parameters()) if n in names)).hexdigest()[:16]
563
+ n_params = sum(p.numel() for p in model.parameters()); cap = args.smoke or None
564
+
565
+ for t in tests:
566
+ t["gold"] = t["states"][-1]
567
+ per_cell = collections.defaultdict(list)
568
+ for t in tests:
569
+ per_cell[(t["cat"], t["d"])].append(t)
570
+ first = lambda n: [t for k in per_cell for t in per_cell[k][:min(n, cap or n)]]
571
+ mon_items = first(args.monitor_per_cell); MON = EvalSet(mon_items, e0, r0, device)
572
+ at_items = [dict(x=a["x"], rels=[a["r"]], gold=a["y"], cat="atomic", d=1) for a in atomic][:cap]; AT = EvalSet(at_items, e0, r0, device)
573
+ mk2 = lambda rows, cat: [dict(x=q["x"], rels=[q["r"], q["s"]], gold=rel[q["s"]][rel[q["r"]][q["x"]]], cat=cat, d=2) for q in rows][:cap]
574
+ VD2 = EvalSet(mk2(val_d2, "val_d2"), e0, r0, device); FIT = EvalSet(mk2(train_d2[::10], "train_d2"), e0, r0, device)
575
+ MAIN_FIXED = 8 # arm P: pre-registered common budget (its largest training R)
576
+
577
+ def main_protocol(ES):
578
+ """the arm's own pre-registered inference protocol -> dict(correct[, T, stopped])"""
579
+ if arm == "D":
580
+ b = FIXED_L if (FIXED_L and args.untied_read != "at_d") else "d" # fixed-depth protocols read after all L steps, whatever the depth
581
+ return dict(correct=eval_fixed(model, ES, [b])[b])
582
+ if arm == "P":
583
+ return dict(correct=eval_fixed(model, ES, [MAIN_FIXED])[MAIN_FIXED])
584
+ c, T, s = eval_halt(model, ES, args.b_eval, args.stop_thr)
585
+ dep = np.array([len(it["rels"]) for it in ES.items]) # 'one loop = one atomic operation' predicts a stop exactly at loop d
586
+ return dict(correct=c & s, T=T.astype(np.float64), stopped=s, timeout=~s, wrong_stop=s & ~c, correct_at_stop_or_cap=c,
587
+ T_eq_d=s & (T == dep), T_lt_d=s & (T < dep), T_gt_d=s & (T > dep))
588
+
589
+ def shallow():
590
+ out = {}
591
+ for name, ES in (("atomic", AT), ("val_d2", VD2), ("train_d2_fit", FIT)):
592
+ r = main_protocol(ES); out[name] = {k: float(np.mean(v)) for k, v in r.items()}
593
+ return out
594
+
595
+ def final_eval(tag="", ckpt_update=None):
596
+ torch.backends.cuda.matmul.allow_tf32 = False
597
+ torch.backends.cudnn.allow_tf32 = False
598
+ sfx = f"_{tag}" if tag else ""; path = os.path.join(args.save_dir, f"final_eval{sfx}.json")
599
+ if os.path.exists(path):
600
+ existing = json.load(open(path))
601
+ assert existing["checkpoint_update"] == ckpt_update and existing.get("main"), "Existing eval has incompatible checkpoint"
602
+ print(f"final_eval{sfx}.json exists, skipping"); return
603
+ if stop_requested[0]: sys.exit(75)
604
+ in_final_eval[0] = True
605
+ model.eval(); t0 = time.time(); res = dict(arm=arm, seed=args.seed, lam=args.lam, b_eval=args.b_eval, stop_thr=args.stop_thr, tag=tag, checkpoint_update=ckpt_update)
606
+ items = first(args.final_n); FULL = EvalSet(items, e0, r0, device); raw = dict(ids=np.array([t["id"] for t in items]))
607
+ r = main_protocol(FULL); res["main"] = by_cell(items, **r); raw.update({f"main_{k}": v for k, v in r.items()})
608
+ res["shallow"] = shallow()
609
+ # diagnostics on the SAME weights: R = d for every arm, the common fixed-budget grid on the first 500 per cell, the paper's KL / entropy stop rule
610
+ rd = r["correct"] if (arm == "D" and not (FIXED_L and args.untied_read != "at_d")) else eval_fixed(model, FULL, ["d"])["d"]; res["forced_R_eq_d"] = by_cell(items, correct=rd); raw["forced_R_eq_d"] = rd
611
+ g_items = first(args.diagnostic_n); G = EvalSet(g_items, e0, r0, device)
612
+ rg = eval_fixed(model, G, ["d", "d+2", "2d"])
613
+ res["matched_budget_diagnostics"] = {str(b): by_cell(g_items, correct=rg[b]) for b in rg}
614
+ res["eval_seconds"] = time.time() - t0
615
+ raw_path = os.path.join(args.save_dir, f"final_answers{sfx}.npz")
616
+ with open(raw_path + ".tmp", "wb") as output: np.savez_compressed(output, **raw)
617
+ os.replace(raw_path + ".tmp", raw_path)
618
+ with open(path + ".tmp", "w") as output: json.dump(res, output, indent=1)
619
+ os.replace(path + ".tmp", path)
620
+ in_final_eval[0] = False
621
+ m = res["main"]; print(f"[final main{sfx}] " + " ".join(f"{k}:{m[k]['correct']:.3f}" for k in m if not k.startswith("all_seen")), flush=True)
622
+
623
+ capturable = device.type == "cuda"
624
+ opt = torch.optim.AdamW(model.parameters(), lr=(torch.tensor(args.lr, device=device) if capturable else args.lr), betas=(0.9, 0.999), eps=args.adam_eps,
625
+ weight_decay=args.weight_decay, capturable=capturable)
626
+ s_at = Stream(len(atomic), args.seed, "atomic"); s_w2 = Stream(len(train_w2), args.seed, "w2"); s_d2 = Stream(len(train_d2), args.seed, "d2")
627
+ r_rng = named_rng(args.seed, "halt_R")
628
+ ctr = dict(atomic=0, w2_sources=0, w2_queries=0, d2=0, row_loops=0, R_hist={})
629
+ start = 1; best = None
630
+ P = lambda f: os.path.join(args.save_dir, f); last_path = P("last.pt")
631
+ resume_state = None
632
+ if (args.resume or args.eval_only) and os.path.exists(last_path):
633
+ resume_state = torch.load(last_path, map_location=device, weights_only=False)
634
+ model.load_state_dict(resume_state["model"]); s_at.load(resume_state["s_at"]); s_w2.load(resume_state["s_w2"]); s_d2.load(resume_state["s_d2"])
635
+ r_rng.bit_generator.state = resume_state["r_rng"]; ctr = resume_state["ctr"]; start = resume_state["update"] + 1; best = resume_state["best"]
636
+ print(f"resumed from update {resume_state['update']}")
637
+ elif not args.eval_only:
638
+ for f in ("metrics.jsonl", "train_log.jsonl"):
639
+ open(P(f), "w").close()
640
+ src = hashlib.sha256(b"".join(open(os.path.join(os.path.dirname(os.path.abspath(__file__)), f), "rb").read() for f in ("train_halt.py", "model/loop_gpt.py", "objectives.py", "streams.py"))).hexdigest()[:16]
641
+ json.dump(dict(theory_args=vars(args), arm=arm, seed=args.seed, n_params=n_params, stop_head_params=model.stop.weight.numel() + 1, backbone_init_hash=hsh(set(model.backbone_names())),
642
+ model=dict(d_model=args.d_model, n_head=args.n_head, n_layer=args.n_layer, position=args.pos, rope_base=args.rope_base, max_pos=args.max_pos, attn_window=args.attn_window, reembed=int(args.reembed), nr_init=bool(args.nr_init or args.ready_state), ready_state=bool(args.ready_state), untied=(dict(layers=args.n_layer, group=args.untied_group, steps=args.n_layer // args.untied_group, read=args.untied_read) if args.untied else None), tied_unroll=(dict(loops=args.unroll, read=args.untied_read) if args.unroll and not args.untied else None), tied_lm_head=True, residual_out_proj_init=0.0, dropout=0.0, ln_eps=1e-5, act="gelu_tanh", readout="last valid input position"),
643
+ batch=dict(queries=N_ROWS, atomic=N_AT, w2_sources=N_W2SRC, w2_rendering=("JOINT: 64 sequences of two independent queries per update, each read at its own last token" if args.joint else "two independent atomic queries per source"), d2=N_D2, loss="mean of 128 final-answer CEs"),
644
+ optimizer=dict(name="AdamW", lr=args.lr, weight_decay=args.weight_decay, warmup=args.warmup, schedule={"none": "linear warm-up then constant", "linear": "linear warm-up, then linear decay to 0 at the last update", "cosine": "linear warm-up, then cosine decay to 0 at the last update"}[args.lr_decay], adam_eps=args.adam_eps, clip=args.clip, label_smoothing=0.0, precision=args.train_precision, evaluation_precision="fp32"),
645
+ updates=args.updates, halting=dict(b_train=args.b_train, b_eval=args.b_eval, lam=args.lam, stop_thr=args.stop_thr, stop_bias=args.stop_bias, tail="all remaining mass at the last unrolled loop (truncation mass)") if arm == "H" else None,
646
+ P_budget="R ~ clip(Poisson(4), 2, 8) per batch; main test budget R = 8" if arm == "P" else None,
647
+ pack=(dict(tokens_per_sequence=args.pack, what=args.pack_what, prefix_max=args.pack_prefix_max, prefix_dist=args.offset_dist, note="n independent shallow queries side by side in one causal sequence, no separator, each read at its own last relation token") if args.pack else None),
648
+ offset_max=args.offset_max, offset_dist=args.offset_dist, delay_max=args.delay_max, delay_start=args.delay_start, persist=args.persist, anchor=(dict(weight=(0.0 if args.anchor_mode == "off" else args.anchor), mode=args.anchor_mode, extra_loops=args.anchor_extra) if args.anchor else None),
649
+ consist_loss=(dict(mode=args.cs_mode, components=list(CS_COMPONENTS[args.cs_mode]), root_shared=("root" in CS_COMPONENTS[args.cs_mode]), weight=args.consist_w, depth=args.consist_depth, rows=args.consist_rows, start=args.consist_start, **simple_cs_metadata(args.cs_mode, args.cs_noise_scale, args.cs_noise_kind, args.cs_noise_scope)) if args.consist_w else None),
650
+ format_loss=(dict(weight=args.format_w, depth=args.format_depth, rows=args.format_rows, start=args.format_start, bounded=bool(args.format_bounded), note="label-free: waiting -> own token embedding, fired -> nearest entity embedding, random chains") if args.format_w else None),
651
+ deep=(dict(depths=DEPTHS_X, chains_per_depth=args.deep_n, rows_per_depth_per_update=args.deep_rows, labels=args.deep_labels, start=args.deep_start, pool_seed=args.deep_seed,
652
+ note="chains walk the trained relation graph; test-matrix chains excluded; read after loop d") if args.deep >= 3 else None),
653
+ prefix_supervision=(dict(loops=2 + args.extra_loops, rule="prefix of depth p supervised at every loop t >= p; position-1 target = atomic fact") if args.prefix_sup else None),
654
+ selection="final checkpoint is primary; best.pt = highest val_d2 accuracy under the arm's own protocol (earliest on ties); deep tests never used",
655
+ data_hashes=meta.get("hashes"), atomic_sha256=hashlib.sha256(open(args.atomic, "rb").read()).hexdigest(), source_hash=src, cuda_graph=bool(args.cuda_graph)),
656
+ open(P("manifest.json"), "w"), indent=1)
657
+ print(f"arm {arm} seed {args.seed} | params {n_params:,} | device {device} | cuda_graph {args.cuda_graph}", flush=True)
658
+ def eval_best():
659
+ bp = os.path.join(args.save_dir, "best.pt")
660
+ if os.path.exists(bp):
661
+ ck = torch.load(bp, map_location=device, weights_only=False); model.load_state_dict(ck["model"]); final_eval("best", ck.get("update"))
662
+
663
+ if args.eval_only:
664
+ if args.eval_ckpt:
665
+ ck = torch.load(args.eval_ckpt, map_location=device, weights_only=False); model.load_state_dict(ck["model"])
666
+ final_eval(args.eval_tag or os.path.splitext(os.path.basename(args.eval_ckpt))[0], ck.get("update"))
667
+ else:
668
+ final_eval(args.eval_tag, resume_state.get("update") if resume_state else None)
669
+ return
670
+
671
+ params = [p for p in model.parameters()]
672
+ OFF = args.offset_max; assert not OFF or (arm == "D" and not (args.joint or args.prefix_sup or args.unroll or args.reembed)), "--offset_max: plain arm D (or its vanilla control) only"
673
+ tok = torch.full((N_ROWS, 3 + OFF), PAD, dtype=torch.long, device=device); last = torch.zeros(N_ROWS, dtype=torch.long, device=device); tgt = torch.zeros(N_ROWS, dtype=torch.long, device=device)
674
+ off_rng = named_rng(args.seed, "halt_offset"); ar_s = torch.arange(3, device=device)[None]
675
+ tgt1 = torch.zeros(N_ROWS, dtype=torch.long, device=device); last[N_AT + 2 * N_W2SRC:] = 2; last[:N_AT + 2 * N_W2SRC] = 1
676
+
677
+ n_deep = args.deep_rows * len(DEPTHS_X); Ld = args.deep + 1
678
+ tok_d = torch.full((max(1, n_deep), max(3, Ld) + args.offset_max), PAD, dtype=torch.long, device=device); ar_d = torch.arange(max(3, Ld), device=device)[None]
679
+ if resume_state is not None and resume_state.get("off_rng"):
680
+ off_rng.bit_generator.state = resume_state["off_rng"]
681
+ assert not (OFF and args.deep_labels == "self"), "--offset_max with self-generated labels is not implemented"
682
+ if args.anchor and args.delay_max:
683
+ args.cuda_graph = 0 # one eager step per random delay instead of delay_max + 1 captured graphs
684
+ assert not args.anchor or (arm == "D" and not (args.persist or args.joint or args.prefix_sup or args.unroll or args.reembed or args.deep_labels == "self")), "--anchor: plain arm D (or its vanilla control)"
685
+ REL_T = torch.tensor(np.array(rel), device=device); pv_s = torch.zeros((N_ROWS, 3 + OFF), dtype=torch.long, device=device); pv_d = torch.zeros_like(tok_d)
686
+ off_s = torch.zeros(N_ROWS, dtype=torch.long, device=device); dep_s = torch.ones(N_ROWS, dtype=torch.long, device=device); off_d = torch.zeros(max(1, n_deep), dtype=torch.long, device=device)
687
+ dep_d = torch.ones(max(1, n_deep), dtype=torch.long, device=device)
688
+ PACK = args.pack; assert not PACK or (arm == "D" and not (OFF or args.deep or args.delay_max or args.persist or args.joint or args.prefix_sup or args.untied or args.unroll or args.reembed)), "--pack: plain arm D, shallow data"
689
+ n_pk = N_ROWS if args.pack_what == "all" else N_AT + 2 * N_W2SRC; tot_pk = 2 * (N_AT + 2 * N_W2SRC) + (3 * N_D2 if args.pack_what == "all" else 0)
690
+ PR, PW, PS = (-(-tot_pk // max(1, PACK)), PACK + 2, PACK // 2 + 2) if PACK else (1, 3, 1)
691
+ PPM = args.pack_prefix_max
692
+ if PPM:
693
+ assert PACK and 0 < PPM <= PACK - 3, "--pack_prefix_max needs --pack and must leave room for a d2 query"
694
+ PR = n_pk; args.cuda_graph = 0 # worst case one query per sequence; the number of sequences varies -> eager training step
695
+ pk_n = [PR]; pref_sample = offset_sampler(args.offset_dist, PPM if PPM else max(1, args.offset_max))
696
+ pk = dict(tok=torch.full((PR, PW), PAD, dtype=torch.long, device=device), valid=torch.zeros((PR, PW), dtype=torch.bool, device=device), pv=torch.zeros((PR, PW), dtype=torch.long, device=device),
697
+ fire=torch.zeros((PR, PW), dtype=torch.long, device=device), seg_last=torch.zeros((PR, PS), dtype=torch.long, device=device), seg_depth=torch.ones((PR, PS), dtype=torch.long, device=device),
698
+ seg_gold=torch.zeros((PR, PS), dtype=torch.long, device=device), seg_valid=torch.zeros((PR, PS), dtype=torch.bool, device=device))
699
+ pack_rng = named_rng(args.seed, "halt_pack"); ar_pk = torch.arange(n_pk, device=device)
700
+ if resume_state is not None and resume_state.get("pack_rng"):
701
+ pack_rng.bit_generator.state = resume_state["pack_rng"]
702
+
703
+ def fill_packed(rows):
704
+ if PPM:
705
+ qn = rows[:n_pk][torch.from_numpy(pack_rng.permutation(n_pk)).to(device)].cpu().numpy()
706
+ pk_n[0] = pack_rows_prefix(qn, PACK, PAD, pk, pref_sample(pack_rng, n_pk), pack_rng.integers(e0, e0 + E, (n_pk, PPM)))
707
+ else:
708
+ pack_rows(rows[:n_pk][torch.from_numpy(pack_rng.permutation(n_pk)).to(device)], PACK, PAD, pk)
709
+
710
+ CST = args.consist_w > 0
711
+ assert not CST or (arm == "D" and args.anchor > 0 and not (args.pack or args.deep or args.unroll or args.reembed)), "--consist_w: arm D with --anchor, plain rows (+ --offset_max)"
712
+ cst_tok = torch.zeros((max(1, args.consist_rows), args.consist_depth + 1), dtype=torch.long, device=device); cst_t = torch.zeros(max(1, args.consist_rows), dtype=torch.long, device=device)
713
+ cst_ent = torch.zeros_like(cst_tok); cst_rng = named_rng(args.seed, "halt_consist"); cst_w = torch.zeros((), device=device)
714
+ # CUDA warm-up/capture precedes actual sampling. Use valid fixed token
715
+ # IDs without consuming any formal RNG stream (zero may be padding).
716
+ if args.cs_mode in SIMPLE_CS_MODES:
717
+ cst_tok.fill_(r0)
718
+ if args.cs_mode == "t0_wait_denoise":
719
+ cst_tok[:, 0].fill_(e0)
720
+ if resume_state is not None and resume_state.get("cst_rng"):
721
+ cst_rng.bit_generator.state = resume_state["cst_rng"]
722
+ # Fixed buffers are captured, but draws happen only in the real update
723
+ # loop. Therefore warm-up/capture cannot consume the formal noise stream.
724
+ cst_noise = (torch.zeros((args.consist_rows, args.consist_depth + 1, args.d_model), device=device)
725
+ if args.cs_mode == "t0_wait_denoise" else None)
726
+ cst_alpha = (torch.zeros((args.consist_rows, args.consist_depth + 1, 1), device=device)
727
+ if args.cs_mode == "t0_wait_denoise" and args.cs_noise_kind in ("nearest", "isotropic") else None)
728
+ cst_relation_ids = torch.arange(r0, r0 + meta["num_relations"], device=device)
729
+ cst_candidate_ids = torch.cat((torch.arange(e0, e0 + E, device=device), cst_relation_ids))
730
+ cst_noise_rng = make_cs_noise_generator(args.seed, device) if cst_noise is not None else None
731
+ if resume_state is not None and cst_noise_rng is not None:
732
+ if "cst_noise_rng" not in resume_state or resume_state["cst_noise_rng"] is None:
733
+ raise ValueError("denoising resume checkpoint is missing its independent noise RNG state")
734
+ cst_noise_rng.set_state(resume_state["cst_noise_rng"].cpu())
735
+ FMT = args.format_w > 0
736
+ assert not FMT or (arm == "D" and args.anchor > 0 and not (args.pack or args.deep or args.untied or args.unroll or args.reembed)), "--format_w: arm D with --anchor, plain rows (+ --offset_max)"
737
+ fmt_tok = torch.zeros((max(1, args.format_rows), args.format_depth + 1), dtype=torch.long, device=device); fmt_rng = named_rng(args.seed, "halt_format"); fmt_w = torch.zeros((), device=device)
738
+ if resume_state is not None and resume_state.get("fmt_rng"):
739
+ fmt_rng.bit_generator.state = resume_state["fmt_rng"]
740
+ DP_AUG = bool(args.delay_max or args.persist)
741
+ assert not DP_AUG or (arm == "D" and not ((OFF and not args.anchor) or args.joint or args.prefix_sup or args.untied or args.unroll or args.reembed or args.deep_labels == "self")), "--delay_max / --persist: plain arm D"
742
+ delay_rng = named_rng(args.seed, "halt_delay")
743
+ if resume_state is not None and resume_state.get("delay_rng"):
744
+ delay_rng.bit_generator.state = resume_state["delay_rng"]
745
+ last_d = torch.ones(max(1, n_deep), dtype=torch.long, device=device)
746
+ tgt_d = torch.zeros(max(1, n_deep), dtype=torch.long, device=device); gold_d = torch.zeros(max(1, n_deep), dtype=torch.long, device=device)
747
+ deep_w = torch.zeros((), device=device) # 0 before --deep_start, 1 afterwards (a tensor, so the captured graph sees the switch)
748
+ s_deep = {d: Stream(args.deep_n, args.seed, f"deep{d}") for d in DEPTHS_X}
749
+ if resume_state is not None and resume_state.get("s_deep"):
750
+ for d in DEPTHS_X:
751
+ s_deep[d].load(resume_state["s_deep"][str(d)])
752
+
753
+ @torch.no_grad()
754
+ def self_labels():
755
+ """targets of the deep rows from the model's OWN atomic skill, composed by an external loop: y <- argmax model([y, r_j]) read after one loop, j = 1..d"""
756
+ was = model.training; model.eval(); y = tok_d[:, 0].clone(); one = torch.ones(n_deep, dtype=torch.long, device=device)
757
+ for j in range(1, Ld):
758
+ act = last_d >= j; q = torch.stack([y, tok_d[:, j].clamp(min=0)], 1); o, _ = model.step_state(model.embed_state(q))
759
+ W_ = model.wte.weight[e0:e0 + E]; pred = e0 + (model.read(o, one) @ W_.t()).argmax(-1); y = torch.where(act, pred, y)
760
+ model.train(was); return y
761
+
762
+ assert not (args.joint or args.prefix_sup) or arm == "D", "--joint / --prefix_sup are implemented for arm D"
763
+ assert not (args.joint and args.prefix_sup) and not (args.reembed and (arm == "P" or args.joint)) and not (arm == "H" and args.reembed not in (0, 2, 3))
764
+ S = N_ROWS // 2; jtok = torch.zeros((S, 6), dtype=torch.long, device=device)
765
+ jbuf = {k: torch.zeros(S, dtype=torch.long, device=device) for k in ("pos_a", "pos_b", "dep_a", "dep_b", "tgt_a", "tgt_b")}
766
+ pair_rng = named_rng(args.seed, "halt_pair"); ar3 = torch.arange(3, device=device)[None]
767
+ if resume_state is not None and resume_state.get("pair_rng"):
768
+ pair_rng.bit_generator.state = resume_state["pair_rng"]
769
+
770
+ def fill_joint(rows):
771
+ """rows (128, 5) = [t0, t1, t2 | PAD, last, gold] in the order [32 atomic | 32 w2-derived (source-wise consecutive) | 64 d2]."""
772
+ n_w = 2 * N_W2SRC; w = rows[N_AT:N_AT + n_w]; rest = torch.cat([rows[:N_AT], rows[N_AT + n_w:]])
773
+ rest = rest[torch.from_numpy(pair_rng.permutation(len(rest))).to(device)]
774
+ a = torch.cat([w[0::2], rest[0::2]]); b = torch.cat([w[1::2], rest[1::2]]); la = a[:, 3] + 1
775
+ jtok.fill_(PAD); jtok[:, :3] = a[:, :3]; jtok.scatter_(1, la[:, None] + ar3, b[:, :3])
776
+ for k, val in (("pos_a", a[:, 3]), ("pos_b", la + b[:, 3]), ("dep_a", a[:, 3]), ("dep_b", b[:, 3]), ("tgt_a", a[:, 4]), ("tgt_b", b[:, 4])):
777
+ jbuf[k].copy_(val)
778
+
779
+ def fwd_bwd(R):
780
+ with torch.autocast("cuda", dtype=torch.bfloat16, enabled=(args.train_precision == "bf16" and device.type == "cuda")):
781
+ if args.joint:
782
+ loss, st = loss_fn_joint(model, jtok, **jbuf)
783
+ elif FIXED_L and not args.anchor:
784
+ loss, st = loss_fn_untied(model, tok, last, tgt, args.untied_read, FIXED_L, dep_s)
785
+ elif args.prefix_sup:
786
+ loss, st = loss_fn_prefix(model, tok, last, tgt, tgt1, 2 + args.extra_loops)
787
+ elif PACK:
788
+ n1_ = N_AT + 2 * N_W2SRC; ex = args.anchor_extra if args.anchor else 0
789
+ nn_ = pk_n[0]
790
+ cep, anc = loss_packed(model, pk["tok"][:nn_], pk["valid"][:nn_], pk["pv"][:nn_], pk["fire"][:nn_], pk["seg_last"][:nn_], pk["seg_depth"][:nn_], pk["seg_gold"][:nn_], pk["seg_valid"][:nn_], args.anchor, ex)
791
+ sv = pk["seg_valid"][:nn_]; is1 = (pk["seg_depth"][:nn_] == 1) & sv; tot = cep.sum(); st = dict(ce_atomic=(cep * is1).sum() / is1.sum().clamp(min=1), ce_d2=(cep * (sv & ~is1)).sum() / (sv & ~is1).sum().clamp(min=1))
792
+ if args.pack_what == "atomic":
793
+ if args.anchor:
794
+ ce2, anc2 = loss_anchor(model, tok[n1_:], dep_s[n1_:], tgt[n1_:], pv_s[n1_:], ex, off_s[n1_:], 2); anc = (anc + anc2) / 2
795
+ else:
796
+ h_ = model.step(model.step(model.embed(tok[n1_:]))); ce2 = F.cross_entropy(model.logits(model.read(h_, last[n1_:])).float(), tgt[n1_:], reduction="none")
797
+ tot = tot + ce2.sum(); st["ce_d2"] = ce2.mean()
798
+ loss = tot / N_ROWS + args.anchor * anc
799
+ if args.anchor:
800
+ st["anchor"] = anc.detach()
801
+ elif args.anchor:
802
+ n1_ = N_AT + 2 * N_W2SRC
803
+ ce, anc, ast = loss_shallow(model, tok, dep_s, tgt, pv_s, off=off_s, anchor_mode=args.anchor_mode)
804
+ st = dict(ce_atomic=_m(ce[:N_AT]), ce_w2=_m(ce[N_AT:n1_]), ce_d2=_m(ce[n1_:]), anchor=anc.detach(), **ast)
805
+ tot, cnt = ce.sum(), N_ROWS
806
+ if n_deep:
807
+ ced, ancd = loss_anchor(model, tok_d, dep_d, tgt_d, pv_d, args.anchor_extra, off_d, args.deep, (R or 0) if args.delay_max else 0, PAD); tot = tot + deep_w * ced.sum(); cnt = cnt + deep_w * n_deep; anc = (anc + deep_w * ancd) / (1 + deep_w)
808
+ st = dict(st, ce_deep=ced.detach().reshape(len(DEPTHS_X), -1).mean(1), anchor_deep=ancd.detach())
809
+ loss = tot / cnt + args.anchor * anc
810
+ if CST:
811
+ cl, cst_stats = loss_consistency(model, cst_tok, cst_t, cst_ent, mode=args.cs_mode,
812
+ noise=cst_noise, noise_scale=args.cs_noise_scale,
813
+ noise_kind=args.cs_noise_kind, noise_scope=args.cs_noise_scope,
814
+ alpha=cst_alpha, candidate_ids=cst_candidate_ids, relation_ids=cst_relation_ids)
815
+ loss = loss + cst_w * args.consist_w * cl
816
+ st = dict(st, consist=cl.detach(), **cst_stats)
817
+ if FMT:
818
+ fl = loss_format(model, fmt_tok, args.anchor_extra, e0, e0 + E, bool(args.format_bounded)); loss = loss + fmt_w * args.format_w * fl; st = dict(st, format=fl.detach())
819
+ elif DP_AUG and not args.anchor:
820
+ n1_ = N_AT + 2 * N_W2SRC; ce = loss_delay_persist(model, tok, last, tgt, R or 0, args.persist, PAD); st = dict(ce_atomic=_m(ce[:N_AT]), ce_w2=_m(ce[N_AT:n1_]), ce_d2=_m(ce[n1_:]))
821
+ tot, cnt = ce.sum(), N_ROWS
822
+ if n_deep:
823
+ ced = loss_delay_persist(model, tok_d, last_d, tgt_d, R or 0, args.persist, PAD); tot = tot + deep_w * ced.sum(); cnt = cnt + deep_w * n_deep
824
+ st = dict(st, ce_deep=ced.detach().reshape(len(DEPTHS_X), -1).mean(1))
825
+ loss = tot / cnt
826
+ else:
827
+ loss, st = loss_fn(model, arm, tok, last, tgt, R=R, b_train=args.b_train, lam=args.lam)
828
+ if n_deep:
829
+ ced = loss_deep(model, tok_d, last_d, tgt_d, args.deep_rows, DEPTHS_X)
830
+ loss = (loss * N_ROWS + deep_w * ced.sum()) / (N_ROWS + deep_w * n_deep) # still a per-query mean
831
+ st = dict(st, ce_deep=ced.detach().reshape(len(DEPTHS_X), -1).mean(1))
832
+ loss.backward(); gn = torch.nn.utils.clip_grad_norm_(params, args.clip, foreach=True); opt.step()
833
+ return loss, st, gn
834
+
835
+ def set_tf32(on):
836
+ torch.backends.cuda.matmul.allow_tf32 = bool(on); torch.backends.cudnn.allow_tf32 = bool(on)
837
+
838
+ def train_step(R):
839
+ opt.zero_grad(set_to_none=True)
840
+ return fwd_bwd(R)
841
+
842
+ graphs = {}
843
+ R_values = list(range(2, 9)) if arm == "P" else (list(range(0, args.delay_max + 1)) if args.delay_max else [None])
844
+ set_tf32(args.train_precision == "tf32")
845
+ if args.cuda_graph and device.type == "cuda":
846
+ snap = {k: v.clone() for k, v in model.state_dict().items()}
847
+ side = torch.cuda.Stream(); side.wait_stream(torch.cuda.current_stream())
848
+ with torch.cuda.stream(side):
849
+ for R in R_values:
850
+ for _ in range(3):
851
+ train_step(R)
852
+ torch.cuda.current_stream().wait_stream(side)
853
+ for R in R_values:
854
+ g = torch.cuda.CUDAGraph(); opt.zero_grad(set_to_none=True)
855
+ with torch.cuda.graph(g):
856
+ loss, st, gn = fwd_bwd(R)
857
+ graphs[R] = (g, loss, st, gn)
858
+ model.load_state_dict(snap) # in place: undo the warm-up / capture updates exactly
859
+ for p_, s_ in opt.state.items():
860
+ for v_ in s_.values():
861
+ v_.zero_()
862
+ if resume_state is not None: # optimizer state restored IN PLACE (graph-safe); eager path identical
863
+ if not opt.state:
864
+ train_step(R_values[0]); model.load_state_dict(resume_state["model"])
865
+ saved = resume_state["opt"]["state"]
866
+ for i, p_ in enumerate(params):
867
+ if i in saved and p_ in opt.state:
868
+ for k_, v_ in saved[i].items():
869
+ opt.state[p_][k_].copy_(v_)
870
+ if resume_state is not None:
871
+ if "torch_rng" in resume_state: torch.set_rng_state(resume_state["torch_rng"].cpu())
872
+ if "cuda_rng" in resume_state and device.type == "cuda": torch.cuda.set_rng_state_all([z.cpu() for z in resume_state["cuda_rng"]])
873
+ del resume_state
874
+
875
+ t0 = time.time(); log = open(P("train_log.jsonl"), "a")
876
+ for u in range(start, args.updates + 1):
877
+ model.train()
878
+ take = lambda st_, n: torch.from_numpy(st_.take(n) if n else np.zeros(0, dtype=np.int64)).to(device) # a zero share (--mix) leaves that stream untouched
879
+ ia = take(s_at, N_AT); iw = take(s_w2, N_W2SRC); idd = take(s_d2, N_D2)
880
+ rows = torch.cat([T_at[ia], T_w2[iw].reshape(-1, 6), T_d2[idd]]); tgt.copy_(rows[:, 4]); tgt1.copy_(rows[:, 5])
881
+ if OFF:
882
+ ks = torch.from_numpy(pref_sample(off_rng, N_ROWS)).to(device); off_s.copy_(ks)
883
+ if args.offset_fill == "entity":
884
+ tok.copy_(torch.from_numpy(off_rng.integers(e0, e0 + E, tuple(tok.shape))).to(device))
885
+ else:
886
+ tok.fill_(PAD)
887
+ tok.scatter_(1, ks[:, None] + ar_s, rows[:, :3]); last.copy_(rows[:, 3] + ks)
888
+ else:
889
+ tok.copy_(rows[:, :3]); last.copy_(rows[:, 3])
890
+ if args.joint:
891
+ fill_joint(rows)
892
+ if PACK:
893
+ fill_packed(rows)
894
+ if CST:
895
+ cst_w.fill_(1.0 if u >= args.consist_start else 0.0); n_ = args.consist_rows; Dc = args.consist_depth
896
+ if args.cs_mode in SIMPLE_CS_MODES:
897
+ c_tok, c_t, c_ent = sample_simple_cs_tokens(cst_rng, args.cs_mode, n_, Dc, e0, E, r0, meta["num_relations"])
898
+ cst_tok.copy_(torch.from_numpy(c_tok).to(device)); cst_t.copy_(torch.from_numpy(c_t).to(device)); cst_ent.copy_(torch.from_numpy(c_ent).to(device))
899
+ if cst_noise is not None:
900
+ fill_cs_noise_(cst_noise, cst_noise_rng)
901
+ if cst_alpha is not None:
902
+ fill_cs_alpha_(cst_alpha, cst_noise_rng)
903
+ else:
904
+ cst_tok[:, 0] = torch.from_numpy(cst_rng.integers(e0, e0 + E, n_)).to(device); cst_tok[:, 1:] = torch.from_numpy(cst_rng.integers(r0, r0 + meta["num_relations"], (n_, Dc))).to(device)
905
+ cst_t.copy_(torch.from_numpy(cst_rng.integers(0, Dc, n_)).to(device)); cst_ent.copy_(torch.from_numpy(cst_rng.integers(e0, e0 + E, (n_, Dc + 1))).to(device))
906
+ if FMT:
907
+ fmt_w.fill_(1.0 if u >= args.format_start else 0.0)
908
+ fmt_tok[:, 0] = torch.from_numpy(fmt_rng.integers(e0, e0 + E, args.format_rows)).to(device); fmt_tok[:, 1:] = torch.from_numpy(fmt_rng.integers(r0, r0 + meta["num_relations"], (args.format_rows, args.format_depth))).to(device)
909
+ dep_s.copy_(rows[:, 3])
910
+ if args.anchor:
911
+ pv_s.zero_(); pv_s.scatter_(1, off_s[:, None] + ar_s, torch.stack([torch.zeros_like(rows[:, 5]), rows[:, 5], rows[:, 4]], 1))
912
+ if n_deep:
913
+ dr = torch.cat([deep_pool[d][torch.from_numpy(s_deep[d].take(args.deep_rows)).to(device)] for d in DEPTHS_X])
914
+ if OFF:
915
+ kd = torch.from_numpy(pref_sample(off_rng, n_deep)).to(device); off_d.copy_(kd)
916
+ if args.offset_fill == "entity":
917
+ tok_d.copy_(torch.from_numpy(off_rng.integers(e0, e0 + E, tuple(tok_d.shape))).to(device))
918
+ else:
919
+ tok_d.fill_(PAD)
920
+ tok_d.scatter_(1, kd[:, None] + ar_d[:, :Ld], dr[:, :Ld]); last_d.copy_(dr[:, Ld] + kd)
921
+ else:
922
+ tok_d.copy_(dr[:, :Ld]); last_d.copy_(dr[:, Ld])
923
+ gold_d.copy_(dr[:, Ld + 1]); deep_w.fill_(1.0 if u >= args.deep_start else 0.0)
924
+ tgt_d.copy_(self_labels() if (args.deep_labels == "self" and u >= args.deep_start) else gold_d)
925
+ if args.anchor:
926
+ dep_d.copy_(dr[:, Ld]); y_ = dr[:, 0] - e0; pvu = torch.zeros((n_deep, Ld), dtype=torch.long, device=device)
927
+ for j_ in range(1, Ld):
928
+ r_ = (dr[:, j_] - r0).clamp(0, REL_T.shape[0] - 1); y_ = torch.where(dep_d >= j_, REL_T[r_, y_], y_); pvu[:, j_] = e0 + y_
929
+ pv_d.zero_(); pv_d.scatter_(1, off_d[:, None] + ar_d[:, :Ld], pvu)
930
+ if args.deep_labels == "self" and u % 100 == 0:
931
+ ctr["self_label_acc"] = float((tgt_d == gold_d).float().mean())
932
+ R = int(np.clip(r_rng.poisson(4), 2, 8)) # the stream advances in every arm; only P uses it
933
+ Ru = R if arm == "P" else ((int(delay_rng.integers(0, args.delay_max + 1)) if u >= args.delay_start else 0) if args.delay_max else None)
934
+ lr_now = args.lr * min(1.0, u / max(1, args.warmup))
935
+ if args.lr_decay == "linear" and u > args.warmup:
936
+ lr_now = args.lr * max(0.0, (args.updates - u) / max(1, args.updates - args.warmup))
937
+ elif args.lr_decay == "cosine" and u > args.warmup: # linear warm-up over --warmup updates, then half a cosine down to 0 at the last update
938
+ lr_now = args.lr * 0.5 * (1.0 + math.cos(math.pi * (u - args.warmup) / max(1, args.updates - args.warmup)))
939
+ if capturable:
940
+ opt.param_groups[0]["lr"].fill_(lr_now)
941
+ else:
942
+ opt.param_groups[0]["lr"] = lr_now
943
+ set_tf32(args.train_precision == "tf32")
944
+ if graphs:
945
+ g, loss, st, gn = graphs[Ru]; g.replay()
946
+ else:
947
+ loss, st, gn = train_step(Ru)
948
+ set_tf32(False) # everything outside the training step (all evaluations) is strict fp32
949
+ ctr["atomic"] += N_AT; ctr["w2_sources"] += N_W2SRC; ctr["w2_queries"] += 2 * N_W2SRC; ctr["d2"] += N_D2
950
+ ctr["row_loops"] += {"D": N_ROWS * 4, "H": N_ROWS * args.b_train, "P": N_ROWS * R}[arm]
951
+ if arm == "P":
952
+ ctr["R_hist"][str(R)] = ctr["R_hist"].get(str(R), 0) + 1
953
+ if u % 100 == 0 or u == 1:
954
+ rec = dict(update=u, lr=lr_now, elapsed=time.time() - t0, loss=float(loss), grad_norm=float(gn), **{k: (v.tolist() if v.dim() else float(v)) for k, v in st.items()})
955
+ log.write(json.dumps(rec) + "\n"); log.flush()
956
+ if not math.isfinite(rec["loss"]) or not math.isfinite(rec["grad_norm"]):
957
+ raise RuntimeError("nonfinite training loss or gradient; preserve previous checkpoint")
958
+ ev = not stop_requested[0] and ((u % args.eval_every == 0) or u == args.updates or (u in (1000, 2000, 5000)) or (args.smoke and u % 20 == 0))
959
+ if ev:
960
+ checkpoint_record = save_eval_checkpoint(model, u, args.save_dir, args.rope_base, args.max_pos)
961
+ model.eval(); te = time.time(); rec = dict(update=u, elapsed=time.time() - t0, shallow=shallow(), checkpoint=checkpoint_record)
962
+ r = main_protocol(MON); rec["monitor"] = by_cell(mon_items, **r)
963
+ rb = eval_fixed(model, MON, ["d", "d+1", "d+2", "2d"]) # diagnostics on the same weights (never the main metric of H / P)
964
+ rec["monitor_budgets"] = {b: {k: v["correct"] for k, v in by_cell(mon_items, correct=rb[b]).items()} for b in rb}
965
+ if arm == "H":
966
+ rec["monitor_forced_R_eq_d"] = rec["monitor_budgets"]["d"]
967
+ rec["counters"] = dict(ctr); rec["eval_seconds"] = time.time() - te; rec["peak_mem_gb"] = torch.cuda.max_memory_allocated() / 2 ** 30 if device.type == "cuda" else None
968
+ with open(P("metrics.jsonl"), "a") as f:
969
+ f.write(json.dumps(rec) + "\n")
970
+ sh = rec["shallow"]; m = rec["monitor"]
971
+ print(f"u {u} | atomic {sh['atomic']['correct']:.3f} val_d2 {sh['val_d2']['correct']:.3f} fit {sh['train_d2_fit']['correct']:.3f} | " + " ".join(f"{k.split('|')[1]}:{m[k]['correct']:.2f}" for k in m if k.startswith("one_new"))
972
+ + (f" | T(d2/d8/d128) {m['one_new|d2']['T']:.1f}/{m['one_new|d8']['T']:.1f}/{m['one_new|d128']['T']:.1f} timeout(d8) {m['one_new|d8']['timeout']:.2f}" if arm == "H" and "one_new|d128" in m else "")
973
+ + f" | eval {rec['eval_seconds']:.0f}s | {time.time() - t0:.0f}s", flush=True)
974
+ score = sh["val_d2"]["correct"]
975
+ if best is None or score > best[0]:
976
+ best = [score, u]; torch.save(dict(model=model.state_dict(), update=u), P("best.pt"))
977
+ if stop_requested[0] or u % args.eval_every == 0 or u == args.updates or u == args.stop_after or (args.ckpt_every and u % args.ckpt_every == 0):
978
+ if (u % 100000 == 0 or u == args.updates) and not ev:
979
+ torch.save(dict(model=model.state_dict(), update=u), P(f"ckpt_{u:07d}.pt"))
980
+ osd = opt.state_dict(); osd = dict(state={i: {k: v.clone() for k, v in s.items()} for i, s in osd["state"].items()})
981
+ torch.save(dict(model=model.state_dict(), opt=osd, update=u, s_at=s_at.state(), s_w2=s_w2.state(), s_d2=s_d2.state(), r_rng=r_rng.bit_generator.state, pair_rng=pair_rng.bit_generator.state, ctr=ctr, best=best,
982
+ s_deep={str(d): s_deep[d].state() for d in DEPTHS_X}, off_rng=off_rng.bit_generator.state,
983
+ delay_rng=delay_rng.bit_generator.state, pack_rng=pack_rng.bit_generator.state, fmt_rng=fmt_rng.bit_generator.state, cst_rng=cst_rng.bit_generator.state,
984
+ cst_noise_rng=cst_noise_rng.get_state() if cst_noise_rng is not None else None,
985
+ torch_rng=torch.get_rng_state(), cuda_rng=torch.cuda.get_rng_state_all() if device.type == "cuda" else []), last_path + ".tmp")
986
+ os.replace(last_path + ".tmp", last_path)
987
+ if stop_requested[0]:
988
+ print(f"checkpointed on signal at update {u}", flush=True)
989
+ sys.exit(75)
990
+ if u == args.stop_after:
991
+ print(f"stopped at update {u} ({(time.time() - t0) / max(1, u - start + 1) * 1000:.1f} ms / update incl. evals)", flush=True); return
992
+ final_eval(ckpt_update=args.updates); print("done", flush=True)
993
+
994
+
995
+ if __name__ == "__main__":
996
+ main()
for Minegishi/wait_denoising_fixed_t0/README.md ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ # Wait isotropic denoising / fixed t=0
2
+
3
+ Original run `R10k_T0iso_s7`. RoPE base 10000, seed 7, full shallow anchor and random entity prefix 0–128; fixed 200k updates.
4
+
5
+ All 23 evaluation checkpoints are provided; the primary final checkpoint is `ckpt_0200000.pt`. Formal final errors: 3/120000.
6
+
7
+ See [the shared README](../README.md) for the exact training objective, model, data, loading example and line plots. `manifest.json` contains the training settings; `train_command.json` contains the command relative to the package root. Named checkpoints contain weights and positional metadata, not optimizer/RNG states.
for Minegishi/wait_denoising_fixed_t0/checkpoints.json ADDED
@@ -0,0 +1,237 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "evaluated_updates": [
3
+ 1000,
4
+ 2000,
5
+ 5000,
6
+ 10000,
7
+ 20000,
8
+ 30000,
9
+ 40000,
10
+ 50000,
11
+ 60000,
12
+ 70000,
13
+ 80000,
14
+ 90000,
15
+ 100000,
16
+ 110000,
17
+ 120000,
18
+ 130000,
19
+ 140000,
20
+ 150000,
21
+ 160000,
22
+ 170000,
23
+ 180000,
24
+ 190000,
25
+ 200000
26
+ ],
27
+ "all_saved": true,
28
+ "checkpoints": [
29
+ {
30
+ "file": "ckpt_0001000.pt",
31
+ "update": 1000,
32
+ "sha256": "b04cbfdb2a1e7a54e1c3857ef6d70d87542c02cd8f03baf4e874d8a998915e97",
33
+ "bytes": 115042564,
34
+ "rope_base": 10000.0,
35
+ "max_pos": 4097,
36
+ "kind": "model_weights_for_exact_evaluation"
37
+ },
38
+ {
39
+ "file": "ckpt_0002000.pt",
40
+ "update": 2000,
41
+ "sha256": "09775249c367ccf150c3fd773f92b295743a30dc7db1bc910b9b1f29ba125f68",
42
+ "bytes": 115042564,
43
+ "rope_base": 10000.0,
44
+ "max_pos": 4097,
45
+ "kind": "model_weights_for_exact_evaluation"
46
+ },
47
+ {
48
+ "file": "ckpt_0005000.pt",
49
+ "update": 5000,
50
+ "sha256": "cbb6360b83c05c86ba437a856901801e36acf74776beefffe41db5a319305cd6",
51
+ "bytes": 115042564,
52
+ "rope_base": 10000.0,
53
+ "max_pos": 4097,
54
+ "kind": "model_weights_for_exact_evaluation"
55
+ },
56
+ {
57
+ "file": "ckpt_0010000.pt",
58
+ "update": 10000,
59
+ "sha256": "aefc1a2e8d13155cf68e5cdbe2b56648b666bede651114928563fe53aee388d8",
60
+ "bytes": 115042564,
61
+ "rope_base": 10000.0,
62
+ "max_pos": 4097,
63
+ "kind": "model_weights_for_exact_evaluation"
64
+ },
65
+ {
66
+ "file": "ckpt_0020000.pt",
67
+ "update": 20000,
68
+ "sha256": "107b8652838126379d513366c4e684c2de800ed5582deae281f48b691966e505",
69
+ "bytes": 115042564,
70
+ "rope_base": 10000.0,
71
+ "max_pos": 4097,
72
+ "kind": "model_weights_for_exact_evaluation"
73
+ },
74
+ {
75
+ "file": "ckpt_0030000.pt",
76
+ "update": 30000,
77
+ "sha256": "584e9f1b1fb3f1ce52fa7c95b5a8dc85c8696f4e9aa829c583dc415bd1144f79",
78
+ "bytes": 115042564,
79
+ "rope_base": 10000.0,
80
+ "max_pos": 4097,
81
+ "kind": "model_weights_for_exact_evaluation"
82
+ },
83
+ {
84
+ "file": "ckpt_0040000.pt",
85
+ "update": 40000,
86
+ "sha256": "259bdfe027e50d18508c6e332add29ebca312fce38c21f6fce9bba2be6b86ce5",
87
+ "bytes": 115042564,
88
+ "rope_base": 10000.0,
89
+ "max_pos": 4097,
90
+ "kind": "model_weights_for_exact_evaluation"
91
+ },
92
+ {
93
+ "file": "ckpt_0050000.pt",
94
+ "update": 50000,
95
+ "sha256": "f8e4cd617629f536541f238f8e7df5e3b89f15db47c4c3025b5b7757a58a51fa",
96
+ "bytes": 115042564,
97
+ "rope_base": 10000.0,
98
+ "max_pos": 4097,
99
+ "kind": "model_weights_for_exact_evaluation"
100
+ },
101
+ {
102
+ "file": "ckpt_0060000.pt",
103
+ "update": 60000,
104
+ "sha256": "f6dd0d27bcb48a93839be15dbccbf37b20db15b5bb5144ba8ef5eda581a0176b",
105
+ "bytes": 115042564,
106
+ "rope_base": 10000.0,
107
+ "max_pos": 4097,
108
+ "kind": "model_weights_for_exact_evaluation"
109
+ },
110
+ {
111
+ "file": "ckpt_0070000.pt",
112
+ "update": 70000,
113
+ "sha256": "afe2b3f44168be01e54532886d2d1439bcb167f0f437c826c45e14668ce1caa0",
114
+ "bytes": 115042564,
115
+ "rope_base": 10000.0,
116
+ "max_pos": 4097,
117
+ "kind": "model_weights_for_exact_evaluation"
118
+ },
119
+ {
120
+ "file": "ckpt_0080000.pt",
121
+ "update": 80000,
122
+ "sha256": "2800b11b460d0e5de02d7009c25a0e4d93b555c7c63f10a59e89739383043118",
123
+ "bytes": 115042564,
124
+ "rope_base": 10000.0,
125
+ "max_pos": 4097,
126
+ "kind": "model_weights_for_exact_evaluation"
127
+ },
128
+ {
129
+ "file": "ckpt_0090000.pt",
130
+ "update": 90000,
131
+ "sha256": "61ba4ff1b03eed8710610004e9b1e68db3886813e2fbd68c83d58da444a9dafe",
132
+ "bytes": 115042564,
133
+ "rope_base": 10000.0,
134
+ "max_pos": 4097,
135
+ "kind": "model_weights_for_exact_evaluation"
136
+ },
137
+ {
138
+ "file": "ckpt_0100000.pt",
139
+ "update": 100000,
140
+ "sha256": "e600093705a1477f6a092135438a4346dd1a4bc5c97430f2d179fdbb7b73a8d7",
141
+ "bytes": 115042564,
142
+ "rope_base": 10000.0,
143
+ "max_pos": 4097,
144
+ "kind": "model_weights_for_exact_evaluation"
145
+ },
146
+ {
147
+ "file": "ckpt_0110000.pt",
148
+ "update": 110000,
149
+ "sha256": "1791c9e8a4be8f84d2fdd713b598439dd7ca3332ae0b9efdc20c53dbd86c07e6",
150
+ "bytes": 115042564,
151
+ "rope_base": 10000.0,
152
+ "max_pos": 4097,
153
+ "kind": "model_weights_for_exact_evaluation"
154
+ },
155
+ {
156
+ "file": "ckpt_0120000.pt",
157
+ "update": 120000,
158
+ "sha256": "c6767cfb80c10f38ca4bac059d5240211929dbe5a5cadb4b48ebb07a196566bf",
159
+ "bytes": 115042564,
160
+ "rope_base": 10000.0,
161
+ "max_pos": 4097,
162
+ "kind": "model_weights_for_exact_evaluation"
163
+ },
164
+ {
165
+ "file": "ckpt_0130000.pt",
166
+ "update": 130000,
167
+ "sha256": "1bf915b7a6cb858a51b534b9afd289421e0fdbb3c6f534831f0c4254143fb890",
168
+ "bytes": 115042564,
169
+ "rope_base": 10000.0,
170
+ "max_pos": 4097,
171
+ "kind": "model_weights_for_exact_evaluation"
172
+ },
173
+ {
174
+ "file": "ckpt_0140000.pt",
175
+ "update": 140000,
176
+ "sha256": "a00e672829198b3cba32d4241aa914b067cbd82f4f46f1a1219075f23ce22f8a",
177
+ "bytes": 115042564,
178
+ "rope_base": 10000.0,
179
+ "max_pos": 4097,
180
+ "kind": "model_weights_for_exact_evaluation"
181
+ },
182
+ {
183
+ "file": "ckpt_0150000.pt",
184
+ "update": 150000,
185
+ "sha256": "c7b347a167d12485afddc0943dc92ce8326e3eb5a099ce5e4aafb4c2c0b81d8c",
186
+ "bytes": 115042564,
187
+ "rope_base": 10000.0,
188
+ "max_pos": 4097,
189
+ "kind": "model_weights_for_exact_evaluation"
190
+ },
191
+ {
192
+ "file": "ckpt_0160000.pt",
193
+ "update": 160000,
194
+ "sha256": "bbd8ee69c630dd68e2ae25cb2eacb4bd1d3ffcadadd37a7bbc9938cc791ecb3a",
195
+ "bytes": 115042564,
196
+ "rope_base": 10000.0,
197
+ "max_pos": 4097,
198
+ "kind": "model_weights_for_exact_evaluation"
199
+ },
200
+ {
201
+ "file": "ckpt_0170000.pt",
202
+ "update": 170000,
203
+ "sha256": "76a4c1cd41d687ef06a1ccf6b67b2c56399c1ea2bf1103c51c144b3eddb6e20f",
204
+ "bytes": 115042564,
205
+ "rope_base": 10000.0,
206
+ "max_pos": 4097,
207
+ "kind": "model_weights_for_exact_evaluation"
208
+ },
209
+ {
210
+ "file": "ckpt_0180000.pt",
211
+ "update": 180000,
212
+ "sha256": "d767ae2ced030111845d93e7f35587d7ecfc7d8ce43622d47df0d9c93638b191",
213
+ "bytes": 115042564,
214
+ "rope_base": 10000.0,
215
+ "max_pos": 4097,
216
+ "kind": "model_weights_for_exact_evaluation"
217
+ },
218
+ {
219
+ "file": "ckpt_0190000.pt",
220
+ "update": 190000,
221
+ "sha256": "a44ac51a1decb6919d48d1d02fcef7f4e525759760b913fffa12e064754d2c29",
222
+ "bytes": 115042564,
223
+ "rope_base": 10000.0,
224
+ "max_pos": 4097,
225
+ "kind": "model_weights_for_exact_evaluation"
226
+ },
227
+ {
228
+ "file": "ckpt_0200000.pt",
229
+ "update": 200000,
230
+ "sha256": "6be182b9f912f741dfb0ba491879812e0813cc8adb264c76fa0cf5a9bbc64e22",
231
+ "bytes": 115042564,
232
+ "rope_base": 10000.0,
233
+ "max_pos": 4097,
234
+ "kind": "model_weights_for_exact_evaluation"
235
+ }
236
+ ]
237
+ }
for Minegishi/wait_denoising_fixed_t0/config.json ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "id": "R10k_T0iso_s7",
3
+ "lane": 1,
4
+ "seed": 7,
5
+ "cs_mode": "t0_wait_denoise",
6
+ "cs_weight": 1,
7
+ "cs_noise_scale": 0.2,
8
+ "cs_noise_kind": "isotropic",
9
+ "cs_noise_scope": "wait",
10
+ "anchor_mode": "full",
11
+ "cs_depth": 32,
12
+ "cs_rows": 32,
13
+ "cs_start": 20000,
14
+ "updates": 200000,
15
+ "rope_base": 10000,
16
+ "max_pos": 4097,
17
+ "eval_every": 10000,
18
+ "checkpoint_at_every_evaluation": true
19
+ }
for Minegishi/wait_denoising_fixed_t0/final_eval.json ADDED
@@ -0,0 +1,513 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "arm": "D",
3
+ "seed": 7,
4
+ "lam": 0.0,
5
+ "b_eval": 256,
6
+ "stop_thr": 0.5,
7
+ "tag": "",
8
+ "checkpoint_update": 200000,
9
+ "main": {
10
+ "all_seen|d128": {
11
+ "n": 5000,
12
+ "correct": 0.9998
13
+ },
14
+ "all_seen|d16": {
15
+ "n": 5000,
16
+ "correct": 1.0
17
+ },
18
+ "all_seen|d2": {
19
+ "n": 5000,
20
+ "correct": 1.0
21
+ },
22
+ "all_seen|d3": {
23
+ "n": 5000,
24
+ "correct": 1.0
25
+ },
26
+ "all_seen|d32": {
27
+ "n": 5000,
28
+ "correct": 1.0
29
+ },
30
+ "all_seen|d4": {
31
+ "n": 5000,
32
+ "correct": 1.0
33
+ },
34
+ "all_seen|d64": {
35
+ "n": 5000,
36
+ "correct": 1.0
37
+ },
38
+ "all_seen|d8": {
39
+ "n": 5000,
40
+ "correct": 1.0
41
+ },
42
+ "one_new|d128": {
43
+ "n": 5000,
44
+ "correct": 1.0
45
+ },
46
+ "one_new|d16": {
47
+ "n": 5000,
48
+ "correct": 0.9998
49
+ },
50
+ "one_new|d2": {
51
+ "n": 5000,
52
+ "correct": 1.0
53
+ },
54
+ "one_new|d3": {
55
+ "n": 5000,
56
+ "correct": 1.0
57
+ },
58
+ "one_new|d32": {
59
+ "n": 5000,
60
+ "correct": 1.0
61
+ },
62
+ "one_new|d4": {
63
+ "n": 5000,
64
+ "correct": 1.0
65
+ },
66
+ "one_new|d64": {
67
+ "n": 5000,
68
+ "correct": 1.0
69
+ },
70
+ "one_new|d8": {
71
+ "n": 5000,
72
+ "correct": 1.0
73
+ },
74
+ "random|d128": {
75
+ "n": 5000,
76
+ "correct": 1.0
77
+ },
78
+ "random|d16": {
79
+ "n": 5000,
80
+ "correct": 1.0
81
+ },
82
+ "random|d2": {
83
+ "n": 5000,
84
+ "correct": 1.0
85
+ },
86
+ "random|d3": {
87
+ "n": 5000,
88
+ "correct": 1.0
89
+ },
90
+ "random|d32": {
91
+ "n": 5000,
92
+ "correct": 1.0
93
+ },
94
+ "random|d4": {
95
+ "n": 5000,
96
+ "correct": 1.0
97
+ },
98
+ "random|d64": {
99
+ "n": 5000,
100
+ "correct": 1.0
101
+ },
102
+ "random|d8": {
103
+ "n": 5000,
104
+ "correct": 0.9998
105
+ }
106
+ },
107
+ "shallow": {
108
+ "atomic": {
109
+ "correct": 1.0
110
+ },
111
+ "val_d2": {
112
+ "correct": 1.0
113
+ },
114
+ "train_d2_fit": {
115
+ "correct": 1.0
116
+ }
117
+ },
118
+ "forced_R_eq_d": {
119
+ "all_seen|d128": {
120
+ "n": 5000,
121
+ "correct": 0.9998
122
+ },
123
+ "all_seen|d16": {
124
+ "n": 5000,
125
+ "correct": 1.0
126
+ },
127
+ "all_seen|d2": {
128
+ "n": 5000,
129
+ "correct": 1.0
130
+ },
131
+ "all_seen|d3": {
132
+ "n": 5000,
133
+ "correct": 1.0
134
+ },
135
+ "all_seen|d32": {
136
+ "n": 5000,
137
+ "correct": 1.0
138
+ },
139
+ "all_seen|d4": {
140
+ "n": 5000,
141
+ "correct": 1.0
142
+ },
143
+ "all_seen|d64": {
144
+ "n": 5000,
145
+ "correct": 1.0
146
+ },
147
+ "all_seen|d8": {
148
+ "n": 5000,
149
+ "correct": 1.0
150
+ },
151
+ "one_new|d128": {
152
+ "n": 5000,
153
+ "correct": 1.0
154
+ },
155
+ "one_new|d16": {
156
+ "n": 5000,
157
+ "correct": 0.9998
158
+ },
159
+ "one_new|d2": {
160
+ "n": 5000,
161
+ "correct": 1.0
162
+ },
163
+ "one_new|d3": {
164
+ "n": 5000,
165
+ "correct": 1.0
166
+ },
167
+ "one_new|d32": {
168
+ "n": 5000,
169
+ "correct": 1.0
170
+ },
171
+ "one_new|d4": {
172
+ "n": 5000,
173
+ "correct": 1.0
174
+ },
175
+ "one_new|d64": {
176
+ "n": 5000,
177
+ "correct": 1.0
178
+ },
179
+ "one_new|d8": {
180
+ "n": 5000,
181
+ "correct": 1.0
182
+ },
183
+ "random|d128": {
184
+ "n": 5000,
185
+ "correct": 1.0
186
+ },
187
+ "random|d16": {
188
+ "n": 5000,
189
+ "correct": 1.0
190
+ },
191
+ "random|d2": {
192
+ "n": 5000,
193
+ "correct": 1.0
194
+ },
195
+ "random|d3": {
196
+ "n": 5000,
197
+ "correct": 1.0
198
+ },
199
+ "random|d32": {
200
+ "n": 5000,
201
+ "correct": 1.0
202
+ },
203
+ "random|d4": {
204
+ "n": 5000,
205
+ "correct": 1.0
206
+ },
207
+ "random|d64": {
208
+ "n": 5000,
209
+ "correct": 1.0
210
+ },
211
+ "random|d8": {
212
+ "n": 5000,
213
+ "correct": 0.9998
214
+ }
215
+ },
216
+ "matched_budget_diagnostics": {
217
+ "d": {
218
+ "all_seen|d128": {
219
+ "n": 100,
220
+ "correct": 1.0
221
+ },
222
+ "all_seen|d16": {
223
+ "n": 100,
224
+ "correct": 1.0
225
+ },
226
+ "all_seen|d2": {
227
+ "n": 100,
228
+ "correct": 1.0
229
+ },
230
+ "all_seen|d3": {
231
+ "n": 100,
232
+ "correct": 1.0
233
+ },
234
+ "all_seen|d32": {
235
+ "n": 100,
236
+ "correct": 1.0
237
+ },
238
+ "all_seen|d4": {
239
+ "n": 100,
240
+ "correct": 1.0
241
+ },
242
+ "all_seen|d64": {
243
+ "n": 100,
244
+ "correct": 1.0
245
+ },
246
+ "all_seen|d8": {
247
+ "n": 100,
248
+ "correct": 1.0
249
+ },
250
+ "one_new|d128": {
251
+ "n": 100,
252
+ "correct": 1.0
253
+ },
254
+ "one_new|d16": {
255
+ "n": 100,
256
+ "correct": 1.0
257
+ },
258
+ "one_new|d2": {
259
+ "n": 100,
260
+ "correct": 1.0
261
+ },
262
+ "one_new|d3": {
263
+ "n": 100,
264
+ "correct": 1.0
265
+ },
266
+ "one_new|d32": {
267
+ "n": 100,
268
+ "correct": 1.0
269
+ },
270
+ "one_new|d4": {
271
+ "n": 100,
272
+ "correct": 1.0
273
+ },
274
+ "one_new|d64": {
275
+ "n": 100,
276
+ "correct": 1.0
277
+ },
278
+ "one_new|d8": {
279
+ "n": 100,
280
+ "correct": 1.0
281
+ },
282
+ "random|d128": {
283
+ "n": 100,
284
+ "correct": 1.0
285
+ },
286
+ "random|d16": {
287
+ "n": 100,
288
+ "correct": 1.0
289
+ },
290
+ "random|d2": {
291
+ "n": 100,
292
+ "correct": 1.0
293
+ },
294
+ "random|d3": {
295
+ "n": 100,
296
+ "correct": 1.0
297
+ },
298
+ "random|d32": {
299
+ "n": 100,
300
+ "correct": 1.0
301
+ },
302
+ "random|d4": {
303
+ "n": 100,
304
+ "correct": 1.0
305
+ },
306
+ "random|d64": {
307
+ "n": 100,
308
+ "correct": 1.0
309
+ },
310
+ "random|d8": {
311
+ "n": 100,
312
+ "correct": 1.0
313
+ }
314
+ },
315
+ "d+2": {
316
+ "all_seen|d128": {
317
+ "n": 100,
318
+ "correct": 1.0
319
+ },
320
+ "all_seen|d16": {
321
+ "n": 100,
322
+ "correct": 1.0
323
+ },
324
+ "all_seen|d2": {
325
+ "n": 100,
326
+ "correct": 1.0
327
+ },
328
+ "all_seen|d3": {
329
+ "n": 100,
330
+ "correct": 1.0
331
+ },
332
+ "all_seen|d32": {
333
+ "n": 100,
334
+ "correct": 1.0
335
+ },
336
+ "all_seen|d4": {
337
+ "n": 100,
338
+ "correct": 1.0
339
+ },
340
+ "all_seen|d64": {
341
+ "n": 100,
342
+ "correct": 1.0
343
+ },
344
+ "all_seen|d8": {
345
+ "n": 100,
346
+ "correct": 1.0
347
+ },
348
+ "one_new|d128": {
349
+ "n": 100,
350
+ "correct": 1.0
351
+ },
352
+ "one_new|d16": {
353
+ "n": 100,
354
+ "correct": 1.0
355
+ },
356
+ "one_new|d2": {
357
+ "n": 100,
358
+ "correct": 1.0
359
+ },
360
+ "one_new|d3": {
361
+ "n": 100,
362
+ "correct": 1.0
363
+ },
364
+ "one_new|d32": {
365
+ "n": 100,
366
+ "correct": 1.0
367
+ },
368
+ "one_new|d4": {
369
+ "n": 100,
370
+ "correct": 1.0
371
+ },
372
+ "one_new|d64": {
373
+ "n": 100,
374
+ "correct": 1.0
375
+ },
376
+ "one_new|d8": {
377
+ "n": 100,
378
+ "correct": 1.0
379
+ },
380
+ "random|d128": {
381
+ "n": 100,
382
+ "correct": 1.0
383
+ },
384
+ "random|d16": {
385
+ "n": 100,
386
+ "correct": 1.0
387
+ },
388
+ "random|d2": {
389
+ "n": 100,
390
+ "correct": 1.0
391
+ },
392
+ "random|d3": {
393
+ "n": 100,
394
+ "correct": 1.0
395
+ },
396
+ "random|d32": {
397
+ "n": 100,
398
+ "correct": 1.0
399
+ },
400
+ "random|d4": {
401
+ "n": 100,
402
+ "correct": 1.0
403
+ },
404
+ "random|d64": {
405
+ "n": 100,
406
+ "correct": 1.0
407
+ },
408
+ "random|d8": {
409
+ "n": 100,
410
+ "correct": 1.0
411
+ }
412
+ },
413
+ "2d": {
414
+ "all_seen|d128": {
415
+ "n": 100,
416
+ "correct": 1.0
417
+ },
418
+ "all_seen|d16": {
419
+ "n": 100,
420
+ "correct": 1.0
421
+ },
422
+ "all_seen|d2": {
423
+ "n": 100,
424
+ "correct": 1.0
425
+ },
426
+ "all_seen|d3": {
427
+ "n": 100,
428
+ "correct": 1.0
429
+ },
430
+ "all_seen|d32": {
431
+ "n": 100,
432
+ "correct": 1.0
433
+ },
434
+ "all_seen|d4": {
435
+ "n": 100,
436
+ "correct": 1.0
437
+ },
438
+ "all_seen|d64": {
439
+ "n": 100,
440
+ "correct": 1.0
441
+ },
442
+ "all_seen|d8": {
443
+ "n": 100,
444
+ "correct": 1.0
445
+ },
446
+ "one_new|d128": {
447
+ "n": 100,
448
+ "correct": 1.0
449
+ },
450
+ "one_new|d16": {
451
+ "n": 100,
452
+ "correct": 1.0
453
+ },
454
+ "one_new|d2": {
455
+ "n": 100,
456
+ "correct": 1.0
457
+ },
458
+ "one_new|d3": {
459
+ "n": 100,
460
+ "correct": 1.0
461
+ },
462
+ "one_new|d32": {
463
+ "n": 100,
464
+ "correct": 1.0
465
+ },
466
+ "one_new|d4": {
467
+ "n": 100,
468
+ "correct": 1.0
469
+ },
470
+ "one_new|d64": {
471
+ "n": 100,
472
+ "correct": 1.0
473
+ },
474
+ "one_new|d8": {
475
+ "n": 100,
476
+ "correct": 1.0
477
+ },
478
+ "random|d128": {
479
+ "n": 100,
480
+ "correct": 1.0
481
+ },
482
+ "random|d16": {
483
+ "n": 100,
484
+ "correct": 1.0
485
+ },
486
+ "random|d2": {
487
+ "n": 100,
488
+ "correct": 1.0
489
+ },
490
+ "random|d3": {
491
+ "n": 100,
492
+ "correct": 1.0
493
+ },
494
+ "random|d32": {
495
+ "n": 100,
496
+ "correct": 1.0
497
+ },
498
+ "random|d4": {
499
+ "n": 100,
500
+ "correct": 1.0
501
+ },
502
+ "random|d64": {
503
+ "n": 100,
504
+ "correct": 1.0
505
+ },
506
+ "random|d8": {
507
+ "n": 100,
508
+ "correct": 1.0
509
+ }
510
+ }
511
+ },
512
+ "eval_seconds": 497.4614384174347
513
+ }
for Minegishi/wait_denoising_fixed_t0/manifest.json ADDED
@@ -0,0 +1,188 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "theory_args": {
3
+ "data_dir": "reproducibility/data/chain_loop_k8",
4
+ "atomic": "reproducibility/data/atomic_joint_2026-09-19/train_atomic.json",
5
+ "save_dir": "wait_denoising_fixed_t0",
6
+ "arm": "D",
7
+ "seed": 7,
8
+ "d_model": 768,
9
+ "n_head": 12,
10
+ "n_layer": 4,
11
+ "reembed": 0,
12
+ "prefix_sup": 0,
13
+ "extra_loops": 2,
14
+ "joint": 0,
15
+ "attn_window": 0,
16
+ "pos": "rope",
17
+ "lr": 0.0001,
18
+ "weight_decay": 0.01,
19
+ "warmup": 2000,
20
+ "clip": 1.0,
21
+ "updates": 200000,
22
+ "lr_decay": "linear",
23
+ "adam_eps": 1e-08,
24
+ "mix": "",
25
+ "deep": 0,
26
+ "anchor": 1.0,
27
+ "consist_w": 1.0,
28
+ "consist_depth": 32,
29
+ "consist_rows": 32,
30
+ "consist_start": 20000,
31
+ "format_w": 0.0,
32
+ "format_start": 0,
33
+ "format_bounded": 1,
34
+ "format_depth": 32,
35
+ "format_rows": 16,
36
+ "anchor_extra": 2,
37
+ "delay_max": 0,
38
+ "delay_start": 0,
39
+ "persist": 0,
40
+ "pack": 0,
41
+ "pack_prefix_max": 0,
42
+ "offset_dist": "uniform",
43
+ "pack_what": "all",
44
+ "offset_fill": "entity",
45
+ "offset_max": 128,
46
+ "deep_list": "",
47
+ "deep_n": 10000,
48
+ "deep_rows": 16,
49
+ "deep_labels": "truth",
50
+ "deep_start": 0,
51
+ "deep_seed": 20260921,
52
+ "nr_init": 0,
53
+ "ready_state": 0,
54
+ "untied": 0,
55
+ "unroll": 0,
56
+ "untied_group": 1,
57
+ "untied_read": "final",
58
+ "eval_ckpt": "",
59
+ "eval_tag": "",
60
+ "b_train": 8,
61
+ "b_eval": 256,
62
+ "lam": 0.0,
63
+ "stop_thr": 0.5,
64
+ "stop_bias": 0.0,
65
+ "eval_every": 10000,
66
+ "monitor_per_cell": 100,
67
+ "cuda_graph": 1,
68
+ "resume": false,
69
+ "eval_only": false,
70
+ "train_precision": "tf32",
71
+ "ckpt_every": 5000,
72
+ "stop_after": 0,
73
+ "smoke": 0,
74
+ "cs_mode": "t0_wait_denoise",
75
+ "cs_noise_scale": 0.2,
76
+ "cs_noise_kind": "isotropic",
77
+ "cs_noise_scope": "wait",
78
+ "anchor_mode": "full",
79
+ "final_n": 5000,
80
+ "diagnostic_n": 100,
81
+ "rope_base": 10000.0,
82
+ "max_pos": 4097
83
+ },
84
+ "arm": "D",
85
+ "seed": 7,
86
+ "n_params": 28756225,
87
+ "stop_head_params": 769,
88
+ "backbone_init_hash": "7fa2889005e8978a",
89
+ "model": {
90
+ "d_model": 768,
91
+ "n_head": 12,
92
+ "n_layer": 4,
93
+ "position": "rope",
94
+ "rope_base": 10000.0,
95
+ "max_pos": 4097,
96
+ "attn_window": 0,
97
+ "reembed": 0,
98
+ "nr_init": false,
99
+ "ready_state": false,
100
+ "untied": null,
101
+ "tied_unroll": null,
102
+ "tied_lm_head": true,
103
+ "residual_out_proj_init": 0.0,
104
+ "dropout": 0.0,
105
+ "ln_eps": 1e-05,
106
+ "act": "gelu_tanh",
107
+ "readout": "last valid input position"
108
+ },
109
+ "batch": {
110
+ "queries": 128,
111
+ "atomic": 32,
112
+ "w2_sources": 16,
113
+ "w2_rendering": "two independent atomic queries per source",
114
+ "d2": 64,
115
+ "loss": "mean of 128 final-answer CEs"
116
+ },
117
+ "optimizer": {
118
+ "name": "AdamW",
119
+ "lr": 0.0001,
120
+ "weight_decay": 0.01,
121
+ "warmup": 2000,
122
+ "schedule": "linear warm-up, then linear decay to 0 at the last update",
123
+ "adam_eps": 1e-08,
124
+ "clip": 1.0,
125
+ "label_smoothing": 0.0,
126
+ "precision": "tf32",
127
+ "evaluation_precision": "fp32"
128
+ },
129
+ "updates": 200000,
130
+ "halting": null,
131
+ "P_budget": null,
132
+ "pack": null,
133
+ "offset_max": 128,
134
+ "offset_dist": "uniform",
135
+ "delay_max": 0,
136
+ "delay_start": 0,
137
+ "persist": 0,
138
+ "anchor": {
139
+ "weight": 1.0,
140
+ "mode": "full",
141
+ "extra_loops": 2
142
+ },
143
+ "consist_loss": {
144
+ "mode": "t0_wait_denoise",
145
+ "components": [
146
+ "wait"
147
+ ],
148
+ "root_shared": false,
149
+ "weight": 1.0,
150
+ "depth": 32,
151
+ "rows": 32,
152
+ "start": 20000,
153
+ "denominator": "rows_times_depth",
154
+ "internal_weight": 0.5,
155
+ "tfire": 0,
156
+ "token_layout": "entity_then_32_relations",
157
+ "constrained_positions": "2..32",
158
+ "input_noise_positions": "2..32",
159
+ "noise_relative_norm": 0.2,
160
+ "noise_kind": "isotropic",
161
+ "noise_scope": "wait",
162
+ "noise_distribution": "normalized Gaussian direction with identical nearest-distance radius",
163
+ "noise_rng": "independent torch.Generator; fill outside CUDA graph; checkpoint state",
164
+ "target": "clean detached input embeddings",
165
+ "input_detached": true,
166
+ "note": "noise mask may include root/frontier; loss mask always only waiting positions 2..32; shallow full anchor unchanged",
167
+ "noise_radius": "uniform_fraction_of_nearest_embedding_distance",
168
+ "noise_alpha_distribution": "Uniform[0, noise_scale)",
169
+ "nearest_candidates": "all_entities_and_relations_excluding_self_and_special_tokens",
170
+ "nearest_metric": "Euclidean L2 of current detached embeddings; direct cdist (no TF32 distance matmul)"
171
+ },
172
+ "format_loss": null,
173
+ "deep": null,
174
+ "prefix_supervision": null,
175
+ "selection": "final checkpoint is primary; best.pt = highest val_d2 accuracy under the arm's own protocol (earliest on ties); deep tests never used",
176
+ "data_hashes": {
177
+ "train_w2": "c17bf311238b3178",
178
+ "train_d2": "e4263675cca07f9d",
179
+ "train_d3": "4f53cda18c2baa0c",
180
+ "val_d2": "2184e30bb884158d",
181
+ "val_w2": "7bf2660c2f724f75",
182
+ "tests": "a66eeb4380ef83fc",
183
+ "vocab": "4aaf5b155079d7a1"
184
+ },
185
+ "atomic_sha256": "4e5f40b8dcd91769b0ada34c3050be09935f0d5e9225edfe95b9d2e7cb4e6e08",
186
+ "source_hash": "84c28353735005c4",
187
+ "cuda_graph": true
188
+ }
for Minegishi/wait_denoising_fixed_t0/metrics.jsonl ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {"update": 1000, "elapsed": 121.90584659576416, "shallow": {"atomic": {"correct": 0.0436}, "val_d2": {"correct": 0.001}, "train_d2_fit": {"correct": 0.0025}}, "checkpoint": {"file": "ckpt_0001000.pt", "update": 1000, "sha256": "b04cbfdb2a1e7a54e1c3857ef6d70d87542c02cd8f03baf4e874d8a998915e97", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.01}, "all_seen|d16": {"n": 100, "correct": 0.0}, "all_seen|d2": {"n": 100, "correct": 0.01}, "all_seen|d3": {"n": 100, "correct": 0.0}, "all_seen|d32": {"n": 100, "correct": 0.0}, "all_seen|d4": {"n": 100, "correct": 0.01}, "all_seen|d64": {"n": 100, "correct": 0.01}, "all_seen|d8": {"n": 100, "correct": 0.0}, "one_new|d128": {"n": 100, "correct": 0.0}, "one_new|d16": {"n": 100, "correct": 0.0}, "one_new|d2": {"n": 100, "correct": 0.0}, "one_new|d3": {"n": 100, "correct": 0.0}, "one_new|d32": {"n": 100, "correct": 0.01}, "one_new|d4": {"n": 100, "correct": 0.0}, "one_new|d64": {"n": 100, "correct": 0.0}, "one_new|d8": {"n": 100, "correct": 0.0}, "random|d128": {"n": 100, "correct": 0.0}, "random|d16": {"n": 100, "correct": 0.0}, "random|d2": {"n": 100, "correct": 0.0}, "random|d3": {"n": 100, "correct": 0.0}, "random|d32": {"n": 100, "correct": 0.0}, "random|d4": {"n": 100, "correct": 0.0}, "random|d64": {"n": 100, "correct": 0.0}, "random|d8": {"n": 100, "correct": 0.0}}, "monitor_budgets": {"d": {"all_seen|d128": 0.01, "all_seen|d16": 0.0, "all_seen|d2": 0.01, "all_seen|d3": 0.0, "all_seen|d32": 0.0, "all_seen|d4": 0.01, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.01, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.0, "random|d64": 0.0, "random|d8": 0.0}, "d+1": {"all_seen|d128": 0.01, "all_seen|d16": 0.0, "all_seen|d2": 0.0, "all_seen|d3": 0.01, "all_seen|d32": 0.0, "all_seen|d4": 0.01, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.01, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.01, "random|d64": 0.0, "random|d8": 0.0}, "d+2": {"all_seen|d128": 0.01, "all_seen|d16": 0.0, "all_seen|d2": 0.0, "all_seen|d3": 0.0, "all_seen|d32": 0.0, "all_seen|d4": 0.01, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.01, "one_new|d32": 0.01, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.0, "random|d64": 0.0, "random|d8": 0.0}, "2d": {"all_seen|d128": 0.01, "all_seen|d16": 0.0, "all_seen|d2": 0.0, "all_seen|d3": 0.0, "all_seen|d32": 0.0, "all_seen|d4": 0.0, "all_seen|d64": 0.0, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.0, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.01, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.0, "random|d64": 0.01, "random|d8": 0.0}}, "counters": {"atomic": 32000, "w2_sources": 16000, "w2_queries": 32000, "d2": 64000, "row_loops": 512000, "R_hist": {}}, "eval_seconds": 29.482265949249268, "peak_mem_gb": 15.119765281677246}
2
+ {"update": 2000, "elapsed": 273.4459092617035, "shallow": {"atomic": {"correct": 0.0889}, "val_d2": {"correct": 0.003}, "train_d2_fit": {"correct": 0.002625}}, "checkpoint": {"file": "ckpt_0002000.pt", "update": 2000, "sha256": "09775249c367ccf150c3fd773f92b295743a30dc7db1bc910b9b1f29ba125f68", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.0}, "all_seen|d16": {"n": 100, "correct": 0.0}, "all_seen|d2": {"n": 100, "correct": 0.0}, "all_seen|d3": {"n": 100, "correct": 0.0}, "all_seen|d32": {"n": 100, "correct": 0.01}, "all_seen|d4": {"n": 100, "correct": 0.0}, "all_seen|d64": {"n": 100, "correct": 0.01}, "all_seen|d8": {"n": 100, "correct": 0.0}, "one_new|d128": {"n": 100, "correct": 0.0}, "one_new|d16": {"n": 100, "correct": 0.0}, "one_new|d2": {"n": 100, "correct": 0.0}, "one_new|d3": {"n": 100, "correct": 0.0}, "one_new|d32": {"n": 100, "correct": 0.0}, "one_new|d4": {"n": 100, "correct": 0.0}, "one_new|d64": {"n": 100, "correct": 0.0}, "one_new|d8": {"n": 100, "correct": 0.0}, "random|d128": {"n": 100, "correct": 0.0}, "random|d16": {"n": 100, "correct": 0.0}, "random|d2": {"n": 100, "correct": 0.0}, "random|d3": {"n": 100, "correct": 0.0}, "random|d32": {"n": 100, "correct": 0.0}, "random|d4": {"n": 100, "correct": 0.01}, "random|d64": {"n": 100, "correct": 0.0}, "random|d8": {"n": 100, "correct": 0.0}}, "monitor_budgets": {"d": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.0, "all_seen|d3": 0.0, "all_seen|d32": 0.01, "all_seen|d4": 0.0, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.0, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.01, "random|d64": 0.0, "random|d8": 0.0}, "d+1": {"all_seen|d128": 0.0, "all_seen|d16": 0.01, "all_seen|d2": 0.0, "all_seen|d3": 0.0, "all_seen|d32": 0.01, "all_seen|d4": 0.0, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.01, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.01, "random|d64": 0.0, "random|d8": 0.0}, "d+2": {"all_seen|d128": 0.0, "all_seen|d16": 0.01, "all_seen|d2": 0.0, "all_seen|d3": 0.02, "all_seen|d32": 0.01, "all_seen|d4": 0.0, "all_seen|d64": 0.01, "all_seen|d8": 0.01, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.01, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.0, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.0, "random|d64": 0.0, "random|d8": 0.0}, "2d": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.0, "all_seen|d3": 0.02, "all_seen|d32": 0.0, "all_seen|d4": 0.0, "all_seen|d64": 0.0, "all_seen|d8": 0.01, "one_new|d128": 0.0, "one_new|d16": 0.01, "one_new|d2": 0.0, "one_new|d3": 0.0, "one_new|d32": 0.0, "one_new|d4": 0.0, "one_new|d64": 0.0, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.01, "random|d2": 0.0, "random|d3": 0.0, "random|d32": 0.0, "random|d4": 0.0, "random|d64": 0.0, "random|d8": 0.0}}, "counters": {"atomic": 64000, "w2_sources": 32000, "w2_queries": 64000, "d2": 128000, "row_loops": 1024000, "R_hist": {}}, "eval_seconds": 28.92754030227661, "peak_mem_gb": 15.119765281677246}
3
+ {"update": 5000, "elapsed": 667.7526476383209, "shallow": {"atomic": {"correct": 0.996}, "val_d2": {"correct": 0.761}, "train_d2_fit": {"correct": 0.803375}}, "checkpoint": {"file": "ckpt_0005000.pt", "update": 5000, "sha256": "cbb6360b83c05c86ba437a856901801e36acf74776beefffe41db5a319305cd6", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.0}, "all_seen|d16": {"n": 100, "correct": 0.0}, "all_seen|d2": {"n": 100, "correct": 0.8}, "all_seen|d3": {"n": 100, "correct": 0.09}, "all_seen|d32": {"n": 100, "correct": 0.0}, "all_seen|d4": {"n": 100, "correct": 0.0}, "all_seen|d64": {"n": 100, "correct": 0.0}, "all_seen|d8": {"n": 100, "correct": 0.0}, "one_new|d128": {"n": 100, "correct": 0.01}, "one_new|d16": {"n": 100, "correct": 0.0}, "one_new|d2": {"n": 100, "correct": 0.83}, "one_new|d3": {"n": 100, "correct": 0.14}, "one_new|d32": {"n": 100, "correct": 0.0}, "one_new|d4": {"n": 100, "correct": 0.0}, "one_new|d64": {"n": 100, "correct": 0.01}, "one_new|d8": {"n": 100, "correct": 0.0}, "random|d128": {"n": 100, "correct": 0.0}, "random|d16": {"n": 100, "correct": 0.01}, "random|d2": {"n": 100, "correct": 0.77}, "random|d3": {"n": 100, "correct": 0.14}, "random|d32": {"n": 100, "correct": 0.0}, "random|d4": {"n": 100, "correct": 0.0}, "random|d64": {"n": 100, "correct": 0.01}, "random|d8": {"n": 100, "correct": 0.0}}, "monitor_budgets": {"d": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.8, "all_seen|d3": 0.09, "all_seen|d32": 0.0, "all_seen|d4": 0.0, "all_seen|d64": 0.0, "all_seen|d8": 0.0, "one_new|d128": 0.01, "one_new|d16": 0.0, "one_new|d2": 0.83, "one_new|d3": 0.14, "one_new|d32": 0.0, "one_new|d4": 0.0, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.01, "random|d2": 0.77, "random|d3": 0.14, "random|d32": 0.0, "random|d4": 0.0, "random|d64": 0.01, "random|d8": 0.0}, "d+1": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.89, "all_seen|d3": 0.41, "all_seen|d32": 0.0, "all_seen|d4": 0.04, "all_seen|d64": 0.0, "all_seen|d8": 0.0, "one_new|d128": 0.01, "one_new|d16": 0.0, "one_new|d2": 0.88, "one_new|d3": 0.39, "one_new|d32": 0.0, "one_new|d4": 0.02, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.01, "random|d2": 0.88, "random|d3": 0.37, "random|d32": 0.0, "random|d4": 0.04, "random|d64": 0.01, "random|d8": 0.0}, "d+2": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.85, "all_seen|d3": 0.49, "all_seen|d32": 0.0, "all_seen|d4": 0.11, "all_seen|d64": 0.0, "all_seen|d8": 0.0, "one_new|d128": 0.01, "one_new|d16": 0.0, "one_new|d2": 0.83, "one_new|d3": 0.45, "one_new|d32": 0.0, "one_new|d4": 0.11, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.01, "random|d2": 0.85, "random|d3": 0.46, "random|d32": 0.0, "random|d4": 0.12, "random|d64": 0.01, "random|d8": 0.0}, "2d": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.85, "all_seen|d3": 0.46, "all_seen|d32": 0.0, "all_seen|d4": 0.14, "all_seen|d64": 0.0, "all_seen|d8": 0.0, "one_new|d128": 0.01, "one_new|d16": 0.0, "one_new|d2": 0.83, "one_new|d3": 0.41, "one_new|d32": 0.0, "one_new|d4": 0.11, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.0, "random|d16": 0.01, "random|d2": 0.85, "random|d3": 0.38, "random|d32": 0.0, "random|d4": 0.1, "random|d64": 0.01, "random|d8": 0.01}}, "counters": {"atomic": 160000, "w2_sources": 80000, "w2_queries": 160000, "d2": 320000, "row_loops": 2560000, "R_hist": {}}, "eval_seconds": 29.225916624069214, "peak_mem_gb": 15.119765281677246}
4
+ {"update": 10000, "elapsed": 1305.9499275684357, "shallow": {"atomic": {"correct": 0.9972}, "val_d2": {"correct": 0.887}, "train_d2_fit": {"correct": 0.903375}}, "checkpoint": {"file": "ckpt_0010000.pt", "update": 10000, "sha256": "aefc1a2e8d13155cf68e5cdbe2b56648b666bede651114928563fe53aee388d8", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.0}, "all_seen|d16": {"n": 100, "correct": 0.0}, "all_seen|d2": {"n": 100, "correct": 0.89}, "all_seen|d3": {"n": 100, "correct": 0.4}, "all_seen|d32": {"n": 100, "correct": 0.0}, "all_seen|d4": {"n": 100, "correct": 0.14}, "all_seen|d64": {"n": 100, "correct": 0.01}, "all_seen|d8": {"n": 100, "correct": 0.0}, "one_new|d128": {"n": 100, "correct": 0.0}, "one_new|d16": {"n": 100, "correct": 0.03}, "one_new|d2": {"n": 100, "correct": 0.96}, "one_new|d3": {"n": 100, "correct": 0.39}, "one_new|d32": {"n": 100, "correct": 0.0}, "one_new|d4": {"n": 100, "correct": 0.1}, "one_new|d64": {"n": 100, "correct": 0.01}, "one_new|d8": {"n": 100, "correct": 0.0}, "random|d128": {"n": 100, "correct": 0.01}, "random|d16": {"n": 100, "correct": 0.0}, "random|d2": {"n": 100, "correct": 0.92}, "random|d3": {"n": 100, "correct": 0.4}, "random|d32": {"n": 100, "correct": 0.0}, "random|d4": {"n": 100, "correct": 0.03}, "random|d64": {"n": 100, "correct": 0.0}, "random|d8": {"n": 100, "correct": 0.01}}, "monitor_budgets": {"d": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.89, "all_seen|d3": 0.4, "all_seen|d32": 0.0, "all_seen|d4": 0.14, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.03, "one_new|d2": 0.96, "one_new|d3": 0.39, "one_new|d32": 0.0, "one_new|d4": 0.1, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.01, "random|d16": 0.0, "random|d2": 0.92, "random|d3": 0.4, "random|d32": 0.0, "random|d4": 0.03, "random|d64": 0.0, "random|d8": 0.01}, "d+1": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.9, "all_seen|d3": 0.63, "all_seen|d32": 0.0, "all_seen|d4": 0.32, "all_seen|d64": 0.01, "all_seen|d8": 0.0, "one_new|d128": 0.0, "one_new|d16": 0.03, "one_new|d2": 0.95, "one_new|d3": 0.59, "one_new|d32": 0.0, "one_new|d4": 0.28, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.01, "random|d16": 0.0, "random|d2": 0.91, "random|d3": 0.6, "random|d32": 0.0, "random|d4": 0.17, "random|d64": 0.0, "random|d8": 0.01}, "d+2": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.83, "all_seen|d3": 0.66, "all_seen|d32": 0.0, "all_seen|d4": 0.39, "all_seen|d64": 0.01, "all_seen|d8": 0.01, "one_new|d128": 0.0, "one_new|d16": 0.03, "one_new|d2": 0.91, "one_new|d3": 0.58, "one_new|d32": 0.0, "one_new|d4": 0.36, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.01, "random|d16": 0.0, "random|d2": 0.89, "random|d3": 0.64, "random|d32": 0.0, "random|d4": 0.23, "random|d64": 0.0, "random|d8": 0.02}, "2d": {"all_seen|d128": 0.0, "all_seen|d16": 0.0, "all_seen|d2": 0.83, "all_seen|d3": 0.56, "all_seen|d32": 0.0, "all_seen|d4": 0.36, "all_seen|d64": 0.01, "all_seen|d8": 0.01, "one_new|d128": 0.0, "one_new|d16": 0.03, "one_new|d2": 0.91, "one_new|d3": 0.56, "one_new|d32": 0.0, "one_new|d4": 0.34, "one_new|d64": 0.01, "one_new|d8": 0.0, "random|d128": 0.01, "random|d16": 0.0, "random|d2": 0.89, "random|d3": 0.58, "random|d32": 0.0, "random|d4": 0.25, "random|d64": 0.0, "random|d8": 0.01}}, "counters": {"atomic": 320000, "w2_sources": 160000, "w2_queries": 320000, "d2": 640000, "row_loops": 5120000, "R_hist": {}}, "eval_seconds": 29.23065757751465, "peak_mem_gb": 15.119765281677246}
5
+ {"update": 20000, "elapsed": 2545.9333856105804, "shallow": {"atomic": {"correct": 0.997}, "val_d2": {"correct": 0.967}, "train_d2_fit": {"correct": 0.967625}}, "checkpoint": {"file": "ckpt_0020000.pt", "update": 20000, "sha256": "107b8652838126379d513366c4e684c2de800ed5582deae281f48b691966e505", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.01}, "all_seen|d16": {"n": 100, "correct": 0.01}, "all_seen|d2": {"n": 100, "correct": 0.96}, "all_seen|d3": {"n": 100, "correct": 0.75}, "all_seen|d32": {"n": 100, "correct": 0.0}, "all_seen|d4": {"n": 100, "correct": 0.47}, "all_seen|d64": {"n": 100, "correct": 0.0}, "all_seen|d8": {"n": 100, "correct": 0.03}, "one_new|d128": {"n": 100, "correct": 0.0}, "one_new|d16": {"n": 100, "correct": 0.0}, "one_new|d2": {"n": 100, "correct": 0.96}, "one_new|d3": {"n": 100, "correct": 0.65}, "one_new|d32": {"n": 100, "correct": 0.0}, "one_new|d4": {"n": 100, "correct": 0.48}, "one_new|d64": {"n": 100, "correct": 0.0}, "one_new|d8": {"n": 100, "correct": 0.02}, "random|d128": {"n": 100, "correct": 0.02}, "random|d16": {"n": 100, "correct": 0.0}, "random|d2": {"n": 100, "correct": 0.97}, "random|d3": {"n": 100, "correct": 0.79}, "random|d32": {"n": 100, "correct": 0.0}, "random|d4": {"n": 100, "correct": 0.44}, "random|d64": {"n": 100, "correct": 0.0}, "random|d8": {"n": 100, "correct": 0.0}}, "monitor_budgets": {"d": {"all_seen|d128": 0.01, "all_seen|d16": 0.01, "all_seen|d2": 0.96, "all_seen|d3": 0.75, "all_seen|d32": 0.0, "all_seen|d4": 0.47, "all_seen|d64": 0.0, "all_seen|d8": 0.03, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.96, "one_new|d3": 0.65, "one_new|d32": 0.0, "one_new|d4": 0.48, "one_new|d64": 0.0, "one_new|d8": 0.02, "random|d128": 0.02, "random|d16": 0.0, "random|d2": 0.97, "random|d3": 0.79, "random|d32": 0.0, "random|d4": 0.44, "random|d64": 0.0, "random|d8": 0.0}, "d+1": {"all_seen|d128": 0.01, "all_seen|d16": 0.01, "all_seen|d2": 0.95, "all_seen|d3": 0.87, "all_seen|d32": 0.0, "all_seen|d4": 0.69, "all_seen|d64": 0.0, "all_seen|d8": 0.07, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.96, "one_new|d3": 0.83, "one_new|d32": 0.0, "one_new|d4": 0.66, "one_new|d64": 0.0, "one_new|d8": 0.04, "random|d128": 0.02, "random|d16": 0.0, "random|d2": 0.96, "random|d3": 0.89, "random|d32": 0.0, "random|d4": 0.68, "random|d64": 0.0, "random|d8": 0.05}, "d+2": {"all_seen|d128": 0.01, "all_seen|d16": 0.01, "all_seen|d2": 0.93, "all_seen|d3": 0.86, "all_seen|d32": 0.0, "all_seen|d4": 0.71, "all_seen|d64": 0.0, "all_seen|d8": 0.09, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.94, "one_new|d3": 0.78, "one_new|d32": 0.0, "one_new|d4": 0.73, "one_new|d64": 0.0, "one_new|d8": 0.06, "random|d128": 0.02, "random|d16": 0.0, "random|d2": 0.95, "random|d3": 0.86, "random|d32": 0.0, "random|d4": 0.72, "random|d64": 0.0, "random|d8": 0.06}, "2d": {"all_seen|d128": 0.01, "all_seen|d16": 0.01, "all_seen|d2": 0.93, "all_seen|d3": 0.84, "all_seen|d32": 0.0, "all_seen|d4": 0.7, "all_seen|d64": 0.0, "all_seen|d8": 0.12, "one_new|d128": 0.0, "one_new|d16": 0.0, "one_new|d2": 0.94, "one_new|d3": 0.76, "one_new|d32": 0.0, "one_new|d4": 0.69, "one_new|d64": 0.0, "one_new|d8": 0.09, "random|d128": 0.02, "random|d16": 0.0, "random|d2": 0.95, "random|d3": 0.82, "random|d32": 0.0, "random|d4": 0.67, "random|d64": 0.0, "random|d8": 0.07}}, "counters": {"atomic": 640000, "w2_sources": 320000, "w2_queries": 640000, "d2": 1280000, "row_loops": 10240000, "R_hist": {}}, "eval_seconds": 29.21105718612671, "peak_mem_gb": 15.119765281677246}
6
+ {"update": 30000, "elapsed": 3785.5603907108307, "shallow": {"atomic": {"correct": 0.9975}, "val_d2": {"correct": 0.984}, "train_d2_fit": {"correct": 0.988375}}, "checkpoint": {"file": "ckpt_0030000.pt", "update": 30000, "sha256": "584e9f1b1fb3f1ce52fa7c95b5a8dc85c8696f4e9aa829c583dc415bd1144f79", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.03}, "all_seen|d16": {"n": 100, "correct": 0.58}, "all_seen|d2": {"n": 100, "correct": 0.99}, "all_seen|d3": {"n": 100, "correct": 0.98}, "all_seen|d32": {"n": 100, "correct": 0.44}, "all_seen|d4": {"n": 100, "correct": 0.92}, "all_seen|d64": {"n": 100, "correct": 0.13}, "all_seen|d8": {"n": 100, "correct": 0.77}, "one_new|d128": {"n": 100, "correct": 0.03}, "one_new|d16": {"n": 100, "correct": 0.59}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 0.96}, "one_new|d32": {"n": 100, "correct": 0.39}, "one_new|d4": {"n": 100, "correct": 0.88}, "one_new|d64": {"n": 100, "correct": 0.1}, "one_new|d8": {"n": 100, "correct": 0.77}, "random|d128": {"n": 100, "correct": 0.0}, "random|d16": {"n": 100, "correct": 0.57}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 0.96}, "random|d32": {"n": 100, "correct": 0.3}, "random|d4": {"n": 100, "correct": 0.89}, "random|d64": {"n": 100, "correct": 0.07}, "random|d8": {"n": 100, "correct": 0.59}}, "monitor_budgets": {"d": {"all_seen|d128": 0.03, "all_seen|d16": 0.58, "all_seen|d2": 0.99, "all_seen|d3": 0.98, "all_seen|d32": 0.44, "all_seen|d4": 0.92, "all_seen|d64": 0.13, "all_seen|d8": 0.77, "one_new|d128": 0.03, "one_new|d16": 0.59, "one_new|d2": 1.0, "one_new|d3": 0.96, "one_new|d32": 0.39, "one_new|d4": 0.88, "one_new|d64": 0.1, "one_new|d8": 0.77, "random|d128": 0.0, "random|d16": 0.57, "random|d2": 1.0, "random|d3": 0.96, "random|d32": 0.3, "random|d4": 0.89, "random|d64": 0.07, "random|d8": 0.59}, "d+1": {"all_seen|d128": 0.03, "all_seen|d16": 0.58, "all_seen|d2": 0.97, "all_seen|d3": 0.96, "all_seen|d32": 0.44, "all_seen|d4": 0.92, "all_seen|d64": 0.12, "all_seen|d8": 0.74, "one_new|d128": 0.03, "one_new|d16": 0.6, "one_new|d2": 1.0, "one_new|d3": 0.98, "one_new|d32": 0.39, "one_new|d4": 0.87, "one_new|d64": 0.11, "one_new|d8": 0.76, "random|d128": 0.0, "random|d16": 0.58, "random|d2": 1.0, "random|d3": 0.94, "random|d32": 0.3, "random|d4": 0.9, "random|d64": 0.06, "random|d8": 0.58}, "d+2": {"all_seen|d128": 0.03, "all_seen|d16": 0.58, "all_seen|d2": 0.97, "all_seen|d3": 0.95, "all_seen|d32": 0.45, "all_seen|d4": 0.91, "all_seen|d64": 0.12, "all_seen|d8": 0.74, "one_new|d128": 0.03, "one_new|d16": 0.6, "one_new|d2": 0.98, "one_new|d3": 0.96, "one_new|d32": 0.39, "one_new|d4": 0.89, "one_new|d64": 0.11, "one_new|d8": 0.75, "random|d128": 0.0, "random|d16": 0.55, "random|d2": 0.97, "random|d3": 0.91, "random|d32": 0.3, "random|d4": 0.88, "random|d64": 0.05, "random|d8": 0.57}, "2d": {"all_seen|d128": 0.03, "all_seen|d16": 0.57, "all_seen|d2": 0.97, "all_seen|d3": 0.93, "all_seen|d32": 0.44, "all_seen|d4": 0.85, "all_seen|d64": 0.12, "all_seen|d8": 0.7, "one_new|d128": 0.03, "one_new|d16": 0.57, "one_new|d2": 0.98, "one_new|d3": 0.95, "one_new|d32": 0.38, "one_new|d4": 0.83, "one_new|d64": 0.1, "one_new|d8": 0.72, "random|d128": 0.0, "random|d16": 0.5, "random|d2": 0.97, "random|d3": 0.91, "random|d32": 0.29, "random|d4": 0.82, "random|d64": 0.05, "random|d8": 0.54}}, "counters": {"atomic": 960000, "w2_sources": 480000, "w2_queries": 960000, "d2": 1920000, "row_loops": 15360000, "R_hist": {}}, "eval_seconds": 29.173584938049316, "peak_mem_gb": 15.119765281677246}
7
+ {"update": 40000, "elapsed": 5026.457386732101, "shallow": {"atomic": {"correct": 0.9971}, "val_d2": {"correct": 0.99}, "train_d2_fit": {"correct": 0.99}}, "checkpoint": {"file": "ckpt_0040000.pt", "update": 40000, "sha256": "259bdfe027e50d18508c6e332add29ebca312fce38c21f6fce9bba2be6b86ce5", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.0}, "all_seen|d16": {"n": 100, "correct": 0.65}, "all_seen|d2": {"n": 100, "correct": 0.97}, "all_seen|d3": {"n": 100, "correct": 0.99}, "all_seen|d32": {"n": 100, "correct": 0.31}, "all_seen|d4": {"n": 100, "correct": 0.96}, "all_seen|d64": {"n": 100, "correct": 0.02}, "all_seen|d8": {"n": 100, "correct": 0.86}, "one_new|d128": {"n": 100, "correct": 0.0}, "one_new|d16": {"n": 100, "correct": 0.71}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 0.93}, "one_new|d32": {"n": 100, "correct": 0.31}, "one_new|d4": {"n": 100, "correct": 0.93}, "one_new|d64": {"n": 100, "correct": 0.04}, "one_new|d8": {"n": 100, "correct": 0.89}, "random|d128": {"n": 100, "correct": 0.0}, "random|d16": {"n": 100, "correct": 0.52}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 0.96}, "random|d32": {"n": 100, "correct": 0.27}, "random|d4": {"n": 100, "correct": 0.96}, "random|d64": {"n": 100, "correct": 0.06}, "random|d8": {"n": 100, "correct": 0.89}}, "monitor_budgets": {"d": {"all_seen|d128": 0.0, "all_seen|d16": 0.65, "all_seen|d2": 0.97, "all_seen|d3": 0.99, "all_seen|d32": 0.31, "all_seen|d4": 0.96, "all_seen|d64": 0.02, "all_seen|d8": 0.86, "one_new|d128": 0.0, "one_new|d16": 0.71, "one_new|d2": 1.0, "one_new|d3": 0.93, "one_new|d32": 0.31, "one_new|d4": 0.93, "one_new|d64": 0.04, "one_new|d8": 0.89, "random|d128": 0.0, "random|d16": 0.52, "random|d2": 1.0, "random|d3": 0.96, "random|d32": 0.27, "random|d4": 0.96, "random|d64": 0.06, "random|d8": 0.89}, "d+1": {"all_seen|d128": 0.01, "all_seen|d16": 0.66, "all_seen|d2": 0.97, "all_seen|d3": 0.98, "all_seen|d32": 0.31, "all_seen|d4": 0.94, "all_seen|d64": 0.02, "all_seen|d8": 0.86, "one_new|d128": 0.0, "one_new|d16": 0.7, "one_new|d2": 0.97, "one_new|d3": 0.93, "one_new|d32": 0.31, "one_new|d4": 0.93, "one_new|d64": 0.04, "one_new|d8": 0.88, "random|d128": 0.0, "random|d16": 0.52, "random|d2": 1.0, "random|d3": 0.96, "random|d32": 0.26, "random|d4": 0.95, "random|d64": 0.06, "random|d8": 0.89}, "d+2": {"all_seen|d128": 0.0, "all_seen|d16": 0.65, "all_seen|d2": 0.97, "all_seen|d3": 0.98, "all_seen|d32": 0.3, "all_seen|d4": 0.94, "all_seen|d64": 0.02, "all_seen|d8": 0.86, "one_new|d128": 0.0, "one_new|d16": 0.7, "one_new|d2": 0.95, "one_new|d3": 0.92, "one_new|d32": 0.31, "one_new|d4": 0.93, "one_new|d64": 0.04, "one_new|d8": 0.88, "random|d128": 0.0, "random|d16": 0.52, "random|d2": 1.0, "random|d3": 0.96, "random|d32": 0.26, "random|d4": 0.94, "random|d64": 0.06, "random|d8": 0.89}, "2d": {"all_seen|d128": 0.0, "all_seen|d16": 0.61, "all_seen|d2": 0.97, "all_seen|d3": 0.96, "all_seen|d32": 0.3, "all_seen|d4": 0.86, "all_seen|d64": 0.02, "all_seen|d8": 0.79, "one_new|d128": 0.0, "one_new|d16": 0.69, "one_new|d2": 0.95, "one_new|d3": 0.89, "one_new|d32": 0.31, "one_new|d4": 0.89, "one_new|d64": 0.04, "one_new|d8": 0.84, "random|d128": 0.0, "random|d16": 0.5, "random|d2": 1.0, "random|d3": 0.94, "random|d32": 0.24, "random|d4": 0.9, "random|d64": 0.05, "random|d8": 0.86}}, "counters": {"atomic": 1280000, "w2_sources": 640000, "w2_queries": 1280000, "d2": 2560000, "row_loops": 20480000, "R_hist": {}}, "eval_seconds": 29.214928150177002, "peak_mem_gb": 15.119765281677246}
8
+ {"update": 50000, "elapsed": 6266.025934457779, "shallow": {"atomic": {"correct": 0.9973}, "val_d2": {"correct": 0.991}, "train_d2_fit": {"correct": 0.992875}}, "checkpoint": {"file": "ckpt_0050000.pt", "update": 50000, "sha256": "f8e4cd617629f536541f238f8e7df5e3b89f15db47c4c3025b5b7757a58a51fa", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.27}, "all_seen|d16": {"n": 100, "correct": 0.84}, "all_seen|d2": {"n": 100, "correct": 0.99}, "all_seen|d3": {"n": 100, "correct": 0.99}, "all_seen|d32": {"n": 100, "correct": 0.64}, "all_seen|d4": {"n": 100, "correct": 0.97}, "all_seen|d64": {"n": 100, "correct": 0.55}, "all_seen|d8": {"n": 100, "correct": 0.92}, "one_new|d128": {"n": 100, "correct": 0.32}, "one_new|d16": {"n": 100, "correct": 0.84}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 0.95}, "one_new|d32": {"n": 100, "correct": 0.68}, "one_new|d4": {"n": 100, "correct": 0.97}, "one_new|d64": {"n": 100, "correct": 0.58}, "one_new|d8": {"n": 100, "correct": 0.92}, "random|d128": {"n": 100, "correct": 0.19}, "random|d16": {"n": 100, "correct": 0.78}, "random|d2": {"n": 100, "correct": 0.99}, "random|d3": {"n": 100, "correct": 0.99}, "random|d32": {"n": 100, "correct": 0.7}, "random|d4": {"n": 100, "correct": 0.96}, "random|d64": {"n": 100, "correct": 0.46}, "random|d8": {"n": 100, "correct": 0.96}}, "monitor_budgets": {"d": {"all_seen|d128": 0.27, "all_seen|d16": 0.84, "all_seen|d2": 0.99, "all_seen|d3": 0.99, "all_seen|d32": 0.64, "all_seen|d4": 0.97, "all_seen|d64": 0.55, "all_seen|d8": 0.92, "one_new|d128": 0.32, "one_new|d16": 0.84, "one_new|d2": 1.0, "one_new|d3": 0.95, "one_new|d32": 0.68, "one_new|d4": 0.97, "one_new|d64": 0.58, "one_new|d8": 0.92, "random|d128": 0.19, "random|d16": 0.78, "random|d2": 0.99, "random|d3": 0.99, "random|d32": 0.7, "random|d4": 0.96, "random|d64": 0.46, "random|d8": 0.96}, "d+1": {"all_seen|d128": 0.27, "all_seen|d16": 0.83, "all_seen|d2": 0.99, "all_seen|d3": 0.99, "all_seen|d32": 0.64, "all_seen|d4": 0.97, "all_seen|d64": 0.55, "all_seen|d8": 0.93, "one_new|d128": 0.32, "one_new|d16": 0.84, "one_new|d2": 0.99, "one_new|d3": 0.93, "one_new|d32": 0.67, "one_new|d4": 0.98, "one_new|d64": 0.58, "one_new|d8": 0.89, "random|d128": 0.18, "random|d16": 0.79, "random|d2": 0.99, "random|d3": 0.97, "random|d32": 0.7, "random|d4": 0.94, "random|d64": 0.46, "random|d8": 0.95}, "d+2": {"all_seen|d128": 0.27, "all_seen|d16": 0.83, "all_seen|d2": 0.99, "all_seen|d3": 0.99, "all_seen|d32": 0.64, "all_seen|d4": 0.97, "all_seen|d64": 0.55, "all_seen|d8": 0.92, "one_new|d128": 0.32, "one_new|d16": 0.84, "one_new|d2": 0.99, "one_new|d3": 0.93, "one_new|d32": 0.66, "one_new|d4": 0.97, "one_new|d64": 0.58, "one_new|d8": 0.88, "random|d128": 0.18, "random|d16": 0.78, "random|d2": 0.98, "random|d3": 0.96, "random|d32": 0.7, "random|d4": 0.93, "random|d64": 0.46, "random|d8": 0.94}, "2d": {"all_seen|d128": 0.27, "all_seen|d16": 0.82, "all_seen|d2": 0.99, "all_seen|d3": 0.96, "all_seen|d32": 0.64, "all_seen|d4": 0.91, "all_seen|d64": 0.53, "all_seen|d8": 0.91, "one_new|d128": 0.32, "one_new|d16": 0.83, "one_new|d2": 0.99, "one_new|d3": 0.92, "one_new|d32": 0.62, "one_new|d4": 0.83, "one_new|d64": 0.56, "one_new|d8": 0.83, "random|d128": 0.17, "random|d16": 0.74, "random|d2": 0.98, "random|d3": 0.94, "random|d32": 0.68, "random|d4": 0.9, "random|d64": 0.45, "random|d8": 0.92}}, "counters": {"atomic": 1600000, "w2_sources": 800000, "w2_queries": 1600000, "d2": 3200000, "row_loops": 25600000, "R_hist": {}}, "eval_seconds": 29.140637397766113, "peak_mem_gb": 15.119765281677246}
9
+ {"update": 60000, "elapsed": 7508.697112560272, "shallow": {"atomic": {"correct": 0.9989}, "val_d2": {"correct": 0.995}, "train_d2_fit": {"correct": 0.999}}, "checkpoint": {"file": "ckpt_0060000.pt", "update": 60000, "sha256": "f6dd0d27bcb48a93839be15dbccbf37b20db15b5bb5144ba8ef5eda581a0176b", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.46}, "all_seen|d16": {"n": 100, "correct": 0.86}, "all_seen|d2": {"n": 100, "correct": 0.99}, "all_seen|d3": {"n": 100, "correct": 0.99}, "all_seen|d32": {"n": 100, "correct": 0.74}, "all_seen|d4": {"n": 100, "correct": 0.96}, "all_seen|d64": {"n": 100, "correct": 0.57}, "all_seen|d8": {"n": 100, "correct": 0.93}, "one_new|d128": {"n": 100, "correct": 0.41}, "one_new|d16": {"n": 100, "correct": 0.88}, "one_new|d2": {"n": 100, "correct": 0.99}, "one_new|d3": {"n": 100, "correct": 0.99}, "one_new|d32": {"n": 100, "correct": 0.85}, "one_new|d4": {"n": 100, "correct": 0.98}, "one_new|d64": {"n": 100, "correct": 0.61}, "one_new|d8": {"n": 100, "correct": 0.94}, "random|d128": {"n": 100, "correct": 0.33}, "random|d16": {"n": 100, "correct": 0.89}, "random|d2": {"n": 100, "correct": 0.99}, "random|d3": {"n": 100, "correct": 0.98}, "random|d32": {"n": 100, "correct": 0.73}, "random|d4": {"n": 100, "correct": 0.97}, "random|d64": {"n": 100, "correct": 0.71}, "random|d8": {"n": 100, "correct": 0.92}}, "monitor_budgets": {"d": {"all_seen|d128": 0.46, "all_seen|d16": 0.86, "all_seen|d2": 0.99, "all_seen|d3": 0.99, "all_seen|d32": 0.74, "all_seen|d4": 0.96, "all_seen|d64": 0.57, "all_seen|d8": 0.93, "one_new|d128": 0.41, "one_new|d16": 0.88, "one_new|d2": 0.99, "one_new|d3": 0.99, "one_new|d32": 0.85, "one_new|d4": 0.98, "one_new|d64": 0.61, "one_new|d8": 0.94, "random|d128": 0.33, "random|d16": 0.89, "random|d2": 0.99, "random|d3": 0.98, "random|d32": 0.73, "random|d4": 0.97, "random|d64": 0.71, "random|d8": 0.92}, "d+1": {"all_seen|d128": 0.46, "all_seen|d16": 0.86, "all_seen|d2": 0.99, "all_seen|d3": 0.99, "all_seen|d32": 0.74, "all_seen|d4": 0.95, "all_seen|d64": 0.57, "all_seen|d8": 0.94, "one_new|d128": 0.41, "one_new|d16": 0.88, "one_new|d2": 0.99, "one_new|d3": 0.97, "one_new|d32": 0.85, "one_new|d4": 0.97, "one_new|d64": 0.61, "one_new|d8": 0.94, "random|d128": 0.33, "random|d16": 0.89, "random|d2": 0.98, "random|d3": 0.98, "random|d32": 0.73, "random|d4": 0.97, "random|d64": 0.71, "random|d8": 0.92}, "d+2": {"all_seen|d128": 0.46, "all_seen|d16": 0.86, "all_seen|d2": 0.99, "all_seen|d3": 0.99, "all_seen|d32": 0.74, "all_seen|d4": 0.95, "all_seen|d64": 0.57, "all_seen|d8": 0.93, "one_new|d128": 0.41, "one_new|d16": 0.88, "one_new|d2": 0.98, "one_new|d3": 0.97, "one_new|d32": 0.85, "one_new|d4": 0.96, "one_new|d64": 0.61, "one_new|d8": 0.94, "random|d128": 0.33, "random|d16": 0.89, "random|d2": 0.98, "random|d3": 0.98, "random|d32": 0.73, "random|d4": 0.97, "random|d64": 0.71, "random|d8": 0.92}, "2d": {"all_seen|d128": 0.44, "all_seen|d16": 0.79, "all_seen|d2": 0.99, "all_seen|d3": 0.98, "all_seen|d32": 0.73, "all_seen|d4": 0.91, "all_seen|d64": 0.54, "all_seen|d8": 0.87, "one_new|d128": 0.41, "one_new|d16": 0.85, "one_new|d2": 0.98, "one_new|d3": 0.96, "one_new|d32": 0.84, "one_new|d4": 0.91, "one_new|d64": 0.58, "one_new|d8": 0.92, "random|d128": 0.32, "random|d16": 0.85, "random|d2": 0.98, "random|d3": 0.98, "random|d32": 0.7, "random|d4": 0.95, "random|d64": 0.69, "random|d8": 0.89}}, "counters": {"atomic": 1920000, "w2_sources": 960000, "w2_queries": 1920000, "d2": 3840000, "row_loops": 30720000, "R_hist": {}}, "eval_seconds": 29.140647888183594, "peak_mem_gb": 15.119765281677246}
10
+ {"update": 70000, "elapsed": 8746.916563987732, "shallow": {"atomic": {"correct": 0.9974}, "val_d2": {"correct": 0.996}, "train_d2_fit": {"correct": 0.98975}}, "checkpoint": {"file": "ckpt_0070000.pt", "update": 70000, "sha256": "afe2b3f44168be01e54532886d2d1439bcb167f0f437c826c45e14668ce1caa0", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.43}, "all_seen|d16": {"n": 100, "correct": 0.9}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 0.98}, "all_seen|d32": {"n": 100, "correct": 0.74}, "all_seen|d4": {"n": 100, "correct": 0.98}, "all_seen|d64": {"n": 100, "correct": 0.58}, "all_seen|d8": {"n": 100, "correct": 0.98}, "one_new|d128": {"n": 100, "correct": 0.49}, "one_new|d16": {"n": 100, "correct": 0.88}, "one_new|d2": {"n": 100, "correct": 0.99}, "one_new|d3": {"n": 100, "correct": 0.98}, "one_new|d32": {"n": 100, "correct": 0.77}, "one_new|d4": {"n": 100, "correct": 0.99}, "one_new|d64": {"n": 100, "correct": 0.53}, "one_new|d8": {"n": 100, "correct": 0.95}, "random|d128": {"n": 100, "correct": 0.34}, "random|d16": {"n": 100, "correct": 0.86}, "random|d2": {"n": 100, "correct": 0.99}, "random|d3": {"n": 100, "correct": 0.99}, "random|d32": {"n": 100, "correct": 0.66}, "random|d4": {"n": 100, "correct": 0.99}, "random|d64": {"n": 100, "correct": 0.59}, "random|d8": {"n": 100, "correct": 0.95}}, "monitor_budgets": {"d": {"all_seen|d128": 0.43, "all_seen|d16": 0.9, "all_seen|d2": 1.0, "all_seen|d3": 0.98, "all_seen|d32": 0.74, "all_seen|d4": 0.98, "all_seen|d64": 0.58, "all_seen|d8": 0.98, "one_new|d128": 0.49, "one_new|d16": 0.88, "one_new|d2": 0.99, "one_new|d3": 0.98, "one_new|d32": 0.77, "one_new|d4": 0.99, "one_new|d64": 0.53, "one_new|d8": 0.95, "random|d128": 0.34, "random|d16": 0.86, "random|d2": 0.99, "random|d3": 0.99, "random|d32": 0.66, "random|d4": 0.99, "random|d64": 0.59, "random|d8": 0.95}, "d+1": {"all_seen|d128": 0.43, "all_seen|d16": 0.9, "all_seen|d2": 1.0, "all_seen|d3": 0.99, "all_seen|d32": 0.74, "all_seen|d4": 0.98, "all_seen|d64": 0.58, "all_seen|d8": 0.98, "one_new|d128": 0.49, "one_new|d16": 0.89, "one_new|d2": 0.99, "one_new|d3": 0.97, "one_new|d32": 0.77, "one_new|d4": 0.96, "one_new|d64": 0.53, "one_new|d8": 0.95, "random|d128": 0.33, "random|d16": 0.85, "random|d2": 0.98, "random|d3": 0.99, "random|d32": 0.66, "random|d4": 0.99, "random|d64": 0.59, "random|d8": 0.95}, "d+2": {"all_seen|d128": 0.43, "all_seen|d16": 0.89, "all_seen|d2": 0.99, "all_seen|d3": 0.97, "all_seen|d32": 0.74, "all_seen|d4": 0.98, "all_seen|d64": 0.58, "all_seen|d8": 0.98, "one_new|d128": 0.49, "one_new|d16": 0.88, "one_new|d2": 0.99, "one_new|d3": 0.97, "one_new|d32": 0.77, "one_new|d4": 0.96, "one_new|d64": 0.53, "one_new|d8": 0.95, "random|d128": 0.33, "random|d16": 0.85, "random|d2": 0.98, "random|d3": 0.98, "random|d32": 0.66, "random|d4": 1.0, "random|d64": 0.59, "random|d8": 0.95}, "2d": {"all_seen|d128": 0.38, "all_seen|d16": 0.74, "all_seen|d2": 0.99, "all_seen|d3": 0.95, "all_seen|d32": 0.66, "all_seen|d4": 0.92, "all_seen|d64": 0.53, "all_seen|d8": 0.82, "one_new|d128": 0.45, "one_new|d16": 0.76, "one_new|d2": 0.99, "one_new|d3": 0.88, "one_new|d32": 0.69, "one_new|d4": 0.86, "one_new|d64": 0.43, "one_new|d8": 0.86, "random|d128": 0.3, "random|d16": 0.75, "random|d2": 0.98, "random|d3": 0.95, "random|d32": 0.53, "random|d4": 0.95, "random|d64": 0.52, "random|d8": 0.84}}, "counters": {"atomic": 2240000, "w2_sources": 1120000, "w2_queries": 2240000, "d2": 4480000, "row_loops": 35840000, "R_hist": {}}, "eval_seconds": 29.155139446258545, "peak_mem_gb": 15.119765281677246}
11
+ {"update": 80000, "elapsed": 9985.045416116714, "shallow": {"atomic": {"correct": 0.999}, "val_d2": {"correct": 0.997}, "train_d2_fit": {"correct": 0.992875}}, "checkpoint": {"file": "ckpt_0080000.pt", "update": 80000, "sha256": "2800b11b460d0e5de02d7009c25a0e4d93b555c7c63f10a59e89739383043118", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.44}, "all_seen|d16": {"n": 100, "correct": 0.81}, "all_seen|d2": {"n": 100, "correct": 0.99}, "all_seen|d3": {"n": 100, "correct": 0.98}, "all_seen|d32": {"n": 100, "correct": 0.73}, "all_seen|d4": {"n": 100, "correct": 0.98}, "all_seen|d64": {"n": 100, "correct": 0.6}, "all_seen|d8": {"n": 100, "correct": 0.98}, "one_new|d128": {"n": 100, "correct": 0.42}, "one_new|d16": {"n": 100, "correct": 0.86}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 0.98}, "one_new|d32": {"n": 100, "correct": 0.75}, "one_new|d4": {"n": 100, "correct": 0.97}, "one_new|d64": {"n": 100, "correct": 0.65}, "one_new|d8": {"n": 100, "correct": 0.94}, "random|d128": {"n": 100, "correct": 0.47}, "random|d16": {"n": 100, "correct": 0.84}, "random|d2": {"n": 100, "correct": 0.99}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 0.66}, "random|d4": {"n": 100, "correct": 0.96}, "random|d64": {"n": 100, "correct": 0.62}, "random|d8": {"n": 100, "correct": 0.92}}, "monitor_budgets": {"d": {"all_seen|d128": 0.44, "all_seen|d16": 0.81, "all_seen|d2": 0.99, "all_seen|d3": 0.98, "all_seen|d32": 0.73, "all_seen|d4": 0.98, "all_seen|d64": 0.6, "all_seen|d8": 0.98, "one_new|d128": 0.42, "one_new|d16": 0.86, "one_new|d2": 1.0, "one_new|d3": 0.98, "one_new|d32": 0.75, "one_new|d4": 0.97, "one_new|d64": 0.65, "one_new|d8": 0.94, "random|d128": 0.47, "random|d16": 0.84, "random|d2": 0.99, "random|d3": 1.0, "random|d32": 0.66, "random|d4": 0.96, "random|d64": 0.62, "random|d8": 0.92}, "d+1": {"all_seen|d128": 0.44, "all_seen|d16": 0.8, "all_seen|d2": 1.0, "all_seen|d3": 0.98, "all_seen|d32": 0.73, "all_seen|d4": 0.97, "all_seen|d64": 0.6, "all_seen|d8": 0.98, "one_new|d128": 0.42, "one_new|d16": 0.84, "one_new|d2": 1.0, "one_new|d3": 0.98, "one_new|d32": 0.75, "one_new|d4": 0.96, "one_new|d64": 0.65, "one_new|d8": 0.94, "random|d128": 0.47, "random|d16": 0.83, "random|d2": 0.98, "random|d3": 1.0, "random|d32": 0.65, "random|d4": 0.97, "random|d64": 0.62, "random|d8": 0.91}, "d+2": {"all_seen|d128": 0.44, "all_seen|d16": 0.79, "all_seen|d2": 0.99, "all_seen|d3": 0.96, "all_seen|d32": 0.73, "all_seen|d4": 0.96, "all_seen|d64": 0.59, "all_seen|d8": 0.98, "one_new|d128": 0.42, "one_new|d16": 0.84, "one_new|d2": 1.0, "one_new|d3": 0.97, "one_new|d32": 0.74, "one_new|d4": 0.96, "one_new|d64": 0.65, "one_new|d8": 0.93, "random|d128": 0.47, "random|d16": 0.83, "random|d2": 0.98, "random|d3": 0.99, "random|d32": 0.65, "random|d4": 0.96, "random|d64": 0.62, "random|d8": 0.91}, "2d": {"all_seen|d128": 0.42, "all_seen|d16": 0.69, "all_seen|d2": 0.99, "all_seen|d3": 0.94, "all_seen|d32": 0.66, "all_seen|d4": 0.84, "all_seen|d64": 0.5, "all_seen|d8": 0.76, "one_new|d128": 0.41, "one_new|d16": 0.74, "one_new|d2": 1.0, "one_new|d3": 0.93, "one_new|d32": 0.68, "one_new|d4": 0.86, "one_new|d64": 0.58, "one_new|d8": 0.75, "random|d128": 0.41, "random|d16": 0.67, "random|d2": 0.98, "random|d3": 0.91, "random|d32": 0.53, "random|d4": 0.79, "random|d64": 0.59, "random|d8": 0.79}}, "counters": {"atomic": 2560000, "w2_sources": 1280000, "w2_queries": 2560000, "d2": 5120000, "row_loops": 40960000, "R_hist": {}}, "eval_seconds": 29.24917984008789, "peak_mem_gb": 15.119765281677246}
12
+ {"update": 90000, "elapsed": 11227.305730342865, "shallow": {"atomic": {"correct": 0.9987}, "val_d2": {"correct": 0.995}, "train_d2_fit": {"correct": 0.99875}}, "checkpoint": {"file": "ckpt_0090000.pt", "update": 90000, "sha256": "61ba4ff1b03eed8710610004e9b1e68db3886813e2fbd68c83d58da444a9dafe", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.67}, "all_seen|d16": {"n": 100, "correct": 0.9}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 0.99}, "all_seen|d32": {"n": 100, "correct": 0.89}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 0.82}, "all_seen|d8": {"n": 100, "correct": 0.95}, "one_new|d128": {"n": 100, "correct": 0.59}, "one_new|d16": {"n": 100, "correct": 0.93}, "one_new|d2": {"n": 100, "correct": 0.99}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 0.84}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 0.86}, "one_new|d8": {"n": 100, "correct": 0.99}, "random|d128": {"n": 100, "correct": 0.65}, "random|d16": {"n": 100, "correct": 0.94}, "random|d2": {"n": 100, "correct": 0.99}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 0.88}, "random|d4": {"n": 100, "correct": 0.99}, "random|d64": {"n": 100, "correct": 0.79}, "random|d8": {"n": 100, "correct": 0.98}}, "monitor_budgets": {"d": {"all_seen|d128": 0.67, "all_seen|d16": 0.9, "all_seen|d2": 1.0, "all_seen|d3": 0.99, "all_seen|d32": 0.89, "all_seen|d4": 1.0, "all_seen|d64": 0.82, "all_seen|d8": 0.95, "one_new|d128": 0.59, "one_new|d16": 0.93, "one_new|d2": 0.99, "one_new|d3": 1.0, "one_new|d32": 0.84, "one_new|d4": 1.0, "one_new|d64": 0.86, "one_new|d8": 0.99, "random|d128": 0.65, "random|d16": 0.94, "random|d2": 0.99, "random|d3": 1.0, "random|d32": 0.88, "random|d4": 0.99, "random|d64": 0.79, "random|d8": 0.98}, "d+1": {"all_seen|d128": 0.67, "all_seen|d16": 0.9, "all_seen|d2": 1.0, "all_seen|d3": 0.98, "all_seen|d32": 0.89, "all_seen|d4": 1.0, "all_seen|d64": 0.81, "all_seen|d8": 0.95, "one_new|d128": 0.59, "one_new|d16": 0.93, "one_new|d2": 0.99, "one_new|d3": 0.99, "one_new|d32": 0.84, "one_new|d4": 1.0, "one_new|d64": 0.86, "one_new|d8": 0.99, "random|d128": 0.64, "random|d16": 0.94, "random|d2": 0.99, "random|d3": 1.0, "random|d32": 0.88, "random|d4": 0.98, "random|d64": 0.79, "random|d8": 0.98}, "d+2": {"all_seen|d128": 0.67, "all_seen|d16": 0.89, "all_seen|d2": 0.98, "all_seen|d3": 0.97, "all_seen|d32": 0.89, "all_seen|d4": 1.0, "all_seen|d64": 0.81, "all_seen|d8": 0.96, "one_new|d128": 0.59, "one_new|d16": 0.93, "one_new|d2": 0.98, "one_new|d3": 0.99, "one_new|d32": 0.84, "one_new|d4": 1.0, "one_new|d64": 0.86, "one_new|d8": 0.99, "random|d128": 0.64, "random|d16": 0.93, "random|d2": 0.99, "random|d3": 1.0, "random|d32": 0.88, "random|d4": 0.98, "random|d64": 0.79, "random|d8": 0.98}, "2d": {"all_seen|d128": 0.67, "all_seen|d16": 0.88, "all_seen|d2": 0.98, "all_seen|d3": 0.97, "all_seen|d32": 0.84, "all_seen|d4": 0.93, "all_seen|d64": 0.78, "all_seen|d8": 0.93, "one_new|d128": 0.57, "one_new|d16": 0.91, "one_new|d2": 0.98, "one_new|d3": 0.99, "one_new|d32": 0.84, "one_new|d4": 0.99, "one_new|d64": 0.84, "one_new|d8": 0.95, "random|d128": 0.64, "random|d16": 0.91, "random|d2": 0.99, "random|d3": 1.0, "random|d32": 0.86, "random|d4": 0.96, "random|d64": 0.76, "random|d8": 0.97}}, "counters": {"atomic": 2880000, "w2_sources": 1440000, "w2_queries": 2880000, "d2": 5760000, "row_loops": 46080000, "R_hist": {}}, "eval_seconds": 29.153381824493408, "peak_mem_gb": 15.119765281677246}
13
+ {"update": 100000, "elapsed": 12465.946828842163, "shallow": {"atomic": {"correct": 0.9997}, "val_d2": {"correct": 0.996}, "train_d2_fit": {"correct": 0.99875}}, "checkpoint": {"file": "ckpt_0100000.pt", "update": 100000, "sha256": "e600093705a1477f6a092135438a4346dd1a4bc5c97430f2d179fdbb7b73a8d7", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.82}, "all_seen|d16": {"n": 100, "correct": 0.95}, "all_seen|d2": {"n": 100, "correct": 0.98}, "all_seen|d3": {"n": 100, "correct": 0.99}, "all_seen|d32": {"n": 100, "correct": 0.89}, "all_seen|d4": {"n": 100, "correct": 0.99}, "all_seen|d64": {"n": 100, "correct": 0.88}, "all_seen|d8": {"n": 100, "correct": 0.98}, "one_new|d128": {"n": 100, "correct": 0.79}, "one_new|d16": {"n": 100, "correct": 0.95}, "one_new|d2": {"n": 100, "correct": 0.99}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 0.92}, "one_new|d4": {"n": 100, "correct": 0.99}, "one_new|d64": {"n": 100, "correct": 0.93}, "one_new|d8": {"n": 100, "correct": 0.97}, "random|d128": {"n": 100, "correct": 0.9}, "random|d16": {"n": 100, "correct": 0.92}, "random|d2": {"n": 100, "correct": 0.99}, "random|d3": {"n": 100, "correct": 0.98}, "random|d32": {"n": 100, "correct": 0.91}, "random|d4": {"n": 100, "correct": 0.99}, "random|d64": {"n": 100, "correct": 0.87}, "random|d8": {"n": 100, "correct": 0.99}}, "monitor_budgets": {"d": {"all_seen|d128": 0.82, "all_seen|d16": 0.95, "all_seen|d2": 0.98, "all_seen|d3": 0.99, "all_seen|d32": 0.89, "all_seen|d4": 0.99, "all_seen|d64": 0.88, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.95, "one_new|d2": 0.99, "one_new|d3": 1.0, "one_new|d32": 0.92, "one_new|d4": 0.99, "one_new|d64": 0.93, "one_new|d8": 0.97, "random|d128": 0.9, "random|d16": 0.92, "random|d2": 0.99, "random|d3": 0.98, "random|d32": 0.91, "random|d4": 0.99, "random|d64": 0.87, "random|d8": 0.99}, "d+1": {"all_seen|d128": 0.82, "all_seen|d16": 0.95, "all_seen|d2": 0.98, "all_seen|d3": 0.98, "all_seen|d32": 0.89, "all_seen|d4": 0.99, "all_seen|d64": 0.88, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.95, "one_new|d2": 0.99, "one_new|d3": 1.0, "one_new|d32": 0.92, "one_new|d4": 0.99, "one_new|d64": 0.93, "one_new|d8": 0.97, "random|d128": 0.9, "random|d16": 0.92, "random|d2": 0.99, "random|d3": 0.97, "random|d32": 0.91, "random|d4": 0.99, "random|d64": 0.87, "random|d8": 0.99}, "d+2": {"all_seen|d128": 0.82, "all_seen|d16": 0.95, "all_seen|d2": 0.98, "all_seen|d3": 0.98, "all_seen|d32": 0.89, "all_seen|d4": 0.99, "all_seen|d64": 0.88, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.95, "one_new|d2": 0.99, "one_new|d3": 1.0, "one_new|d32": 0.92, "one_new|d4": 0.99, "one_new|d64": 0.93, "one_new|d8": 0.96, "random|d128": 0.9, "random|d16": 0.92, "random|d2": 0.99, "random|d3": 0.96, "random|d32": 0.91, "random|d4": 0.99, "random|d64": 0.87, "random|d8": 0.99}, "2d": {"all_seen|d128": 0.82, "all_seen|d16": 0.95, "all_seen|d2": 0.98, "all_seen|d3": 0.98, "all_seen|d32": 0.88, "all_seen|d4": 0.99, "all_seen|d64": 0.88, "all_seen|d8": 0.96, "one_new|d128": 0.78, "one_new|d16": 0.95, "one_new|d2": 0.99, "one_new|d3": 1.0, "one_new|d32": 0.92, "one_new|d4": 0.99, "one_new|d64": 0.93, "one_new|d8": 0.95, "random|d128": 0.9, "random|d16": 0.92, "random|d2": 0.99, "random|d3": 0.96, "random|d32": 0.9, "random|d4": 0.99, "random|d64": 0.87, "random|d8": 0.99}}, "counters": {"atomic": 3200000, "w2_sources": 1600000, "w2_queries": 3200000, "d2": 6400000, "row_loops": 51200000, "R_hist": {}}, "eval_seconds": 29.232378244400024, "peak_mem_gb": 15.119765281677246}
14
+ {"update": 110000, "elapsed": 13704.134062290192, "shallow": {"atomic": {"correct": 0.9996}, "val_d2": {"correct": 0.996}, "train_d2_fit": {"correct": 0.99875}}, "checkpoint": {"file": "ckpt_0110000.pt", "update": 110000, "sha256": "1791c9e8a4be8f84d2fdd713b598439dd7ca3332ae0b9efdc20c53dbd86c07e6", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.95}, "all_seen|d16": {"n": 100, "correct": 0.96}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 0.99}, "all_seen|d4": {"n": 100, "correct": 0.99}, "all_seen|d64": {"n": 100, "correct": 0.95}, "all_seen|d8": {"n": 100, "correct": 0.97}, "one_new|d128": {"n": 100, "correct": 0.97}, "one_new|d16": {"n": 100, "correct": 0.99}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 0.99}, "one_new|d32": {"n": 100, "correct": 0.98}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 0.95}, "one_new|d8": {"n": 100, "correct": 0.98}, "random|d128": {"n": 100, "correct": 0.97}, "random|d16": {"n": 100, "correct": 0.98}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 0.98}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 0.98}, "random|d8": {"n": 100, "correct": 0.99}}, "monitor_budgets": {"d": {"all_seen|d128": 0.95, "all_seen|d16": 0.96, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.99, "all_seen|d4": 0.99, "all_seen|d64": 0.95, "all_seen|d8": 0.97, "one_new|d128": 0.97, "one_new|d16": 0.99, "one_new|d2": 1.0, "one_new|d3": 0.99, "one_new|d32": 0.98, "one_new|d4": 1.0, "one_new|d64": 0.95, "one_new|d8": 0.98, "random|d128": 0.97, "random|d16": 0.98, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.98, "random|d4": 1.0, "random|d64": 0.98, "random|d8": 0.99}, "d+1": {"all_seen|d128": 0.95, "all_seen|d16": 0.96, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.99, "all_seen|d4": 0.99, "all_seen|d64": 0.95, "all_seen|d8": 0.97, "one_new|d128": 0.97, "one_new|d16": 0.99, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.98, "one_new|d4": 1.0, "one_new|d64": 0.95, "one_new|d8": 0.98, "random|d128": 0.97, "random|d16": 0.98, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.98, "random|d4": 1.0, "random|d64": 0.98, "random|d8": 0.99}, "d+2": {"all_seen|d128": 0.95, "all_seen|d16": 0.96, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.99, "all_seen|d4": 0.99, "all_seen|d64": 0.95, "all_seen|d8": 0.97, "one_new|d128": 0.97, "one_new|d16": 0.99, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.98, "one_new|d4": 1.0, "one_new|d64": 0.95, "one_new|d8": 0.98, "random|d128": 0.97, "random|d16": 0.98, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.98, "random|d4": 0.99, "random|d64": 0.98, "random|d8": 0.99}, "2d": {"all_seen|d128": 0.95, "all_seen|d16": 0.96, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.99, "all_seen|d4": 0.99, "all_seen|d64": 0.95, "all_seen|d8": 0.97, "one_new|d128": 0.97, "one_new|d16": 0.99, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.98, "one_new|d4": 1.0, "one_new|d64": 0.95, "one_new|d8": 0.98, "random|d128": 0.95, "random|d16": 0.97, "random|d2": 1.0, "random|d3": 0.99, "random|d32": 0.98, "random|d4": 0.98, "random|d64": 0.98, "random|d8": 0.99}}, "counters": {"atomic": 3520000, "w2_sources": 1760000, "w2_queries": 3520000, "d2": 7040000, "row_loops": 56320000, "R_hist": {}}, "eval_seconds": 29.351109266281128, "peak_mem_gb": 15.119765281677246}
15
+ {"update": 120000, "elapsed": 14942.493373155594, "shallow": {"atomic": {"correct": 0.9998}, "val_d2": {"correct": 0.997}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0120000.pt", "update": 120000, "sha256": "c6767cfb80c10f38ca4bac059d5240211929dbe5a5cadb4b48ebb07a196566bf", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.78}, "all_seen|d16": {"n": 100, "correct": 0.98}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 0.93}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 0.92}, "all_seen|d8": {"n": 100, "correct": 0.98}, "one_new|d128": {"n": 100, "correct": 0.79}, "one_new|d16": {"n": 100, "correct": 0.97}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 0.95}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 0.88}, "one_new|d8": {"n": 100, "correct": 0.99}, "random|d128": {"n": 100, "correct": 0.86}, "random|d16": {"n": 100, "correct": 0.96}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 0.98}, "random|d4": {"n": 100, "correct": 0.98}, "random|d64": {"n": 100, "correct": 0.94}, "random|d8": {"n": 100, "correct": 0.99}}, "monitor_budgets": {"d": {"all_seen|d128": 0.78, "all_seen|d16": 0.98, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.93, "all_seen|d4": 1.0, "all_seen|d64": 0.92, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.97, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.95, "one_new|d4": 1.0, "one_new|d64": 0.88, "one_new|d8": 0.99, "random|d128": 0.86, "random|d16": 0.96, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.98, "random|d4": 0.98, "random|d64": 0.94, "random|d8": 0.99}, "d+1": {"all_seen|d128": 0.78, "all_seen|d16": 0.98, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.92, "all_seen|d4": 1.0, "all_seen|d64": 0.92, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.97, "one_new|d2": 1.0, "one_new|d3": 0.99, "one_new|d32": 0.95, "one_new|d4": 1.0, "one_new|d64": 0.88, "one_new|d8": 0.99, "random|d128": 0.86, "random|d16": 0.96, "random|d2": 0.99, "random|d3": 1.0, "random|d32": 0.98, "random|d4": 0.98, "random|d64": 0.94, "random|d8": 0.98}, "d+2": {"all_seen|d128": 0.78, "all_seen|d16": 0.98, "all_seen|d2": 0.98, "all_seen|d3": 1.0, "all_seen|d32": 0.92, "all_seen|d4": 1.0, "all_seen|d64": 0.92, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.97, "one_new|d2": 1.0, "one_new|d3": 0.99, "one_new|d32": 0.95, "one_new|d4": 1.0, "one_new|d64": 0.88, "one_new|d8": 0.99, "random|d128": 0.86, "random|d16": 0.96, "random|d2": 0.99, "random|d3": 0.99, "random|d32": 0.98, "random|d4": 0.98, "random|d64": 0.94, "random|d8": 0.98}, "2d": {"all_seen|d128": 0.78, "all_seen|d16": 0.98, "all_seen|d2": 0.98, "all_seen|d3": 1.0, "all_seen|d32": 0.92, "all_seen|d4": 0.99, "all_seen|d64": 0.92, "all_seen|d8": 0.98, "one_new|d128": 0.79, "one_new|d16": 0.97, "one_new|d2": 1.0, "one_new|d3": 0.99, "one_new|d32": 0.95, "one_new|d4": 0.98, "one_new|d64": 0.88, "one_new|d8": 0.99, "random|d128": 0.86, "random|d16": 0.96, "random|d2": 0.99, "random|d3": 0.99, "random|d32": 0.98, "random|d4": 0.97, "random|d64": 0.94, "random|d8": 0.96}}, "counters": {"atomic": 3840000, "w2_sources": 1920000, "w2_queries": 3840000, "d2": 7680000, "row_loops": 61440000, "R_hist": {}}, "eval_seconds": 29.486525535583496, "peak_mem_gb": 15.119765281677246}
16
+ {"update": 130000, "elapsed": 16183.014260292053, "shallow": {"atomic": {"correct": 0.9998}, "val_d2": {"correct": 0.999}, "train_d2_fit": {"correct": 0.998625}}, "checkpoint": {"file": "ckpt_0130000.pt", "update": 130000, "sha256": "1bf915b7a6cb858a51b534b9afd289421e0fdbb3c6f534831f0c4254143fb890", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 0.85}, "all_seen|d16": {"n": 100, "correct": 0.99}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 0.98}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 0.98}, "all_seen|d8": {"n": 100, "correct": 0.98}, "one_new|d128": {"n": 100, "correct": 0.88}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 0.96}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 0.96}, "one_new|d8": {"n": 100, "correct": 0.98}, "random|d128": {"n": 100, "correct": 0.85}, "random|d16": {"n": 100, "correct": 0.99}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 0.99}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 0.98}, "random|d8": {"n": 100, "correct": 0.99}}, "monitor_budgets": {"d": {"all_seen|d128": 0.85, "all_seen|d16": 0.99, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.98, "all_seen|d4": 1.0, "all_seen|d64": 0.98, "all_seen|d8": 0.98, "one_new|d128": 0.88, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.96, "one_new|d4": 1.0, "one_new|d64": 0.96, "one_new|d8": 0.98, "random|d128": 0.85, "random|d16": 0.99, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.99, "random|d4": 1.0, "random|d64": 0.98, "random|d8": 0.99}, "d+1": {"all_seen|d128": 0.85, "all_seen|d16": 0.98, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.98, "all_seen|d4": 1.0, "all_seen|d64": 0.97, "all_seen|d8": 0.98, "one_new|d128": 0.87, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.96, "one_new|d4": 1.0, "one_new|d64": 0.96, "one_new|d8": 0.98, "random|d128": 0.85, "random|d16": 0.99, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.99, "random|d4": 1.0, "random|d64": 0.98, "random|d8": 0.98}, "d+2": {"all_seen|d128": 0.85, "all_seen|d16": 0.98, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.98, "all_seen|d4": 1.0, "all_seen|d64": 0.97, "all_seen|d8": 0.98, "one_new|d128": 0.87, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.96, "one_new|d4": 1.0, "one_new|d64": 0.96, "one_new|d8": 0.98, "random|d128": 0.85, "random|d16": 0.99, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.99, "random|d4": 1.0, "random|d64": 0.98, "random|d8": 0.98}, "2d": {"all_seen|d128": 0.85, "all_seen|d16": 0.97, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 0.98, "all_seen|d4": 1.0, "all_seen|d64": 0.96, "all_seen|d8": 0.98, "one_new|d128": 0.86, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 0.95, "one_new|d4": 1.0, "one_new|d64": 0.96, "one_new|d8": 0.98, "random|d128": 0.85, "random|d16": 0.99, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 0.99, "random|d4": 1.0, "random|d64": 0.97, "random|d8": 0.98}}, "counters": {"atomic": 4160000, "w2_sources": 2080000, "w2_queries": 4160000, "d2": 8320000, "row_loops": 66560000, "R_hist": {}}, "eval_seconds": 29.17534899711609, "peak_mem_gb": 15.119765281677246}
17
+ {"update": 140000, "elapsed": 17422.63646006584, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0140000.pt", "update": 140000, "sha256": "a00e672829198b3cba32d4241aa914b067cbd82f4f46f1a1219075f23ce22f8a", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 0.99}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 0.99, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 0.99, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 0.99, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 0.99, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 4480000, "w2_sources": 2240000, "w2_queries": 4480000, "d2": 8960000, "row_loops": 71680000, "R_hist": {}}, "eval_seconds": 29.278916358947754, "peak_mem_gb": 15.119765281677246}
18
+ {"update": 150000, "elapsed": 18661.133254766464, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0150000.pt", "update": 150000, "sha256": "c7b347a167d12485afddc0943dc92ce8326e3eb5a099ce5e4aafb4c2c0b81d8c", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 1.0}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 4800000, "w2_sources": 2400000, "w2_queries": 4800000, "d2": 9600000, "row_loops": 76800000, "R_hist": {}}, "eval_seconds": 29.17150855064392, "peak_mem_gb": 15.119765281677246}
19
+ {"update": 160000, "elapsed": 19899.709609031677, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0160000.pt", "update": 160000, "sha256": "bbd8ee69c630dd68e2ae25cb2eacb4bd1d3ffcadadd37a7bbc9938cc791ecb3a", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 1.0}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 5120000, "w2_sources": 2560000, "w2_queries": 5120000, "d2": 10240000, "row_loops": 81920000, "R_hist": {}}, "eval_seconds": 29.216355323791504, "peak_mem_gb": 15.119765281677246}
20
+ {"update": 170000, "elapsed": 21140.495632886887, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0170000.pt", "update": 170000, "sha256": "76a4c1cd41d687ef06a1ccf6b67b2c56399c1ea2bf1103c51c144b3eddb6e20f", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 1.0}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 5440000, "w2_sources": 2720000, "w2_queries": 5440000, "d2": 10880000, "row_loops": 87040000, "R_hist": {}}, "eval_seconds": 29.20540428161621, "peak_mem_gb": 15.119765281677246}
21
+ {"update": 180000, "elapsed": 22380.195405244827, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0180000.pt", "update": 180000, "sha256": "d767ae2ced030111845d93e7f35587d7ecfc7d8ce43622d47df0d9c93638b191", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 1.0}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 5760000, "w2_sources": 2880000, "w2_queries": 5760000, "d2": 11520000, "row_loops": 92160000, "R_hist": {}}, "eval_seconds": 29.201653480529785, "peak_mem_gb": 15.119765281677246}
22
+ {"update": 190000, "elapsed": 23618.75986123085, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0190000.pt", "update": 190000, "sha256": "a44ac51a1decb6919d48d1d02fcef7f4e525759760b913fffa12e064754d2c29", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 1.0}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 6080000, "w2_sources": 3040000, "w2_queries": 6080000, "d2": 12160000, "row_loops": 97280000, "R_hist": {}}, "eval_seconds": 29.26492953300476, "peak_mem_gb": 15.119765281677246}
23
+ {"update": 200000, "elapsed": 24857.36122918129, "shallow": {"atomic": {"correct": 1.0}, "val_d2": {"correct": 1.0}, "train_d2_fit": {"correct": 1.0}}, "checkpoint": {"file": "ckpt_0200000.pt", "update": 200000, "sha256": "6be182b9f912f741dfb0ba491879812e0813cc8adb264c76fa0cf5a9bbc64e22", "bytes": 115042564, "rope_base": 10000.0, "max_pos": 4097, "kind": "model_weights_for_exact_evaluation"}, "monitor": {"all_seen|d128": {"n": 100, "correct": 1.0}, "all_seen|d16": {"n": 100, "correct": 1.0}, "all_seen|d2": {"n": 100, "correct": 1.0}, "all_seen|d3": {"n": 100, "correct": 1.0}, "all_seen|d32": {"n": 100, "correct": 1.0}, "all_seen|d4": {"n": 100, "correct": 1.0}, "all_seen|d64": {"n": 100, "correct": 1.0}, "all_seen|d8": {"n": 100, "correct": 1.0}, "one_new|d128": {"n": 100, "correct": 1.0}, "one_new|d16": {"n": 100, "correct": 1.0}, "one_new|d2": {"n": 100, "correct": 1.0}, "one_new|d3": {"n": 100, "correct": 1.0}, "one_new|d32": {"n": 100, "correct": 1.0}, "one_new|d4": {"n": 100, "correct": 1.0}, "one_new|d64": {"n": 100, "correct": 1.0}, "one_new|d8": {"n": 100, "correct": 1.0}, "random|d128": {"n": 100, "correct": 1.0}, "random|d16": {"n": 100, "correct": 1.0}, "random|d2": {"n": 100, "correct": 1.0}, "random|d3": {"n": 100, "correct": 1.0}, "random|d32": {"n": 100, "correct": 1.0}, "random|d4": {"n": 100, "correct": 1.0}, "random|d64": {"n": 100, "correct": 1.0}, "random|d8": {"n": 100, "correct": 1.0}}, "monitor_budgets": {"d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+1": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "d+2": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}, "2d": {"all_seen|d128": 1.0, "all_seen|d16": 1.0, "all_seen|d2": 1.0, "all_seen|d3": 1.0, "all_seen|d32": 1.0, "all_seen|d4": 1.0, "all_seen|d64": 1.0, "all_seen|d8": 1.0, "one_new|d128": 1.0, "one_new|d16": 1.0, "one_new|d2": 1.0, "one_new|d3": 1.0, "one_new|d32": 1.0, "one_new|d4": 1.0, "one_new|d64": 1.0, "one_new|d8": 1.0, "random|d128": 1.0, "random|d16": 1.0, "random|d2": 1.0, "random|d3": 1.0, "random|d32": 1.0, "random|d4": 1.0, "random|d64": 1.0, "random|d8": 1.0}}, "counters": {"atomic": 6400000, "w2_sources": 3200000, "w2_queries": 6400000, "d2": 12800000, "row_loops": 102400000, "R_hist": {}}, "eval_seconds": 29.17498278617859, "peak_mem_gb": 15.119765281677246}
for Minegishi/wait_denoising_fixed_t0/train_command.json ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ "python",
3
+ "-u",
4
+ "reproducibility/train_halt.py",
5
+ "--data_dir",
6
+ "reproducibility/data/chain_loop_k8",
7
+ "--atomic",
8
+ "reproducibility/data/atomic_joint_2026-09-19/train_atomic.json",
9
+ "--save_dir",
10
+ "new_runs/R10k_T0iso_s7",
11
+ "--arm",
12
+ "D",
13
+ "--seed",
14
+ "7",
15
+ "--updates",
16
+ "200000",
17
+ "--lr_decay",
18
+ "linear",
19
+ "--pos",
20
+ "rope",
21
+ "--rope_base",
22
+ "10000",
23
+ "--max_pos",
24
+ "4097",
25
+ "--anchor",
26
+ "1",
27
+ "--anchor_mode",
28
+ "full",
29
+ "--consist_w",
30
+ "1",
31
+ "--cs_mode",
32
+ "t0_wait_denoise",
33
+ "--consist_depth",
34
+ "32",
35
+ "--consist_rows",
36
+ "32",
37
+ "--consist_start",
38
+ "20000",
39
+ "--offset_max",
40
+ "128",
41
+ "--offset_fill",
42
+ "entity",
43
+ "--cs_noise_scale",
44
+ "0.2",
45
+ "--cs_noise_kind",
46
+ "isotropic",
47
+ "--cs_noise_scope",
48
+ "wait",
49
+ "--train_precision",
50
+ "tf32",
51
+ "--cuda_graph",
52
+ "1",
53
+ "--monitor_per_cell",
54
+ "100",
55
+ "--eval_every",
56
+ "10000",
57
+ "--ckpt_every",
58
+ "5000",
59
+ "--final_n",
60
+ "5000",
61
+ "--diagnostic_n",
62
+ "100"
63
+ ]
for Minegishi/wait_denoising_fixed_t0/train_log.jsonl ADDED
The diff for this file is too large to render. See raw diff