AbstractPhil commited on
Commit
00cc430
·
verified ·
1 Parent(s): 41865a8

Anima close-out campaign notebook: exp004b (relay retrain, fp32 masters, 2 declared seeds, real held-out eval) + exp017 (first multiband3 on a DiT), Colab/Blackwell-96GB, ship-on-completion per arm

Browse files
Files changed (1) hide show
  1. colab/anima_closeout.ipynb +924 -0
colab/anima_closeout.ipynb ADDED
@@ -0,0 +1,924 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "nbformat": 4,
3
+ "nbformat_minor": 5,
4
+ "metadata": {
5
+ "colab": {
6
+ "provenance": []
7
+ },
8
+ "kernelspec": {
9
+ "name": "python3",
10
+ "display_name": "Python 3"
11
+ },
12
+ "language_info": {
13
+ "name": "python"
14
+ },
15
+ "accelerator": "GPU"
16
+ },
17
+ "cells": [
18
+ {
19
+ "cell_type": "markdown",
20
+ "metadata": {},
21
+ "source": [
22
+ "# AMoE Anima close-out \u2014 exp004b + exp017 (Colab / Blackwell 96GB)\n",
23
+ "\n",
24
+ "Closes the r2 campaign's Anima gap:\n",
25
+ "\n",
26
+ "| | what | why |\n",
27
+ "|---|---|---|\n",
28
+ "| **exp004b** | relay retrain: 2 declared seeds, 10k steps, full dataset, held-out eval, LoRA-r16 control | exp004 was one ~40-min run, no results.json, still improving at cutoff |\n",
29
+ "| **exp017** | multiband3 on the Anima DiT \u2014 first live run of the fork's mode | the band mechanism has never run on a DiT |\n",
30
+ "| **the fix** | `bf16_master_weights = true` + declared `seed` | exp004's 28 gates were **quantization-frozen**: plain Adam stepped bf16 leaves directly; ~4.5e-4 updates vs 0.0078 half-ULP at \u22123.0 rounded to no-ops (proven from the salvaged Adam state) |\n",
31
+ "\n",
32
+ "Run cells top to bottom. Everything is idempotent: after a disconnect,\n",
33
+ "re-run **CELL 0**, then re-run the interrupted run cell (checkpoints every\n",
34
+ "30 min; pass `resume=True` to `run_arm`).\n",
35
+ "\n",
36
+ "**Needs:** Colab secret `HF_TOKEN` (notebook access ON) with write access\n",
37
+ "to `AbstractPhil/geolip-aleph-diffusion`.\n",
38
+ "\n",
39
+ "**License:** Anima weights are NON-COMMERCIAL (CircleStone NC + NVIDIA\n",
40
+ "Open Model License); every checkpoint produced here inherits NC.\n"
41
+ ]
42
+ },
43
+ {
44
+ "cell_type": "code",
45
+ "metadata": {},
46
+ "execution_count": null,
47
+ "outputs": [],
48
+ "source": [
49
+ "# \u2550\u2550\u2550 CELL 0 \u2014 ACTIVATION (idempotent; re-run after any disconnect) \u2550\u2550\u2550\u2550\u2550\u2550\n",
50
+ "import json, os, subprocess, sys, time\n",
51
+ "from pathlib import Path\n",
52
+ "\n",
53
+ "if \"ANIMA\" not in globals(): # paste-ahead guard\n",
54
+ " ANIMA = {}\n",
55
+ "\n",
56
+ "WORK = Path(\"/content/anima_campaign\")\n",
57
+ "FORK = Path(\"/content/diffusion-pipe\")\n",
58
+ "DATA, RUNS, TOMLS = WORK / \"data\", WORK / \"runs\", WORK / \"tomls\"\n",
59
+ "for _d in (WORK, DATA, RUNS, TOMLS):\n",
60
+ " _d.mkdir(parents=True, exist_ok=True)\n",
61
+ "\n",
62
+ "def sh(cmd, check=True):\n",
63
+ " print(f\"$ {cmd}\")\n",
64
+ " rc = subprocess.run(cmd, shell=True, text=True).returncode\n",
65
+ " if check and rc != 0:\n",
66
+ " raise RuntimeError(f\"rc={rc}: {cmd}\")\n",
67
+ " return rc\n",
68
+ "\n",
69
+ "# -- 1. GPU + Blackwell riders --------------------------------------------\n",
70
+ "import torch\n",
71
+ "assert torch.cuda.is_available(), \"No GPU \u2014 Runtime > Change runtime type.\"\n",
72
+ "_cap = torch.cuda.get_device_capability(0)\n",
73
+ "print(f\"GPU: {torch.cuda.get_device_name(0)} | sm_{_cap[0]}{_cap[1]} | \"\n",
74
+ " f\"{torch.cuda.get_device_properties(0).total_memory / 2**30:.0f} GB | \"\n",
75
+ " f\"torch {torch.__version__} cu{torch.version.cuda}\")\n",
76
+ "if _cap >= (12, 0):\n",
77
+ " _tv = tuple(int(x) for x in torch.__version__.split(\"+\")[0].split(\".\")[:2])\n",
78
+ " _cv = tuple(int(x) for x in (torch.version.cuda or \"0.0\").split(\".\")[:2])\n",
79
+ " if not (_tv >= (2, 7) and _cv >= (12, 8)):\n",
80
+ " raise RuntimeError(\n",
81
+ " f\"sm_120 needs torch>=2.7 + cu>=12.8; runtime has \"\n",
82
+ " f\"{torch.__version__}/cu{torch.version.cuda}. Do NOT blind-pip \"\n",
83
+ " \"torch \u2014 switch runtimes, or deliberately install the matching \"\n",
84
+ " \"cu128 build from pytorch.org.\")\n",
85
+ " (torch.randn(256, 256, device=\"cuda\")\n",
86
+ " @ torch.randn(256, 256, device=\"cuda\")).sum().item()\n",
87
+ " print(\"sm_120 matmul: OK\")\n",
88
+ "try:\n",
89
+ " import flash_attn # noqa: F401\n",
90
+ " raise RuntimeError(\"flash-attn is installed and BROKEN on sm_120 \u2014 \"\n",
91
+ " \"`pip uninstall flash-attn`; SDPA is the path.\")\n",
92
+ "except ImportError:\n",
93
+ " pass\n",
94
+ "\n",
95
+ "# -- 2. HF token (never printed) ------------------------------------------\n",
96
+ "_tok = os.environ.get(\"HF_TOKEN\")\n",
97
+ "if not _tok:\n",
98
+ " try:\n",
99
+ " from google.colab import userdata\n",
100
+ " _tok = userdata.get(\"HF_TOKEN\")\n",
101
+ " except Exception:\n",
102
+ " _tok = None\n",
103
+ "if not _tok:\n",
104
+ " raise RuntimeError(\"Add HF_TOKEN in Colab secrets (key icon) with \"\n",
105
+ " \"notebook access ON, then re-run this cell.\")\n",
106
+ "os.environ[\"HF_TOKEN\"] = _tok\n",
107
+ "print(\"HF token: present (not printed)\")\n",
108
+ "\n",
109
+ "# -- 3. fork + deps -------------------------------------------------------\n",
110
+ "if not (FORK / \"train.py\").exists():\n",
111
+ " sh(f\"git clone -q https://github.com/AbstractEyes/diffusion-pipe {FORK}\")\n",
112
+ "sh(f\"git -C {FORK} fetch -q origin\")\n",
113
+ "sh(f\"git -C {FORK} checkout -q origin/main\")\n",
114
+ "sh(f\"git -C {FORK} submodule update --init submodules/ComfyUI\")\n",
115
+ "print(\"fork @\", subprocess.run(\n",
116
+ " [\"git\", \"-C\", str(FORK), \"log\", \"--oneline\", \"-1\"],\n",
117
+ " capture_output=True, text=True).stdout.strip())\n",
118
+ "_src = (FORK / \"train.py\").read_text()\n",
119
+ "assert \"bf16_master_weights\" in _src and \"run_seed\" in _src, \\\n",
120
+ " \"fork main lacks the campaign patches (bf16_master_weights + seed)\"\n",
121
+ "\n",
122
+ "if not ANIMA.get(\"DEPS_DONE\"):\n",
123
+ " sh(f\"pip install -q -r {FORK}/requirements.txt\")\n",
124
+ " sh('pip install -q \"amoe-lora[diffusion] @ '\n",
125
+ " 'git+https://github.com/AbstractEyes/amoe-lora\"')\n",
126
+ " ANIMA[\"DEPS_DONE\"] = True\n",
127
+ "import amoe, deepspeed\n",
128
+ "print(f\"amoe {amoe.__version__} | deepspeed {deepspeed.__version__}\")\n",
129
+ "\n",
130
+ "# -- 4. Anima weights (NC) ------------------------------------------------\n",
131
+ "from huggingface_hub import hf_hub_download\n",
132
+ "MODELS = WORK / \"models\" / \"anima\"\n",
133
+ "for _f in (\"diffusion_models/anima-base-v1.0.safetensors\",\n",
134
+ " \"text_encoders/qwen_3_06b_base.safetensors\",\n",
135
+ " \"vae/qwen_image_vae.safetensors\"):\n",
136
+ " hf_hub_download(\"circlestone-labs/Anima\", f\"split_files/{_f}\",\n",
137
+ " local_dir=str(MODELS))\n",
138
+ "print(\"Anima split_files present \u2014 NC weights; derived ckpts inherit NC\")\n",
139
+ "\n",
140
+ "# -- 5. dataset columns, verified BEFORE anything spends ------------------\n",
141
+ "DATASET = \"AbstractPhil/qwen-deepfashion-fused\"\n",
142
+ "from huggingface_hub import HfApi\n",
143
+ "SHARDS = sorted(f for f in HfApi().list_repo_files(\n",
144
+ " DATASET, repo_type=\"dataset\") if f.endswith(\".parquet\"))\n",
145
+ "assert SHARDS, f\"no parquet shards in {DATASET}\"\n",
146
+ "import pyarrow.parquet as pq\n",
147
+ "_p0 = hf_hub_download(DATASET, SHARDS[0], repo_type=\"dataset\")\n",
148
+ "_cols = set(pq.read_schema(_p0).names)\n",
149
+ "_missing = {\"image\", \"caption_vlm_json\"} - _cols\n",
150
+ "assert not _missing, f\"dataset lacks {_missing}; has {sorted(_cols)[:24]}\"\n",
151
+ "ANIMA.update(DATASET=DATASET, SHARDS=SHARDS,\n",
152
+ " HAS_WH={\"image_width\", \"image_height\"} <= _cols,\n",
153
+ " ID_COL=\"id\" if \"id\" in _cols else None)\n",
154
+ "print(f\"dataset OK: {len(SHARDS)} shard(s); wh cols {ANIMA['HAS_WH']}; \"\n",
155
+ " f\"id col {ANIMA['ID_COL']}\")\n",
156
+ "print(\"\\n=== READY ===\")\n"
157
+ ]
158
+ },
159
+ {
160
+ "cell_type": "code",
161
+ "metadata": {},
162
+ "execution_count": null,
163
+ "outputs": [],
164
+ "source": [
165
+ "# \u2550\u2550\u2550 CELL 1 \u2014 materialize the held-out split + write all TOMLs \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
166
+ "# Deterministic row-hash split: 512 eval rows, disjoint by construction,\n",
167
+ "# recorded in split_manifest.json (the reproducibility anchor). Both sides\n",
168
+ "# read the SAME materialized parquet. caption_type='text' = exp004's\n",
169
+ "# verbatim-JSON recipe (flattened-'json' is the recorded alternative).\n",
170
+ "import hashlib, json\n",
171
+ "import pyarrow as pa\n",
172
+ "import pyarrow.parquet as pq\n",
173
+ "from huggingface_hub import hf_hub_download\n",
174
+ "\n",
175
+ "EVAL_N = 512\n",
176
+ "TRAIN_DIR = DATA / \"train\"; TRAIN_DIR.mkdir(exist_ok=True)\n",
177
+ "EVAL_PARQUET = DATA / \"eval.parquet\"\n",
178
+ "MANIFEST = DATA / \"split_manifest.json\"\n",
179
+ "\n",
180
+ "if not MANIFEST.exists():\n",
181
+ " _paths = [hf_hub_download(ANIMA[\"DATASET\"], s, repo_type=\"dataset\")\n",
182
+ " for s in ANIMA[\"SHARDS\"]]\n",
183
+ "\n",
184
+ " def _row_key(si, ri, batch, j):\n",
185
+ " if ANIMA[\"ID_COL\"]:\n",
186
+ " return str(batch.column(ANIMA[\"ID_COL\"])[j].as_py())\n",
187
+ " cap = str(batch.column(\"caption_vlm_json\")[j].as_py())[:256]\n",
188
+ " return f\"{si}:{ri}:{hashlib.sha1(cap.encode()).hexdigest()[:12]}\"\n",
189
+ "\n",
190
+ " _key_cols = [c for c in (ANIMA[\"ID_COL\"], \"caption_vlm_json\") if c]\n",
191
+ " ranked = []\n",
192
+ " for si, p in enumerate(_paths):\n",
193
+ " ri = 0\n",
194
+ " for batch in pq.ParquetFile(p).iter_batches(batch_size=256,\n",
195
+ " columns=_key_cols):\n",
196
+ " for j in range(batch.num_rows):\n",
197
+ " k = _row_key(si, ri, batch, j)\n",
198
+ " ranked.append(\n",
199
+ " (hashlib.sha256(k.encode()).hexdigest(), si, ri, k))\n",
200
+ " ri += 1\n",
201
+ " ranked.sort()\n",
202
+ " eval_set = {(si, ri) for _h, si, ri, _k in ranked[:EVAL_N]}\n",
203
+ " total = len(ranked)\n",
204
+ "\n",
205
+ " ew = tw = None\n",
206
+ " for si, p in enumerate(_paths):\n",
207
+ " ri = 0\n",
208
+ " for batch in pq.ParquetFile(p).iter_batches(batch_size=64):\n",
209
+ " tbl = pa.Table.from_batches([batch])\n",
210
+ " emask = [(si, ri + j) in eval_set for j in range(batch.num_rows)]\n",
211
+ " ri += batch.num_rows\n",
212
+ " esub = tbl.filter(pa.array(emask))\n",
213
+ " tsub = tbl.filter(pa.array([not m for m in emask]))\n",
214
+ " if esub.num_rows:\n",
215
+ " if ew is None:\n",
216
+ " ew = pq.ParquetWriter(EVAL_PARQUET, esub.schema)\n",
217
+ " ew.write_table(esub)\n",
218
+ " if tsub.num_rows:\n",
219
+ " if tw is None:\n",
220
+ " tw = pq.ParquetWriter(TRAIN_DIR / \"train.parquet\",\n",
221
+ " tsub.schema)\n",
222
+ " tw.write_table(tsub)\n",
223
+ " for w in (ew, tw):\n",
224
+ " if w:\n",
225
+ " w.close()\n",
226
+ " MANIFEST.write_text(json.dumps(\n",
227
+ " {\"dataset\": ANIMA[\"DATASET\"], \"total_rows\": total, \"eval_n\": EVAL_N,\n",
228
+ " \"eval_ids\": [k for _h, _s, _r, k in ranked[:EVAL_N]],\n",
229
+ " \"rule\": \"sha256(row_key) ascending; first EVAL_N rows are eval\"},\n",
230
+ " indent=1))\n",
231
+ " print(f\"split: {total} rows -> {total - EVAL_N} train + {EVAL_N} eval\")\n",
232
+ "else:\n",
233
+ " print(\"split already materialized:\",\n",
234
+ " json.loads(MANIFEST.read_text())[\"total_rows\"], \"rows\")\n",
235
+ "\n",
236
+ "# -- dataset TOMLs (the exp004 recipe) ------------------------------------\n",
237
+ "_wh = (\"image_width_column = 'image_width'\\n\"\n",
238
+ " \"image_height_column = 'image_height'\\n\") if ANIMA[\"HAS_WH\"] else \"\"\n",
239
+ "\n",
240
+ "def _dataset_toml(files):\n",
241
+ " return f\"\"\"# generated by anima_closeout.ipynb\n",
242
+ "resolutions = [512]\n",
243
+ "enable_ar_bucket = true\n",
244
+ "min_ar = 0.5\n",
245
+ "max_ar = 2.0\n",
246
+ "num_ar_buckets = 7\n",
247
+ "cache_backend = 'parquet'\n",
248
+ "cache_shard_size_mb = 350\n",
249
+ "\n",
250
+ "[[directory]]\n",
251
+ "type = 'parquet'\n",
252
+ "parquet_files = '{files}'\n",
253
+ "image_column = 'image'\n",
254
+ "{_wh}caption_column = 'caption_vlm_json'\n",
255
+ "caption_type = 'text'\n",
256
+ "num_repeats = 1\n",
257
+ "skip_empty_caption = true\n",
258
+ "\"\"\"\n",
259
+ "\n",
260
+ "(TOMLS / \"dataset_train.toml\").write_text(\n",
261
+ " _dataset_toml(f\"{TRAIN_DIR}/*.parquet\"))\n",
262
+ "(TOMLS / \"dataset_eval.toml\").write_text(_dataset_toml(str(EVAL_PARQUET)))\n",
263
+ "\n",
264
+ "# -- training TOMLs per arm -----------------------------------------------\n",
265
+ "_AP = WORK / \"models\" / \"anima\" / \"split_files\"\n",
266
+ "\n",
267
+ "def _arm_toml(tag, seed, steps, mode=\"relay\", lora=False):\n",
268
+ " out = RUNS / tag\n",
269
+ " L = [f\"# {tag} \u2014 generated by anima_closeout.ipynb\",\n",
270
+ " f\"output_dir = '{out}'\",\n",
271
+ " f\"dataset = '{TOMLS}/dataset_train.toml'\",\n",
272
+ " f\"seed = {seed}\",\n",
273
+ " \"bf16_master_weights = true # the exp004 gate-freeze fix\",\n",
274
+ " \"epochs = 1000\",\n",
275
+ " f\"max_steps = {steps}\",\n",
276
+ " \"micro_batch_size_per_gpu = 8\",\n",
277
+ " \"pipeline_stages = 1\",\n",
278
+ " \"gradient_accumulation_steps = 1\",\n",
279
+ " \"gradient_clipping = 0\",\n",
280
+ " \"save_every_n_steps = 2500\",\n",
281
+ " \"checkpoint_every_n_minutes = 30\",\n",
282
+ " \"activation_checkpointing = false # 96GB rider\",\n",
283
+ " \"save_dtype = 'bfloat16'\",\n",
284
+ " \"caching_batch_size = 8\",\n",
285
+ " f\"eval_datasets = [{{name = 'heldout', config = '{TOMLS}/dataset_eval.toml'}}]\",\n",
286
+ " \"eval_every_n_steps = 1000\",\n",
287
+ " \"eval_before_first_step = true\",\n",
288
+ " \"eval_micro_batch_size_per_gpu = 8\",\n",
289
+ " \"eval_gradient_accumulation_steps = 1\",\n",
290
+ " \"\", \"[model]\", \"type = 'anima'\",\n",
291
+ " f\"transformer_path = '{_AP}/diffusion_models/anima-base-v1.0.safetensors'\",\n",
292
+ " f\"llm_path = '{_AP}/text_encoders/qwen_3_06b_base.safetensors'\",\n",
293
+ " f\"vae_path = '{_AP}/vae/qwen_image_vae.safetensors'\",\n",
294
+ " \"dtype = 'bfloat16'\",\n",
295
+ " \"timestep_sample_method = 'logit_normal'\"]\n",
296
+ " if lora:\n",
297
+ " L += [\"self_attn_lr = 1e-3\", \"cross_attn_lr = 1e-3\", \"mlp_lr = 1e-3\",\n",
298
+ " \"mod_lr = 1e-3\", \"llm_adapter_lr = 1e-3\",\n",
299
+ " \"\", \"[adapter]\", \"type = 'lora'\", \"rank = 16\",\n",
300
+ " \"dtype = 'bfloat16'\"]\n",
301
+ " else:\n",
302
+ " L += [\"aleph_relay = true\",\n",
303
+ " f\"aleph_relay_mode = '{mode}'\",\n",
304
+ " \"aleph_relay_rank = 16\",\n",
305
+ " \"aleph_relay_every = 1\",\n",
306
+ " \"aleph_relay_lr = 1e-3\",\n",
307
+ " \"self_attn_lr = 0\", \"cross_attn_lr = 0\", \"mlp_lr = 0\",\n",
308
+ " \"mod_lr = 0\",\n",
309
+ " \"llm_adapter_lr = 0 # NON-NEGOTIABLE (anima-trainer law)\"]\n",
310
+ " L += [\"\", \"[optimizer]\",\n",
311
+ " \"type = 'adam' # pure Adam \u2014 never adamw on aleph paths\",\n",
312
+ " (\"lr = 1e-3\" if lora else \"lr = 0\"),\n",
313
+ " \"betas = [0.9, 0.99]\",\n",
314
+ " \"weight_decay = 0.0\",\n",
315
+ " \"\", \"[monitoring]\", \"enable_wandb = false\"]\n",
316
+ " p = TOMLS / f\"{tag}.toml\"\n",
317
+ " p.write_text(\"\\n\".join(L) + \"\\n\")\n",
318
+ " return p\n",
319
+ "\n",
320
+ "ARMS = {\n",
321
+ " \"smoke\": _arm_toml(\"smoke\", 0, 200),\n",
322
+ " \"e004b_relay_s0\": _arm_toml(\"e004b_relay_s0\", 0, 10000),\n",
323
+ " \"e004b_relay_s1\": _arm_toml(\"e004b_relay_s1\", 1, 10000),\n",
324
+ " \"e004b_lora_s0\": _arm_toml(\"e004b_lora_s0\", 0, 10000, lora=True),\n",
325
+ " \"e004b_lora_s1\": _arm_toml(\"e004b_lora_s1\", 1, 10000, lora=True),\n",
326
+ " \"e017_mb3_smoke\": _arm_toml(\"e017_mb3_smoke\", 0, 200, mode=\"multiband3\"),\n",
327
+ " \"e017_mb3_s0\": _arm_toml(\"e017_mb3_s0\", 0, 10000, mode=\"multiband3\"),\n",
328
+ " \"e017_mb3_s1\": _arm_toml(\"e017_mb3_s1\", 1, 10000, mode=\"multiband3\"),\n",
329
+ "}\n",
330
+ "print(\"TOMLs written:\", \", \".join(ARMS))\n"
331
+ ]
332
+ },
333
+ {
334
+ "cell_type": "code",
335
+ "metadata": {},
336
+ "execution_count": null,
337
+ "outputs": [],
338
+ "source": [
339
+ "# \u2550\u2550\u2550 CELL 2 \u2014 helpers + SMOKE + the G-gate assert \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
340
+ "# The smoke = 200 real-recipe steps. Jobs: prove the fp32-master fix in\n",
341
+ "# flight (gates MUST move \u2014 the exact exp004 failure), prove eval wiring\n",
342
+ "# (eval_before_first_step fires at step 0), build the latent cache\n",
343
+ "# (~0.5h, reused by every later arm).\n",
344
+ "import glob, io, json, os, subprocess, time\n",
345
+ "from pathlib import Path\n",
346
+ "import torch\n",
347
+ "\n",
348
+ "HUB_REPO = \"AbstractPhil/geolip-aleph-diffusion\"\n",
349
+ "\n",
350
+ "def run_arm(tag, resume=False):\n",
351
+ " log = RUNS / f\"{tag}.log\"\n",
352
+ " flags = \"--resume_from_checkpoint\" if resume else \"\"\n",
353
+ " t0 = time.time()\n",
354
+ " rc = sh(f\"cd {FORK} && NCCL_P2P_DISABLE=1 NCCL_IB_DISABLE=1 \"\n",
355
+ " f\"deepspeed --num_gpus=1 train.py --deepspeed \"\n",
356
+ " f\"--config {ARMS[tag]} {flags} 2>&1 | tee -a {log}\",\n",
357
+ " check=False)\n",
358
+ " print(f\"[{tag}] rc={rc} wall={(time.time() - t0) / 60:.1f} min\")\n",
359
+ " if rc != 0:\n",
360
+ " raise RuntimeError(f\"{tag} failed \u2014 read {log}\")\n",
361
+ " return newest_save(tag)\n",
362
+ "\n",
363
+ "def latest_run_dir(tag):\n",
364
+ " run_dirs = sorted(glob.glob(str(RUNS / tag / \"*\")))\n",
365
+ " assert run_dirs, f\"no run dir for {tag}\"\n",
366
+ " return Path(run_dirs[-1])\n",
367
+ "\n",
368
+ "def newest_save(tag):\n",
369
+ " rd = latest_run_dir(tag)\n",
370
+ " saves = sorted(glob.glob(str(rd / \"epoch*\")) +\n",
371
+ " glob.glob(str(rd / \"step*\")), key=os.path.getmtime)\n",
372
+ " assert saves, f\"{tag}: run wrote no save dir\"\n",
373
+ " return Path(saves[-1])\n",
374
+ "\n",
375
+ "def gate_values(save_dir):\n",
376
+ " blob = torch.load(Path(save_dir) / \"aleph_relay.pt\",\n",
377
+ " map_location=\"cpu\", weights_only=True)\n",
378
+ " st = blob[\"relays\"]\n",
379
+ " out = []\n",
380
+ " for k in sorted(st, key=int):\n",
381
+ " g = st[k].get(\"gate\")\n",
382
+ " if g is None: # multiband: per-band gates\n",
383
+ " g = st[k][\"gates\"]\n",
384
+ " out.extend(torch.as_tensor(g).float().flatten().tolist())\n",
385
+ " return out\n",
386
+ "\n",
387
+ "def eval_curve(tag):\n",
388
+ " \"\"\"The fork's deterministic eval scalars, parsed from tfevents.\"\"\"\n",
389
+ " from tensorboard.backend.event_processing.event_accumulator import \\\n",
390
+ " EventAccumulator\n",
391
+ " acc = EventAccumulator(str(latest_run_dir(tag)),\n",
392
+ " size_guidance={\"scalars\": 0})\n",
393
+ " acc.Reload()\n",
394
+ " out = {}\n",
395
+ " for t in acc.Tags()[\"scalars\"]:\n",
396
+ " if t.startswith(\"heldout/\"):\n",
397
+ " out[t] = [(e.step, e.value) for e in acc.Scalars(t)]\n",
398
+ " assert \"heldout/loss\" in out, \\\n",
399
+ " f\"{tag}: no heldout/loss in tfevents ({acc.Tags()['scalars']})\"\n",
400
+ " return out\n",
401
+ "\n",
402
+ "def ship_arm(package, arm, save_dir, extra=None):\n",
403
+ " \"\"\"Ship-on-completion: ckpt (+canonical safetensors) + TOML + log tail\n",
404
+ " into the hub package immediately.\"\"\"\n",
405
+ " from huggingface_hub import CommitOperationAdd, HfApi\n",
406
+ " from amoe.io.checkpoint import load_diffusion_anchor\n",
407
+ " from amoe.io.safetensors_io import save_anchor_safetensors\n",
408
+ " sd = Path(save_dir)\n",
409
+ " ops = []\n",
410
+ " for f in sd.iterdir():\n",
411
+ " if f.suffix in (\".pt\", \".safetensors\", \".toml\") or \\\n",
412
+ " f.name == \"adapter_config.json\":\n",
413
+ " ops.append((f\"{package}/{arm}/{f.name}\", str(f)))\n",
414
+ " relay_pt = sd / \"aleph_relay.pt\"\n",
415
+ " if relay_pt.exists():\n",
416
+ " ck = load_diffusion_anchor(str(relay_pt), substrate={\n",
417
+ " \"family\": \"cosmos_dit\",\n",
418
+ " \"base_model_id\": \"circlestone-labs/Anima@anima-base-v1.0\"})\n",
419
+ " canon = sd / \"aleph_relay.canonical.safetensors\"\n",
420
+ " save_anchor_safetensors(ck, str(canon))\n",
421
+ " ops.append((f\"{package}/{arm}/aleph_relay.canonical.safetensors\",\n",
422
+ " str(canon)))\n",
423
+ " tag_guess = sd.parent.parent.name\n",
424
+ " log = RUNS / f\"{tag_guess}.log\"\n",
425
+ " if log.exists():\n",
426
+ " tail = \"\\n\".join(log.read_text(errors=\"replace\").splitlines()[-400:])\n",
427
+ " ops.append((f\"{package}/{arm}/log_tail.txt\",\n",
428
+ " io.BytesIO(tail.encode())))\n",
429
+ " for repo_path, content in (extra or {}).items():\n",
430
+ " body = (io.BytesIO(json.dumps(content, indent=1).encode())\n",
431
+ " if isinstance(content, (dict, list)) else content)\n",
432
+ " ops.append((repo_path, body))\n",
433
+ " HfApi().create_commit(\n",
434
+ " repo_id=HUB_REPO,\n",
435
+ " operations=[CommitOperationAdd(path_in_repo=r, path_or_fileobj=s)\n",
436
+ " for r, s in ops],\n",
437
+ " commit_message=f\"{package}/{arm}: ship-on-completion from Colab\")\n",
438
+ " print(f\"[ship] {package}/{arm}: {len(ops)} file(s) -> {HUB_REPO}\")\n",
439
+ "\n",
440
+ "# -- SMOKE ----------------------------------------------------------------\n",
441
+ "_save = run_arm(\"smoke\")\n",
442
+ "_g = gate_values(_save)\n",
443
+ "_moved = [x for x in _g if abs(x + 3.0) > 1e-3]\n",
444
+ "print(f\"G-GATE: {len(_moved)}/{len(_g)} gates moved off -3.0 \"\n",
445
+ " f\"(range {min(_g):+.4f} .. {max(_g):+.4f})\")\n",
446
+ "assert _moved, (\n",
447
+ " \"G-GATE FAILED: gates did not move \u2014 fp32 masters are not effective \"\n",
448
+ " \"here. STOP. The designed fallback is the gate delta-reparam \"\n",
449
+ " \"(gate = -3 + near-zero delta); do not burn full arms.\")\n",
450
+ "_c = eval_curve(\"smoke\")\n",
451
+ "print(f\"eval wiring OK: {len(_c)} heldout tags, \"\n",
452
+ " f\"step-0 loss {_c['heldout/loss'][0][1]:.5f}\")\n",
453
+ "print(\"SMOKE: ALL GREEN \u2014 masters proven, eval proven, cache built\")\n"
454
+ ]
455
+ },
456
+ {
457
+ "cell_type": "markdown",
458
+ "metadata": {},
459
+ "source": [
460
+ "## exp004b \u2014 the relay close-out\n",
461
+ "\n",
462
+ "Preregistration, stated before spend:\n",
463
+ "- **P1** relay \u2265 matched LoRA-r16 on the paired held-out eval\n",
464
+ "- **P2** toggle bit-exact post-train (val battery, CELL 6)\n",
465
+ "- **P3** gates MOVE at full scale (the exp004 failure, asserted per arm)\n",
466
+ "- **P4** the old ~3k-step loss is a floor: the 10k eval curve descends\n",
467
+ " past it\n",
468
+ "\n",
469
+ "LoRA control s1 runs only if the s0 relay-vs-lora gap is within ~2\u00d7 the\n",
470
+ "relay seed spread (results-gated, disclosed either way).\n"
471
+ ]
472
+ },
473
+ {
474
+ "cell_type": "code",
475
+ "metadata": {},
476
+ "execution_count": null,
477
+ "outputs": [],
478
+ "source": [
479
+ "# \u2550\u2550\u2550 CELL 3 \u2014 exp004b: relay s0, relay s1, LoRA control \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
480
+ "relay_saves = {}\n",
481
+ "for _tag in (\"e004b_relay_s0\", \"e004b_relay_s1\"):\n",
482
+ " _save = run_arm(_tag) # run_arm(_tag, resume=True) after\n",
483
+ " _g = gate_values(_save) # a disconnect\n",
484
+ " _mv = [x for x in _g if abs(x + 3.0) > 1e-3]\n",
485
+ " print(f\"[{_tag}] P3 gates moved: {len(_mv)}/{len(_g)} \"\n",
486
+ " f\"(range {min(_g):+.4f}..{max(_g):+.4f})\")\n",
487
+ " assert _mv, f\"{_tag}: P3 FAILED \u2014 gates frozen at full scale\"\n",
488
+ " _arm = _tag.replace(\"e004b_\", \"\")\n",
489
+ " ship_arm(\"exp004b_anima_relay\", _arm, _save,\n",
490
+ " extra={f\"exp004b_anima_relay/{_arm}/eval_curve.json\":\n",
491
+ " eval_curve(_tag)})\n",
492
+ " relay_saves[_tag] = str(_save)\n",
493
+ "\n",
494
+ "lora_saves = {\"e004b_lora_s0\": str(run_arm(\"e004b_lora_s0\"))}\n",
495
+ "ship_arm(\"exp004b_anima_relay\", \"lora_s0\",\n",
496
+ " Path(lora_saves[\"e004b_lora_s0\"]),\n",
497
+ " extra={\"exp004b_anima_relay/lora_s0/eval_curve.json\":\n",
498
+ " eval_curve(\"e004b_lora_s0\")})\n",
499
+ "\n",
500
+ "def _final(tag):\n",
501
+ " return eval_curve(tag)[\"heldout/loss\"][-1][1]\n",
502
+ "\n",
503
+ "_r0, _r1 = _final(\"e004b_relay_s0\"), _final(\"e004b_relay_s1\")\n",
504
+ "_l0 = _final(\"e004b_lora_s0\")\n",
505
+ "_spread = abs(_r0 - _r1)\n",
506
+ "_gap = abs(min(_r0, _r1) - _l0)\n",
507
+ "print(f\"relay s0 {_r0:.5f} | s1 {_r1:.5f} (spread {_spread:.5f}) | \"\n",
508
+ " f\"lora s0 {_l0:.5f} (gap {_gap:.5f})\")\n",
509
+ "if _gap <= 2 * max(_spread, 1e-6):\n",
510
+ " print(\"GATE: gap within 2x seed spread -> LoRA s1 REQUIRED\")\n",
511
+ " lora_saves[\"e004b_lora_s1\"] = str(run_arm(\"e004b_lora_s1\"))\n",
512
+ " ship_arm(\"exp004b_anima_relay\", \"lora_s1\",\n",
513
+ " Path(lora_saves[\"e004b_lora_s1\"]),\n",
514
+ " extra={\"exp004b_anima_relay/lora_s1/eval_curve.json\":\n",
515
+ " eval_curve(\"e004b_lora_s1\")})\n",
516
+ "else:\n",
517
+ " print(\"GATE: ordering decisive at s0 -> LoRA s1 skipped (disclosed)\")\n"
518
+ ]
519
+ },
520
+ {
521
+ "cell_type": "markdown",
522
+ "metadata": {},
523
+ "source": [
524
+ "## exp017 \u2014 multiband3 on the Anima DiT\n",
525
+ "\n",
526
+ "First live run of `aleph_relay_mode='multiband3'` (CPU-smoked only until\n",
527
+ "now). Preregistration:\n",
528
+ "- **P1** band lesions surgical on the paired eval: own-band damage \u226510\u00d7\n",
529
+ " cross-band at matching quantiles (the exp008 instrument on a DiT)\n",
530
+ "- **P2** toggle bit-exact; every-band-lesioned bit-exact\n",
531
+ "- On record: the aggregate loss may well favor a monolith (Law-1) \u2014 exp017\n",
532
+ " certifies the *mechanism*, not the loss number.\n"
533
+ ]
534
+ },
535
+ {
536
+ "cell_type": "code",
537
+ "metadata": {},
538
+ "execution_count": null,
539
+ "outputs": [],
540
+ "source": [
541
+ "# \u2550\u2550\u2550 CELL 4 \u2014 exp017: mb3 smoke, then both seeds \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
542
+ "from amoe.io.checkpoint import load_diffusion_anchor\n",
543
+ "\n",
544
+ "_save = run_arm(\"e017_mb3_smoke\")\n",
545
+ "_ck = load_diffusion_anchor(str(Path(_save) / \"aleph_relay.pt\"))\n",
546
+ "assert _ck.kind == \"multiband3\", \\\n",
547
+ " f\"mb3 smoke saved kind={_ck.kind} \u2014 the mode did not engage\"\n",
548
+ "_g = gate_values(_save)\n",
549
+ "_mv = [x for x in _g if abs(x + 3.0) > 1e-3]\n",
550
+ "print(f\"mb3 smoke: kind OK, {_ck.n_sites} sites, gates moved \"\n",
551
+ " f\"{len(_mv)}/{len(_g)}\")\n",
552
+ "assert _mv, \"mb3 smoke: gates frozen \u2014 same STOP rule as the G-gate\"\n",
553
+ "\n",
554
+ "mb3_saves = {}\n",
555
+ "for _tag in (\"e017_mb3_s0\", \"e017_mb3_s1\"):\n",
556
+ " _save = run_arm(_tag)\n",
557
+ " _arm = _tag.replace(\"e017_\", \"\")\n",
558
+ " ship_arm(\"exp017_anima_multiband\", _arm, _save,\n",
559
+ " extra={f\"exp017_anima_multiband/{_arm}/eval_curve.json\":\n",
560
+ " eval_curve(_tag)})\n",
561
+ " mb3_saves[_tag] = str(_save)\n",
562
+ "print(\"exp017 arms complete:\", list(mb3_saves))\n"
563
+ ]
564
+ },
565
+ {
566
+ "cell_type": "code",
567
+ "metadata": {},
568
+ "execution_count": null,
569
+ "outputs": [],
570
+ "source": [
571
+ "# \u2550\u2550\u2550 CELL 5 \u2014 the paired-val bed source (own cell, reviewable) \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
572
+ "_BED_SRC = r'''\n",
573
+ "# val_bed.py \u2014 paired val battery (subprocess; env-configured, no argparse).\n",
574
+ "# Does what the fork eval cannot: toggle bit-exact asserts, fp32 gate\n",
575
+ "# stats, band-lesion battery \u2014 on IDENTICAL (row, noise, t) triples per\n",
576
+ "# arm. Reuses the fork's eval cache (latents + text embeds) and the fork's\n",
577
+ "# own transformer + llm_adapter call path (mirrored from\n",
578
+ "# cosmos_predict2.py InitialLayer/LLMAdapterLayer).\n",
579
+ "import json, os, sys\n",
580
+ "sys.path.insert(0, \".\") # fork root (cwd)\n",
581
+ "import utils.common # bind fork utils BEFORE ComfyUI\n",
582
+ "sys.path.append(\"submodules/ComfyUI\")\n",
583
+ "os.environ.setdefault(\"MASTER_ADDR\", \"127.0.0.1\")\n",
584
+ "os.environ.setdefault(\"MASTER_PORT\", \"29571\")\n",
585
+ "os.environ.setdefault(\"RANK\", \"0\")\n",
586
+ "os.environ.setdefault(\"WORLD_SIZE\", \"1\")\n",
587
+ "os.environ.setdefault(\"LOCAL_RANK\", \"0\")\n",
588
+ "import deepspeed\n",
589
+ "deepspeed.init_distributed()\n",
590
+ "import torch\n",
591
+ "import pyarrow.parquet as pq\n",
592
+ "\n",
593
+ "CFG = json.load(open(os.environ[\"VB_CONFIG\"]))\n",
594
+ "QUANTILES = [0.1, 0.3, 0.5, 0.7, 0.9]\n",
595
+ "DEV = \"cuda\"\n",
596
+ "\n",
597
+ "# -- read the fork's eval cache (schema discovered loudly) ----------------\n",
598
+ "cache_dir = CFG[\"eval_cache_dir\"]\n",
599
+ "shard_files = sorted(f for f in os.listdir(cache_dir)\n",
600
+ " if f.endswith(\".parquet\"))\n",
601
+ "assert shard_files, f\"no cache shards under {cache_dir}\"\n",
602
+ "schema = pq.read_schema(os.path.join(cache_dir, shard_files[0]))\n",
603
+ "print(\"cache columns:\", schema.names, flush=True)\n",
604
+ "NEEDED = [\"latents\", \"prompt_embeds\", \"attn_mask\", \"t5_input_ids\",\n",
605
+ " \"t5_attn_mask\"]\n",
606
+ "for c in NEEDED:\n",
607
+ " assert c in schema.names or c + \"__shape\" in schema.names, \\\n",
608
+ " f\"cache lacks column {c} \u2014 inspect the printed schema\"\n",
609
+ "\n",
610
+ "def decol(tbl, name, j):\n",
611
+ " raw = tbl.column(name)[j].as_py()\n",
612
+ " shape = tbl.column(name + \"__shape\")[j].as_py()\n",
613
+ " dt = str(tbl.column(name + \"__dtype\")[j].as_py()).replace(\"torch.\", \"\")\n",
614
+ " return torch.frombuffer(bytearray(raw),\n",
615
+ " dtype=getattr(torch, dt)).reshape(shape)\n",
616
+ "\n",
617
+ "rows = []\n",
618
+ "for sf in shard_files:\n",
619
+ " tbl = pq.read_table(os.path.join(cache_dir, sf))\n",
620
+ " for j in range(tbl.num_rows):\n",
621
+ " rows.append({c: decol(tbl, c, j) for c in NEEDED\n",
622
+ " if c + \"__shape\" in schema.names})\n",
623
+ "rows = rows[:int(os.environ.get(\"VB_MAX_ROWS\", \"512\"))]\n",
624
+ "print(f\"eval cache rows used: {len(rows)}\", flush=True)\n",
625
+ "assert rows\n",
626
+ "\n",
627
+ "# -- pipeline (frozen trunk); adapters swapped per arm --------------------\n",
628
+ "from models import cosmos_predict2\n",
629
+ "mc = CFG[\"model_cfg\"]\n",
630
+ "mc[\"model\"][\"dtype\"] = torch.bfloat16\n",
631
+ "pipe = cosmos_predict2.CosmosPredict2Pipeline(mc)\n",
632
+ "pipe.load_diffusion_model()\n",
633
+ "tf = pipe.transformer.to(DEV).eval()\n",
634
+ "from models.aleph_relay import attach_aleph_relays, band_weights\n",
635
+ "\n",
636
+ "def clear_adapters():\n",
637
+ " for b in tf.blocks:\n",
638
+ " for attr in (\"aleph_relay\", \"aleph_w_bands\"):\n",
639
+ " if hasattr(b, attr):\n",
640
+ " delattr(b, attr)\n",
641
+ "\n",
642
+ "# text embeds through llm_adapter ONCE per row (t-independent) \u2014 the exact\n",
643
+ "# call LLMAdapterLayer makes (cosmos_predict2.py:643-649)\n",
644
+ "la = getattr(tf, \"llm_adapter\", None)\n",
645
+ "embeds = []\n",
646
+ "with torch.no_grad():\n",
647
+ " for r in rows:\n",
648
+ " e = r[\"prompt_embeds\"].to(DEV, torch.bfloat16)\n",
649
+ " if e.ndim == 2:\n",
650
+ " e = e.unsqueeze(0)\n",
651
+ " if la is not None:\n",
652
+ " am = r[\"attn_mask\"].to(DEV).reshape(1, -1)\n",
653
+ " ids = r[\"t5_input_ids\"].to(DEV).reshape(1, -1)\n",
654
+ " t5m = r[\"t5_attn_mask\"].to(DEV).reshape(1, -1)\n",
655
+ " e = la(source_hidden_states=e, target_input_ids=ids,\n",
656
+ " target_attention_mask=t5m, source_attention_mask=am)\n",
657
+ " e = e.clone()\n",
658
+ " e[~t5m.bool()] = 0\n",
659
+ " embeds.append(e)\n",
660
+ "print(\"embeds prepared (post-llm_adapter)\", flush=True)\n",
661
+ "\n",
662
+ "# fixed noise bank + the fork's exact eval-t transform:\n",
663
+ "# t = sigmoid(Normal.icdf(q)) (logit_normal, scale 1, no shift \u2014\n",
664
+ "# cosmos_predict2.py:436-444)\n",
665
+ "torch.manual_seed(1400)\n",
666
+ "lats, noises = [], []\n",
667
+ "for r in rows:\n",
668
+ " l = r[\"latents\"].float()\n",
669
+ " while l.ndim < 5:\n",
670
+ " l = l.unsqueeze(0)\n",
671
+ " lats.append(l)\n",
672
+ " noises.append(torch.randn_like(l))\n",
673
+ "from torch.distributions import Normal\n",
674
+ "T_OF_Q = {q: float(torch.sigmoid(Normal(0.0, 1.0).icdf(torch.tensor(q))))\n",
675
+ " for q in QUANTILES}\n",
676
+ "print(\"t(q):\", {q: round(t, 4) for q, t in T_OF_Q.items()}, flush=True)\n",
677
+ "\n",
678
+ "results = {\"quantiles\": QUANTILES, \"t_of_q\": T_OF_Q, \"n_rows\": len(rows),\n",
679
+ " \"arms\": {}}\n",
680
+ "\n",
681
+ "def eval_arm(name, lesion=None, needs_bands=False):\n",
682
+ " per_q = {}\n",
683
+ " with torch.no_grad():\n",
684
+ " for q in QUANTILES:\n",
685
+ " t = T_OF_Q[q]\n",
686
+ " if needs_bands:\n",
687
+ " w = band_weights(torch.tensor([t], device=DEV))\n",
688
+ " for b in tf.blocks:\n",
689
+ " b.aleph_w_bands = w\n",
690
+ " losses = []\n",
691
+ " for l, nz, e in zip(lats, noises, embeds):\n",
692
+ " ld = l.to(DEV, torch.bfloat16)\n",
693
+ " nd = nz.to(DEV, torch.bfloat16)\n",
694
+ " x_t = (1 - t) * ld + t * nd\n",
695
+ " ts = torch.full((ld.shape[0],), t, device=DEV,\n",
696
+ " dtype=torch.bfloat16)\n",
697
+ " pred = tf(x_t, ts, e)\n",
698
+ " if isinstance(pred, (list, tuple)):\n",
699
+ " pred = pred[0]\n",
700
+ " losses.append(((pred.float().cpu() - (nz - l)) ** 2)\n",
701
+ " .mean().item())\n",
702
+ " per_q[str(q)] = sum(losses) / len(losses)\n",
703
+ " results[\"arms\"][name] = {\"loss_by_quantile\": per_q,\n",
704
+ " \"loss_mean\": sum(per_q.values()) / len(per_q),\n",
705
+ " \"lesion\": lesion}\n",
706
+ " print(f\"[{name}] mean {results['arms'][name]['loss_mean']:.6f}\",\n",
707
+ " flush=True)\n",
708
+ "\n",
709
+ "clear_adapters()\n",
710
+ "eval_arm(\"frozen\")\n",
711
+ "f_ref = dict(results[\"arms\"][\"frozen\"][\"loss_by_quantile\"])\n",
712
+ "eval_arm(\"frozen_repeat\")\n",
713
+ "assert results[\"arms\"][\"frozen_repeat\"][\"loss_by_quantile\"] == f_ref, \\\n",
714
+ " \"DETERMINISM FAILED \u2014 paired comparison meaningless; stop\"\n",
715
+ "\n",
716
+ "for name, arm in CFG[\"arms\"].items():\n",
717
+ " clear_adapters()\n",
718
+ " mode = arm.get(\"kind\", \"relay\")\n",
719
+ " attach_aleph_relays(tf, tf.model_channels, relay_path=arm[\"ckpt\"],\n",
720
+ " dtype=torch.bfloat16, mode=mode)\n",
721
+ " gates = []\n",
722
+ " for b in tf.blocks:\n",
723
+ " g = getattr(b.aleph_relay, \"gate\", None)\n",
724
+ " if g is None:\n",
725
+ " g = b.aleph_relay.gates\n",
726
+ " gates.extend(torch.as_tensor(g).float().flatten().tolist())\n",
727
+ " for b in tf.blocks: # toggle assert per arm\n",
728
+ " b.aleph_relay.enabled = False\n",
729
+ " eval_arm(name + \"_toggled_off\", needs_bands=(mode == \"multiband3\"))\n",
730
+ " assert results[\"arms\"][name + \"_toggled_off\"][\"loss_by_quantile\"] \\\n",
731
+ " == f_ref, f\"{name}: TOGGLE LAW VIOLATED\"\n",
732
+ " for b in tf.blocks:\n",
733
+ " b.aleph_relay.enabled = True\n",
734
+ " if arm.get(\"lesion\") is not None:\n",
735
+ " b.aleph_relay.band_enabled[arm[\"lesion\"]] = False\n",
736
+ " eval_arm(name, lesion=arm.get(\"lesion\"),\n",
737
+ " needs_bands=(mode == \"multiband3\"))\n",
738
+ " results[\"arms\"][name][\"gates_fp32\"] = {\n",
739
+ " \"min\": min(gates), \"max\": max(gates),\n",
740
+ " \"moved\": sum(1 for g in gates if abs(g + 3.0) > 1e-3)}\n",
741
+ "\n",
742
+ "json.dump(results, open(CFG[\"out\"], \"w\"), indent=1)\n",
743
+ "print(\"WROTE\", CFG[\"out\"], flush=True)\n",
744
+ "'''\n",
745
+ "VAL_BED = WORK / \"val_bed.py\"\n",
746
+ "VAL_BED.write_text(_BED_SRC)\n",
747
+ "print(f\"val_bed.py written: {VAL_BED} ({len(_BED_SRC)/1000:.1f}k chars)\")\n"
748
+ ]
749
+ },
750
+ {
751
+ "cell_type": "code",
752
+ "metadata": {},
753
+ "execution_count": null,
754
+ "outputs": [],
755
+ "source": [
756
+ "# \u2550\u2550\u2550 CELL 6 \u2014 run the val battery + assemble results.json \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
757
+ "import glob, json\n",
758
+ "try:\n",
759
+ " import tomllib\n",
760
+ "except ImportError:\n",
761
+ " import tomli as tomllib\n",
762
+ "\n",
763
+ "_pq_files = glob.glob(str(DATA / \"**\" / \"*.parquet\"), recursive=True)\n",
764
+ "_cache_dirs = sorted({str(Path(f).parent) for f in _pq_files\n",
765
+ " if \"cache\" in f.lower() and \"eval\" in f.lower()})\n",
766
+ "print(\"eval cache candidates:\", _cache_dirs)\n",
767
+ "assert _cache_dirs, (\"no eval cache found under DATA \u2014 inspect the tree, \"\n",
768
+ " \"set EVAL_CACHE manually, re-run\")\n",
769
+ "EVAL_CACHE = _cache_dirs[0]\n",
770
+ "\n",
771
+ "_full = tomllib.loads(ARMS[\"e004b_relay_s0\"].read_text())\n",
772
+ "_model = {k: v for k, v in _full[\"model\"].items()\n",
773
+ " if not k.startswith(\"aleph_\") and not k.endswith(\"_lr\")}\n",
774
+ "_model_cfg = {\"model\": _model} # the bed attaches adapters itself\n",
775
+ "\n",
776
+ "def _bed(config, out_name):\n",
777
+ " cfgp = WORK / f\"vb_{out_name}.json\"\n",
778
+ " outp = WORK / f\"{out_name}.json\"\n",
779
+ " config.update(out=str(outp), eval_cache_dir=EVAL_CACHE,\n",
780
+ " model_cfg=_model_cfg)\n",
781
+ " cfgp.write_text(json.dumps(config, indent=1))\n",
782
+ " sh(f\"cd {FORK} && VB_CONFIG={cfgp} python {WORK}/val_bed.py \"\n",
783
+ " f\"2>&1 | tee {RUNS}/{out_name}.log\")\n",
784
+ " return json.loads(outp.read_text())\n",
785
+ "\n",
786
+ "def _pt(saves, tag):\n",
787
+ " return str(Path(saves[tag]) / \"aleph_relay.pt\")\n",
788
+ "\n",
789
+ "r4 = _bed({\"arms\": {\n",
790
+ " \"relay_s0\": {\"kind\": \"relay\", \"ckpt\": _pt(relay_saves, \"e004b_relay_s0\")},\n",
791
+ " \"relay_s1\": {\"kind\": \"relay\", \"ckpt\": _pt(relay_saves, \"e004b_relay_s1\")},\n",
792
+ "}}, \"results_exp004b\")\n",
793
+ "# (the LoRA arm is PEFT-format; its numbers come from the fork eval curve)\n",
794
+ "\n",
795
+ "r17 = _bed({\"arms\": {\n",
796
+ " \"mb3_s0\": {\"kind\": \"multiband3\", \"ckpt\": _pt(mb3_saves, \"e017_mb3_s0\")},\n",
797
+ " \"mb3_s0_lesion_b0\": {\"kind\": \"multiband3\", \"lesion\": 0,\n",
798
+ " \"ckpt\": _pt(mb3_saves, \"e017_mb3_s0\")},\n",
799
+ " \"mb3_s0_lesion_b1\": {\"kind\": \"multiband3\", \"lesion\": 1,\n",
800
+ " \"ckpt\": _pt(mb3_saves, \"e017_mb3_s0\")},\n",
801
+ " \"mb3_s0_lesion_b2\": {\"kind\": \"multiband3\", \"lesion\": 2,\n",
802
+ " \"ckpt\": _pt(mb3_saves, \"e017_mb3_s0\")},\n",
803
+ " \"mb3_s1\": {\"kind\": \"multiband3\", \"ckpt\": _pt(mb3_saves, \"e017_mb3_s1\")},\n",
804
+ "}}, \"results_exp017\")\n",
805
+ "\n",
806
+ "def _final(tag):\n",
807
+ " return eval_curve(tag)[\"heldout/loss\"][-1][1]\n",
808
+ "\n",
809
+ "RES4 = {\"bed\": r4,\n",
810
+ " \"fork_eval\": {t: eval_curve(t) for t in\n",
811
+ " (\"e004b_relay_s0\", \"e004b_relay_s1\", \"e004b_lora_s0\")},\n",
812
+ " \"split_manifest\": json.loads((DATA / \"split_manifest.json\")\n",
813
+ " .read_text())}\n",
814
+ "_rl = min(_final(\"e004b_relay_s0\"), _final(\"e004b_relay_s1\"))\n",
815
+ "_curve0 = RES4[\"fork_eval\"][\"e004b_relay_s0\"][\"heldout/loss\"]\n",
816
+ "_at3k = min(v for s, v in _curve0 if s <= 3000)\n",
817
+ "RES4[\"verdict\"] = {\n",
818
+ " \"P1_relay_vs_lora\": \"relay\" if _rl <= _final(\"e004b_lora_s0\") else \"lora\",\n",
819
+ " \"P2_toggle_bit_exact\": True, # the bed asserted, else it raised\n",
820
+ " \"P3_gates_moved\": True, # per-arm asserts in CELL 3\n",
821
+ " \"P4_3k_was_a_floor\": _curve0[-1][1] < _at3k,\n",
822
+ "}\n",
823
+ "(WORK / \"results_exp004b_full.json\").write_text(json.dumps(RES4, indent=1))\n",
824
+ "ship_arm(\"exp004b_anima_relay\", \"results\", newest_save(\"e004b_relay_s0\"),\n",
825
+ " extra={\"exp004b_anima_relay/results.json\": RES4})\n",
826
+ "print(\"exp004b verdict:\", RES4[\"verdict\"])\n",
827
+ "\n",
828
+ "_q = [str(x) for x in r17[\"quantiles\"]]\n",
829
+ "_allon = r17[\"arms\"][\"mb3_s0\"][\"loss_by_quantile\"]\n",
830
+ "_bandq = {0: [\"0.1\", \"0.3\"], 1: [\"0.5\", \"0.7\"], 2: [\"0.9\"]}\n",
831
+ "_surgical = {}\n",
832
+ "for _b, _own_q in _bandq.items():\n",
833
+ " _les = r17[\"arms\"][f\"mb3_s0_lesion_b{_b}\"][\"loss_by_quantile\"]\n",
834
+ " _own = sum(_les[x] - _allon[x] for x in _own_q) / len(_own_q)\n",
835
+ " _cq = [x for x in _q if x not in _own_q]\n",
836
+ " _cross = sum(_les[x] - _allon[x] for x in _cq) / len(_cq)\n",
837
+ " _surgical[_b] = {\"own_damage\": _own, \"cross_damage\": _cross,\n",
838
+ " \"ratio\": (_own / _cross) if _cross > 0 else None}\n",
839
+ "RES17 = {\"bed\": r17, \"surgical\": _surgical,\n",
840
+ " \"P1_surgical_10x\": all(\n",
841
+ " (s[\"ratio\"] is None and s[\"own_damage\"] > 0) or\n",
842
+ " (s[\"ratio\"] is not None and s[\"ratio\"] >= 10)\n",
843
+ " for s in _surgical.values()),\n",
844
+ " \"P2_toggle_bit_exact\": True}\n",
845
+ "(WORK / \"results_exp017_full.json\").write_text(json.dumps(RES17, indent=1))\n",
846
+ "ship_arm(\"exp017_anima_multiband\", \"results\", newest_save(\"e017_mb3_s0\"),\n",
847
+ " extra={\"exp017_anima_multiband/results.json\": RES17})\n",
848
+ "print(\"exp017 surgical:\", json.dumps(_surgical, indent=1))\n"
849
+ ]
850
+ },
851
+ {
852
+ "cell_type": "code",
853
+ "metadata": {},
854
+ "execution_count": null,
855
+ "outputs": [],
856
+ "source": [
857
+ "# \u2550\u2550\u2550 CELL 7 \u2014 production naming + final ship \u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\u2550\n",
858
+ "# v2 adapters into amoe_loras/production/anima/, usage cards written from\n",
859
+ "# the MEASURED verdicts of this run. Local follow-ups: catalog regen,\n",
860
+ "# ComfyUI KNOWN_ADAPTERS + live verify, write-back to the mind.\n",
861
+ "import json\n",
862
+ "from amoe.io.checkpoint import load_diffusion_anchor\n",
863
+ "from amoe.io.safetensors_io import save_anchor_safetensors\n",
864
+ "from huggingface_hub import CommitOperationAdd, HfApi\n",
865
+ "\n",
866
+ "def _named(save, out_name, display, usage, evidence, seed, caveats,\n",
867
+ " bands=None):\n",
868
+ " ck = load_diffusion_anchor(str(Path(save) / \"aleph_relay.pt\"),\n",
869
+ " substrate={\n",
870
+ " \"family\": \"cosmos_dit\",\n",
871
+ " \"base_model_id\": \"circlestone-labs/Anima@anima-base-v1.0\"})\n",
872
+ " ck.meta.update(display_name=display, usage=usage, evidence=evidence,\n",
873
+ " tier=\"capability\", seed=seed, caveats=caveats, nc=True,\n",
874
+ " license=\"NC (CircleStone NC + NVIDIA Open Model License)\",\n",
875
+ " objective={\"kind\": \"flow\"}, recommended_strength=1.0)\n",
876
+ " if bands:\n",
877
+ " ck.meta[\"band_roles\"] = bands\n",
878
+ " p = WORK / f\"{out_name}.safetensors\"\n",
879
+ " save_anchor_safetensors(ck, str(p))\n",
880
+ " return (f\"amoe_loras/production/anima/{out_name}.safetensors\", str(p))\n",
881
+ "\n",
882
+ "_BANDS = [\"fidelity/detail (LOW noise)\", \"continuity/semantics (MID)\",\n",
883
+ " \"diversity/structure (HIGH noise)\"]\n",
884
+ "_v = RES4[\"verdict\"]\n",
885
+ "_ops = []\n",
886
+ "for _seed, _tag in ((0, \"e004b_relay_s0\"), (1, \"e004b_relay_s1\")):\n",
887
+ " _ops.append(_named(\n",
888
+ " relay_saves[_tag], f\"anima-dit-relay-v2-s{_seed}\",\n",
889
+ " f\"Anima 2B DiT Relay v2 (seed {_seed}) \u2014 NON-COMMERCIAL\",\n",
890
+ " \"The exp004b retrain: fp32 master weights (working gates), 10k \"\n",
891
+ " \"steps, full dataset, declared seed, real held-out eval.\",\n",
892
+ " f\"exp004b: P1 relay_vs_lora={_v['P1_relay_vs_lora']}, \"\n",
893
+ " f\"P4 old-3k-was-a-floor={_v['P4_3k_was_a_floor']}; numbers in \"\n",
894
+ " \"exp004b_anima_relay/results.json.\",\n",
895
+ " _seed,\n",
896
+ " [\"Non-commercial (derived from NC weights).\",\n",
897
+ " \"Needs an Anima checkpoint (28 sites, width 2048).\"]))\n",
898
+ "for _seed, _tag in ((0, \"e017_mb3_s0\"), (1, \"e017_mb3_s1\")):\n",
899
+ " _ops.append(_named(\n",
900
+ " mb3_saves[_tag], f\"anima-dit-multiband-v1-s{_seed}\",\n",
901
+ " f\"Anima 2B DiT Multiband v1 (seed {_seed}) \u2014 NON-COMMERCIAL\",\n",
902
+ " \"Three sigma-band experts on the Anima DiT, step-gated \u2014 the \"\n",
903
+ " \"first multiband stack on a DiT-class trunk.\",\n",
904
+ " f\"exp017: P1 surgical(>=10x)={RES17['P1_surgical_10x']}; lesion \"\n",
905
+ " \"numbers in exp017_anima_multiband/results.json.\",\n",
906
+ " _seed,\n",
907
+ " [\"Non-commercial (derived from NC weights).\",\n",
908
+ " \"Needs an Anima checkpoint; band gating needs a step-gated \"\n",
909
+ " \"loader (ComfyUI nodes / StepGatedSampler).\"],\n",
910
+ " bands=_BANDS))\n",
911
+ "\n",
912
+ "HfApi().create_commit(\n",
913
+ " repo_id=HUB_REPO,\n",
914
+ " operations=[CommitOperationAdd(path_in_repo=r, path_or_fileobj=s)\n",
915
+ " for r, s in _ops],\n",
916
+ " commit_message=\"Anima v2 production adapters (exp004b/exp017), usage \"\n",
917
+ " \"cards from measured verdicts\")\n",
918
+ "print(\"shipped:\", [r for r, _s in _ops])\n",
919
+ "print(\"\\nCAMPAIGN COMPLETE \u2014 local follow-ups: catalog regen, ComfyUI \"\n",
920
+ " \"KNOWN_ADAPTERS + live verify, write-back to the mind.\")\n"
921
+ ]
922
+ }
923
+ ]
924
+ }