For Minegishi: training documentation, reproducibility files and line plots 4/4
Browse files- for Minegishi/SHA256SUMS +104 -0
- for Minegishi/reproducibility/data/chain_loop_k8/train_atomic.json +0 -0
- for Minegishi/reproducibility/data/chain_loop_k8/train_d2.json +0 -0
- for Minegishi/reproducibility/data/chain_loop_k8/train_d3.json +1 -0
- for Minegishi/reproducibility/data/chain_loop_k8/train_w2.json +0 -0
- for Minegishi/reproducibility/data/chain_loop_k8/val_d2.json +1 -0
- for Minegishi/reproducibility/data/chain_loop_k8/val_w2.json +1 -0
- for Minegishi/reproducibility/data/chain_loop_k8/vocab.json +1 -0
- for Minegishi/reproducibility/eval_checkpoints.py +26 -0
- for Minegishi/reproducibility/extrapolation.py +289 -0
- for Minegishi/reproducibility/frozen_source_sha256.json +14 -0
- for Minegishi/reproducibility/load_checkpoint.py +35 -0
- for Minegishi/reproducibility/model/__init__.py +5 -0
- for Minegishi/reproducibility/model/gpt2.py +258 -0
- for Minegishi/reproducibility/model/latent_executor.py +213 -0
- for Minegishi/reproducibility/model/loop_gpt.py +216 -0
- for Minegishi/reproducibility/objectives.py +467 -0
- for Minegishi/reproducibility/probes.py +371 -0
- for Minegishi/reproducibility/recipe.py +20 -0
- for Minegishi/reproducibility/streams.py +27 -0
- for Minegishi/reproducibility/train_chain.py +435 -0
- for Minegishi/reproducibility/train_halt.py +996 -0
- for Minegishi/wait_denoising_fixed_t0/README.md +7 -0
- for Minegishi/wait_denoising_fixed_t0/checkpoints.json +237 -0
- for Minegishi/wait_denoising_fixed_t0/config.json +19 -0
- for Minegishi/wait_denoising_fixed_t0/final_eval.json +513 -0
- for Minegishi/wait_denoising_fixed_t0/manifest.json +188 -0
- for Minegishi/wait_denoising_fixed_t0/metrics.jsonl +23 -0
- for Minegishi/wait_denoising_fixed_t0/train_command.json +63 -0
- 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
|
|
|