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- 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 |
+
}
|