AbstractPhil commited on
Commit
afe3c0e
Β·
verified Β·
1 Parent(s): ed8bcf5

Create trainer_v1.py

Browse files
Files changed (1) hide show
  1. trainer_v1.py +1053 -0
trainer_v1.py ADDED
@@ -0,0 +1,1053 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ============================================================================
2
+ # CAPTIONBERT-8192-v2 β€” CONSENSUS DISTILLATION AT CC12M SCALE
3
+ #
4
+ # v2 vs the shipped 500k model, per Phil's 2026-07-31 guidance:
5
+ # - NO ALIGNMENT BANK. v1's bank was additive and experimental; measured on real
6
+ # embeddings its expert-consistency block varied 0.2% across samples and took
7
+ # 0.23% of geo_proj energy while anchor distances took 98.7%. Banks in this
8
+ # format are content extensions β€” an AMOE-LORA is the right carrier, attached
9
+ # as a separate finetune pass on the prefitted core. Not here.
10
+ # - LEGROOM. d 384->512, 6L->12L, ff 1536->2048, heads 6->8. 26.0M -> 58.3M
11
+ # (0.53x bert-base, so the compression story survives). Sized for many
12
+ # overlapping sources at ~36M features/teacher, not one 500k census.
13
+ # - CHAMPION OBJECTIVE. InfoNCE + per-sample MSE against the consensus β€” the
14
+ # consensus_nce_mse form that won the CC12M vision matrix on every task gauge,
15
+ # both seeds. NO shipped rotation needed here: that line aligns to a running
16
+ # mean (frame free), this one aligns to a REFERENCE MEMBER (bert), so the frame
17
+ # is pinned by construction. A frame-fit gauge runs anyway to confirm it.
18
+ # - CULL-PROOF. Colab kills the VM every 24h and takes local disk with it.
19
+ # Full state (model/opt/sched/scaler/step/epoch/chunk-order/RNG) checkpoints on
20
+ # a TIME cadence, and pushes to HF so a cull costs minutes, not the run.
21
+ # - FULL TENSORBOARD. per-step losses + lr + grad-norm, per-eval gauges
22
+ # (mimicry, cos, isotropy, effective rank, CV), histograms, and the alignment
23
+ # report as text.
24
+ #
25
+ # STAGES (each resumable, each gated) β€” carried from the cc12m pipeline:
26
+ # 0 PARITY which caption field was embedded + row alignment. Hard gate.
27
+ # 1 FIT one global whitened-Procrustes map per expert -> bert, stratified
28
+ # random fit, reported OUT-OF-SAMPLE on held-out chunks.
29
+ # 2 TARGETS per-chunk consensus -> fp16, ledgered, expert shards deleted after.
30
+ # 3 TRAIN streams (captions, consensus) pairs, dynamic padding.
31
+ #
32
+ # Colab-cell-safe. HF_TOKEN from Colab secrets (key icon) or env.
33
+ # ============================================================================
34
+
35
+ import gc, json, math, os, random, sys, time, subprocess, shutil
36
+ from dataclasses import dataclass, asdict
37
+ from typing import Any, Dict, List, Optional, Tuple
38
+
39
+ for _p in ("datasets", "transformers", "huggingface_hub", "tensorboard", "safetensors"):
40
+ try:
41
+ __import__(_p)
42
+ except ImportError:
43
+ subprocess.run([sys.executable, "-m", "pip", "install", "-q", _p], check=False)
44
+
45
+ import numpy as np
46
+ import torch
47
+ import torch.nn as nn
48
+ import torch.nn.functional as F
49
+ from huggingface_hub import hf_hub_download, HfApi, create_repo
50
+ from torch.utils.tensorboard import SummaryWriter
51
+
52
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
53
+
54
+
55
+ # ══════════════════════════════════════════════════════════════════
56
+ # BASE CONFIG
57
+ # ══════════════════════════════════════════════════════════════════
58
+
59
+ @dataclass
60
+ class BaseConfig:
61
+ run_name: str = "captionbert-8192-v2"
62
+
63
+ # ── sources ── (list so overlapping datasets can be added later)
64
+ sources: Tuple[Dict[str, Any], ...] = (
65
+ {"repo": "AbstractPhil/conceptual-captions-12m-webdataset-berts",
66
+ "n_chunks": 66, "chunk_rows": 500_000,
67
+ "missing": {"modern": (5, 7, 8, 21, 25, 26, 28, 32, 38, 46)}},
68
+ )
69
+ experts: Tuple[str, ...] = ("bert", "modern", "roberta", "albert", "distil")
70
+ ref_expert: str = "bert"
71
+ ref_hf_name: str = "google-bert/bert-base-uncased"
72
+ require_all_experts: bool = True
73
+ caption_field: Optional[str] = None
74
+ caption_field_candidates: Tuple[str, ...] = (
75
+ "caption_llava", "caption", "caption_llava_short")
76
+
77
+ work_dir: str = "/content/cbv2"
78
+ keep_expert_shards: bool = False
79
+
80
+ # ── hardware allowance (Colab Pro+ / RTX 6000 Pro, measured 2026-07-31) ──
81
+ # disk 235.7GB (~176 free) | RAM 176.9GB | GPU 95.6GB | 401.5 units @ 8.9/h = 45.1h
82
+ # The expert shards are 507GB β€” 2.1x the WHOLE DISK. They are streamed one chunk
83
+ # at a time and deleted; only the 43GB consensus is kept.
84
+ disk_floor_gb: float = 25.0 # abort a chunk if free disk drops below
85
+ ram_resident: bool = True # hold tokens+targets in RAM (48.8GB)
86
+ preflight: bool = True
87
+
88
+ # ── backup (Colab culls at 24h; local disk dies with the VM) ──
89
+ hf_repo: str = "AbstractPhil/captionbert-8192-v2"
90
+ targets_repo: str = "AbstractPhil/captionbert-8192-v2-consensus"
91
+ push_targets: bool = True # 43GB; re-derivable only from a 507GB pull
92
+ hf_push: bool = True
93
+ push_every_min: float = 30.0
94
+ keep_local_ckpts: int = 3
95
+
96
+ # ── stage 0 ──
97
+ parity_chunk: int = 0
98
+ parity_n: int = 64
99
+ parity_min_cos: float = 0.999
100
+
101
+ # ── stage 1 ──
102
+ fit_chunks: Tuple[int, ...] = (0, 11, 22, 33, 44, 55)
103
+ fit_rows_per_chunk: int = 4000 # 24k vs d=768 -> N/d = 31
104
+ holdout_chunks: Tuple[int, ...] = (60, 61)
105
+ fit_seed: int = 0
106
+
107
+ # ── student (LEGROOM) ──
108
+ d_model: int = 512 # was 384
109
+ n_heads: int = 8 # was 6
110
+ n_layers: int = 12 # was 6
111
+ d_ff: int = 2048 # was 1536
112
+ max_len: int = 8192 # name-bearing; costs 4.2M params
113
+ output_dim: int = 768 # consensus space = teacher dim
114
+ dropout: float = 0.1
115
+ pooling: str = "mean" # arm: "cls". teachers are mean-pooled
116
+ max_tokens: int = 256 # dynamic pad ceiling
117
+
118
+ # ── training (sized for 95.6GB GPU: batch size IS the InfoNCE negative count) ──
119
+ epochs: int = 4 # 13.7k steps/ep at 2048 -> ~55k total
120
+ batch_size: int = 2048 # was 512; ~19GB activations, 4x negatives
121
+ lr: float = 6e-4 # sqrt-scaled from 3e-4 @ 512
122
+ min_lr: float = 1e-6
123
+ warmup_steps: int = 2000
124
+ grad_clip: float = 1.0
125
+ seed: int = 42
126
+ amp: bool = True
127
+ num_workers: int = 0 # RAM-resident: no workers needed
128
+ log_every: int = 50
129
+ eval_every: int = 1000
130
+ ckpt_every_min: float = 20.0 # TIME-based: culls are wall-clock
131
+
132
+ # ── loss: the champion form ──
133
+ nce_weight: float = 1.0
134
+ mse_weight: float = 1.0
135
+ nce_temperature: float = 0.07
136
+ cv_weight: float = 0.0 # arm: 0.1 reproduces the v1 stack
137
+ cv_target: float = 0.084
138
+
139
+ # ── stages ──
140
+ run_stage0: bool = True
141
+ run_stage1: bool = True
142
+ run_stage2: bool = True
143
+ run_stage3: bool = True
144
+ resume: bool = True
145
+
146
+
147
+ CFG = BaseConfig()
148
+
149
+
150
+ # ══════════════════════════════════════════════════════════════════
151
+ # HELPERS
152
+ # ══════════════════════════════════════════════════════════════════
153
+
154
+ def line(t=""):
155
+ print("─" * 78 if not t else f"── {t} " + "─" * max(0, 74 - len(t)))
156
+
157
+
158
+ def paths(cfg) -> Dict[str, str]:
159
+ w = cfg.work_dir
160
+ d = {"root": w, "targets": f"{w}/targets", "maps": f"{w}/maps",
161
+ "ckpt": f"{w}/checkpoints", "tb": f"{w}/tensorboard", "shards": f"{w}/shards",
162
+ "config": f"{w}/config"}
163
+ for p in d.values():
164
+ os.makedirs(p, exist_ok=True)
165
+ return d
166
+
167
+
168
+ def src0(cfg) -> Dict[str, Any]:
169
+ return cfg.sources[0]
170
+
171
+
172
+ def usable_chunks(cfg) -> List[int]:
173
+ s = src0(cfg)
174
+ c = set(range(s["n_chunks"]))
175
+ if cfg.require_all_experts:
176
+ for miss in s.get("missing", {}).values():
177
+ c -= set(miss)
178
+ return sorted(c - set(cfg.holdout_chunks))
179
+
180
+
181
+ def fetch(cfg, fname: str) -> str:
182
+ return hf_hub_download(src0(cfg)["repo"], fname, repo_type="dataset",
183
+ local_dir=paths(cfg)["shards"])
184
+
185
+
186
+ def load_captions_chunk(cfg, c: int) -> List[str]:
187
+ raw = json.load(open(fetch(cfg, f"captions_{c:03d}.json")))
188
+ f = cfg.caption_field
189
+ if isinstance(raw, dict):
190
+ return list(raw[f])
191
+ if raw and isinstance(raw[0], dict):
192
+ return [r[f] for r in raw]
193
+ return list(raw)
194
+
195
+
196
+ def load_expert_chunk(cfg, expert: str, c: int) -> torch.Tensor:
197
+ return torch.load(fetch(cfg, f"{expert}_{c:03d}.pt"),
198
+ weights_only=True, map_location="cpu")
199
+
200
+
201
+ def drop_shard(cfg, fname: str):
202
+ if cfg.keep_expert_shards:
203
+ return
204
+ p = os.path.join(paths(cfg)["shards"], fname)
205
+ if os.path.exists(p):
206
+ os.remove(p)
207
+
208
+
209
+ def free_gb(path: str) -> float:
210
+ st = os.statvfs(path)
211
+ return st.f_bavail * st.f_frsize / 1e9
212
+
213
+
214
+ def purge_hf_cache(cfg):
215
+ """
216
+ The expert shards total 507GB against a 235.7GB disk. hf_hub_download with
217
+ local_dir does not populate the global cache on modern hub versions, but a
218
+ stale HF_HOME cache or an older version WILL duplicate every shard and blow
219
+ the disk mid-run. Purge both, every chunk.
220
+ """
221
+ for d in (os.path.join(paths(cfg)["shards"], ".cache"),
222
+ os.environ.get("HF_HUB_CACHE", ""),
223
+ os.path.expanduser("~/.cache/huggingface/hub")):
224
+ if d and os.path.isdir(d):
225
+ for entry in os.listdir(d):
226
+ if entry.startswith("datasets--"):
227
+ shutil.rmtree(os.path.join(d, entry), ignore_errors=True)
228
+
229
+
230
+ def preflight(cfg):
231
+ """Hard-check the allowance before anything expensive starts."""
232
+ line("PREFLIGHT β€” disk / RAM / GPU vs the plan")
233
+ P = paths(cfg)
234
+ disk = free_gb(P["root"])
235
+ s = src0(cfg)
236
+ n_keep = len(usable_chunks(cfg)) + len(cfg.holdout_chunks)
237
+ rows = n_keep * s["chunk_rows"]
238
+ targets_gb = rows * cfg.output_dim * 2 / 1e9
239
+ transient_gb = len(cfg.experts) * 1.536
240
+ caps_gb = s["n_chunks"] * 0.120
241
+ need = targets_gb + caps_gb + transient_gb + 10.0
242
+ print(f" source on HF : {s['n_chunks'] * len(cfg.experts) * 1.536:.0f} GB expert shards "
243
+ f"(streamed one chunk at a time, deleted after)")
244
+ print(f" disk free : {disk:.1f} GB | stage-2 peak need β‰ˆ {need:.1f} GB "
245
+ f"(targets {targets_gb:.1f} + captions {caps_gb:.1f} + transient {transient_gb:.1f})")
246
+ if disk < need:
247
+ raise RuntimeError(
248
+ f"DISK: {disk:.1f} GB free, need β‰ˆ {need:.1f} GB. Free space, reduce chunks, "
249
+ f"or set push_targets=True and drop consensus locally after each push.")
250
+ try:
251
+ import psutil
252
+ ram = psutil.virtual_memory().total / 1e9
253
+ except Exception:
254
+ ram = float("nan")
255
+ ram_need = (rows * 100 * 2 + rows * 8 + rows * cfg.output_dim * 2) / 1e9
256
+ print(f" RAM total : {ram:.1f} GB | ram_resident need β‰ˆ {ram_need:.1f} GB "
257
+ f"(ragged tokens + offsets + fp16 targets)")
258
+ if cfg.ram_resident and ram == ram and ram_need > 0.7 * ram:
259
+ print(f" !! ram_resident wants {ram_need:.1f} GB of {ram:.1f}. "
260
+ f"Set ram_resident=False to stream per chunk from disk instead.")
261
+ if DEVICE == "cuda":
262
+ g = torch.cuda.get_device_properties(0).total_memory / 1e9
263
+ print(f" GPU : {torch.cuda.get_device_name()} {g:.1f} GB | "
264
+ f"batch {cfg.batch_size} -> {cfg.batch_size} InfoNCE negatives")
265
+ print(f" plan : {rows:,} rows, {rows // cfg.batch_size:,} steps/epoch "
266
+ f"x {cfg.epochs} = {rows // cfg.batch_size * cfg.epochs:,} steps")
267
+
268
+
269
+ def effective_rank(x: torch.Tensor) -> float:
270
+ xc = (x - x.mean(0, keepdim=True)).double()
271
+ s2 = torch.linalg.svdvals(xc) ** 2
272
+ return float((s2.sum() ** 2 / (s2 ** 2).sum()).item())
273
+
274
+
275
+ def hf_token() -> Optional[str]:
276
+ t = os.environ.get("HF_TOKEN")
277
+ if t:
278
+ return t
279
+ try:
280
+ from google.colab import userdata
281
+ return userdata.get("HF_TOKEN")
282
+ except Exception:
283
+ return None
284
+
285
+
286
+ # ══════════════════════════════════════════════════════════════════
287
+ # BACKUP β€” a Colab cull must cost minutes, not the run
288
+ # ══════════════════════════════════════════════════════════════════
289
+
290
+ class Backup:
291
+ def __init__(self, cfg):
292
+ self.cfg, self.api, self.ok, self.last = cfg, None, False, 0.0
293
+ if not cfg.hf_push:
294
+ return
295
+ tok = hf_token()
296
+ if not tok:
297
+ print(" [backup] no HF_TOKEN β€” LOCAL ONLY. A cull will lose the run.")
298
+ return
299
+ try:
300
+ create_repo(cfg.hf_repo, token=tok, exist_ok=True, private=True)
301
+ self.api = HfApi(token=tok)
302
+ self.ok = True
303
+ print(f" [backup] -> {cfg.hf_repo} (private)")
304
+ except Exception as e:
305
+ print(f" [backup] disabled: {type(e).__name__}: {str(e)[:100]}")
306
+
307
+ def push(self, force: bool = False, msg: str = "checkpoint"):
308
+ if not self.ok:
309
+ return
310
+ if not force and (time.time() - self.last) / 60 < self.cfg.push_every_min:
311
+ return
312
+ P = paths(self.cfg)
313
+ try:
314
+ for folder, dest in ((P["ckpt"], "checkpoints"), (P["tb"], "tensorboard"),
315
+ (P["maps"], "maps"), (P["config"], "config")):
316
+ if os.path.isdir(folder) and os.listdir(folder):
317
+ self.api.upload_folder(folder_path=folder, path_in_repo=dest,
318
+ repo_id=self.cfg.hf_repo,
319
+ commit_message=f"{msg} ({dest})")
320
+ self.last = time.time()
321
+ print(f" [backup] pushed ({msg})")
322
+ except Exception as e:
323
+ print(f" [backup] push failed: {type(e).__name__}: {str(e)[:100]}")
324
+
325
+ def pull_latest(self) -> Optional[str]:
326
+ """Recover state.pt after a cull."""
327
+ if not self.ok:
328
+ return None
329
+ try:
330
+ p = hf_hub_download(self.cfg.hf_repo, "checkpoints/state.pt",
331
+ token=hf_token(), local_dir=paths(self.cfg)["root"])
332
+ print(f" [backup] recovered {p}")
333
+ return p
334
+ except Exception:
335
+ return None
336
+
337
+
338
+ # ══════════════════════════════════════════════════════════════════
339
+ # STAGE 0 β€” PARITY GATE
340
+ # ══════════════════════════════════════════════════════════════════
341
+
342
+ def stage0_parity(cfg) -> str:
343
+ """
344
+ Which caption field was embedded, and is row i of <expert>_XXX.pt caption i?
345
+ The manifest names three fields and does not say which was used. If the stored
346
+ vectors came from caption_llava and the student trains on caption_llava_short,
347
+ every target is silently wrong. Re-embed with the real reference model, demand
348
+ cos ~ 1.0. Nothing downstream runs until this passes.
349
+ """
350
+ from transformers import AutoModel, AutoTokenizer
351
+ line("STAGE 0 β€” PARITY GATE (caption field + row alignment)")
352
+ stored = load_expert_chunk(cfg, cfg.ref_expert, cfg.parity_chunk)[: cfg.parity_n].float()
353
+ raw = json.load(open(fetch(cfg, f"captions_{cfg.parity_chunk:03d}.json")))
354
+ if isinstance(raw, dict):
355
+ fields = {k: list(v)[: cfg.parity_n] for k, v in raw.items()
356
+ if k in cfg.caption_field_candidates}
357
+ elif raw and isinstance(raw[0], dict):
358
+ fields = {k: [r[k] for r in raw[: cfg.parity_n]]
359
+ for k in raw[0] if k in cfg.caption_field_candidates}
360
+ else:
361
+ fields = {"(flat)": list(raw[: cfg.parity_n])}
362
+ print(f" stored rows {tuple(stored.shape)} | fields {list(fields)}")
363
+
364
+ tok = AutoTokenizer.from_pretrained(cfg.ref_hf_name)
365
+ mdl = AutoModel.from_pretrained(cfg.ref_hf_name).to(DEVICE).eval()
366
+ best, best_cos = None, -1.0
367
+ for f, texts in fields.items():
368
+ with torch.no_grad():
369
+ inp = tok(list(texts), max_length=512, padding=True, truncation=True,
370
+ return_tensors="pt").to(DEVICE)
371
+ h = mdl(**inp).last_hidden_state
372
+ m = inp.attention_mask.unsqueeze(-1).float()
373
+ pooled = ((h * m).sum(1) / m.sum(1).clamp(min=1)).float().cpu()
374
+ cos = F.cosine_similarity(pooled, stored, dim=-1)
375
+ print(f" {f:22s} cos mean {cos.mean():.6f} min {cos.min():.6f}")
376
+ if cos.mean().item() > best_cos:
377
+ best, best_cos = f, cos.mean().item()
378
+ del mdl; gc.collect(); torch.cuda.empty_cache()
379
+ if best_cos < cfg.parity_min_cos:
380
+ raise RuntimeError(
381
+ f"PARITY GATE FAIL: best field '{best}' only reaches cos {best_cos:.6f} "
382
+ f"(need >= {cfg.parity_min_cos}). Either the field is not among "
383
+ f"{cfg.caption_field_candidates}, row order differs, or the extraction used "
384
+ f"different pooling/truncation. DO NOT SPEND GPU TIME until this resolves.")
385
+ print(f" GATE PASS: field = '{best}' at cos {best_cos:.6f}")
386
+ return best
387
+
388
+
389
+ # ══════════════════════════════════════════════════════════════════
390
+ # STAGE 1 β€” GLOBAL WHITENED PROCRUSTES (out-of-sample reported)
391
+ # ══════════════════════════════════════════════════════════════════
392
+
393
+ def symmetric_inv_sqrt(cov: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
394
+ ev, evec = torch.linalg.eigh(cov.double())
395
+ return (evec @ torch.diag(torch.clamp(ev, min=eps).rsqrt()) @ evec.T).float()
396
+
397
+
398
+ def fit_map(S: torch.Tensor, T: torch.Tensor) -> Dict[str, torch.Tensor]:
399
+ N = S.shape[0]
400
+ s_mean, t_mean = S.mean(0, keepdim=True), T.mean(0, keepdim=True)
401
+ Sc, Tc = S - s_mean, T - t_mean
402
+ s_w = symmetric_inv_sqrt((Sc.T @ Sc) / max(N - 1, 1))
403
+ t_w = symmetric_inv_sqrt((Tc.T @ Tc) / max(N - 1, 1))
404
+ U, _, Vt = torch.linalg.svd(
405
+ (F.normalize(Tc @ t_w, dim=-1).T @ F.normalize(Sc @ s_w, dim=-1)).double(),
406
+ full_matrices=False)
407
+ return {"rotation": (U @ Vt).float(), "source_mean": s_mean.squeeze(0),
408
+ "target_mean": t_mean.squeeze(0), "source_whitener": s_w,
409
+ "target_whitener": t_w, "target_unwhitener": torch.linalg.pinv(t_w)}
410
+
411
+
412
+ def apply_map(emb: torch.Tensor, a) -> torch.Tensor:
413
+ x = (emb.float() - a["source_mean"]) @ a["source_whitener"]
414
+ return (x @ a["rotation"].T) @ a["target_unwhitener"]
415
+
416
+
417
+ def score_map(S, T, a) -> Dict[str, float]:
418
+ Sw = F.normalize((S - a["source_mean"]) @ a["source_whitener"], dim=-1)
419
+ Tw = F.normalize((T - a["target_mean"]) @ a["target_whitener"], dim=-1)
420
+ cos = F.cosine_similarity(Sw @ a["rotation"].T, Tw, dim=-1).mean().item()
421
+ n = min(2000, S.shape[0])
422
+ sim = F.normalize(apply_map(S[:n], a), dim=-1) @ F.normalize(T[:n], dim=-1).T
423
+ return {"cos": cos, "r1": (sim.argmax(1) == torch.arange(n)).float().mean().item(),
424
+ "n": int(S.shape[0]), "chance": 1.0 / n}
425
+
426
+
427
+ def stage1_fit(cfg, bk: "Backup"):
428
+ line("STAGE 1 β€” GLOBAL ALIGNMENT (stratified fit, OUT-OF-SAMPLE report)")
429
+ P = paths(cfg)
430
+ mp = f"{P['maps']}/alignment_maps.pt"
431
+ if os.path.exists(mp):
432
+ print(" maps exist, loading"); return torch.load(mp, weights_only=False)
433
+
434
+ g = torch.Generator().manual_seed(cfg.fit_seed)
435
+ fit = {e: [] for e in cfg.experts}
436
+ for c in cfg.fit_chunks:
437
+ idx = None
438
+ for e in cfg.experts:
439
+ X = load_expert_chunk(cfg, e, c)
440
+ if idx is None:
441
+ idx = torch.randperm(X.shape[0], generator=g)[: cfg.fit_rows_per_chunk]
442
+ fit[e].append(X[idx].float()); del X; gc.collect()
443
+ drop_shard(cfg, f"{e}_{c:03d}.pt")
444
+ print(f" fit chunk {c:03d}: {len(idx)} random rows")
445
+ fit = {e: torch.cat(v) for e, v in fit.items()}
446
+ N = fit[cfg.ref_expert].shape[0]
447
+ print(f" fit set {N} rows, d=768 -> N/d = {N/768:.1f}")
448
+
449
+ hold = {e: [] for e in cfg.experts}
450
+ for c in cfg.holdout_chunks:
451
+ for e in cfg.experts:
452
+ X = load_expert_chunk(cfg, e, c)
453
+ hold[e].append(X[: cfg.fit_rows_per_chunk].float()); del X; gc.collect()
454
+ hold = {e: torch.cat(v) for e, v in hold.items()}
455
+
456
+ maps, report, T = {}, {}, fit[cfg.ref_expert]
457
+ for e in cfg.experts:
458
+ a = fit_map(fit[e], T)
459
+ ins, oos = score_map(fit[e], T, a), score_map(hold[e], hold[cfg.ref_expert], a)
460
+ maps[e], report[e] = a, {"in_sample": ins, "out_of_sample": oos}
461
+ tag = " (ref: must read ~1.0)" if e == cfg.ref_expert else ""
462
+ print(f" {e:9s} cos in {ins['cos']:.4f} / OUT {oos['cos']:.4f} "
463
+ f"R@1 in {ins['r1']:.4f} / OUT {oos['r1']:.4f} "
464
+ f"(chance {oos['chance']:.5f}){tag}")
465
+ print(" READ THE 'OUT' COLUMN. A 768x768 rotation is 294,528 free parameters;")
466
+ print(" at low N/d the in-sample cosine reproduces strong numbers from nothing.")
467
+ torch.save(maps, mp)
468
+ json.dump(report, open(f"{P['maps']}/fit_report.json", "w"), indent=2)
469
+ bk.push(force=True, msg="stage1 alignment maps")
470
+ return maps
471
+
472
+
473
+ # ══════════════════════════════════════════════════════════════════
474
+ # STAGE 2 β€” CONSENSUS TARGETS
475
+ # ══════════════════════════════════════════════════════════════════
476
+
477
+ def stage2_targets(cfg, maps, bk: "Backup") -> List[int]:
478
+ line("STAGE 2 β€” CONSENSUS TARGETS (fp16, per chunk, resumable)")
479
+ P = paths(cfg)
480
+ lp = f"{P['targets']}/ledger.json"
481
+ ledger = json.load(open(lp)) if os.path.exists(lp) else {}
482
+ want = sorted(set(usable_chunks(cfg)) | set(cfg.holdout_chunks))
483
+ print(f" {len(want)} chunks with all {len(cfg.experts)} experts | "
484
+ f"streaming {len(want)*len(cfg.experts)*1.536:.0f} GB through "
485
+ f"{free_gb(P['root']):.0f} GB of free disk")
486
+ tapi = None
487
+ if cfg.push_targets and bk.ok:
488
+ try:
489
+ create_repo(cfg.targets_repo, token=hf_token(), exist_ok=True,
490
+ private=True, repo_type="dataset")
491
+ tapi = HfApi(token=hf_token())
492
+ print(f" targets -> {cfg.targets_repo} (dataset, private)")
493
+ except Exception as e:
494
+ print(f" target push disabled: {type(e).__name__}: {str(e)[:80]}")
495
+ for c in want:
496
+ k, out_p = f"{c:03d}", f"{P['targets']}/consensus_{c:03d}.pt"
497
+ if ledger.get(k) and os.path.exists(out_p):
498
+ continue
499
+ if free_gb(P["root"]) < cfg.disk_floor_gb:
500
+ raise RuntimeError(f"DISK FLOOR: {free_gb(P['root']):.1f} GB free at chunk {k}. "
501
+ f"Push and drop earlier consensus files, then resume.")
502
+ acc, n = None, None
503
+ for e in cfg.experts:
504
+ X = load_expert_chunk(cfg, e, c).float()
505
+ if n is None:
506
+ n = X.shape[0]
507
+ elif X.shape[0] != n:
508
+ raise RuntimeError(f"chunk {k}: {e} has {X.shape[0]} rows, expected {n}")
509
+ A = apply_map(X, maps[e])
510
+ acc = A if acc is None else acc + A
511
+ del X, A; gc.collect()
512
+ drop_shard(cfg, f"{e}_{c:03d}.pt")
513
+ purge_hf_cache(cfg)
514
+ cons = F.normalize(acc / len(cfg.experts), dim=-1).half()
515
+ torch.save(cons, out_p)
516
+ er = effective_rank(cons[:4000].float())
517
+ ledger[k] = {"rows": int(cons.shape[0]), "target_erank": er, "ts": time.time()}
518
+ json.dump(ledger, open(lp, "w"), indent=2)
519
+ if tapi is not None:
520
+ try:
521
+ tapi.upload_file(path_or_fileobj=out_p,
522
+ path_in_repo=f"consensus_{k}.pt",
523
+ repo_id=cfg.targets_repo, repo_type="dataset",
524
+ commit_message=f"consensus chunk {k}")
525
+ except Exception as ex:
526
+ print(f" target push failed for {k}: {str(ex)[:70]}")
527
+ print(f" chunk {k}: {cons.shape[0]} targets | TARGET erank {er:.1f}/768 | "
528
+ f"disk free {free_gb(P['root']):.0f} GB")
529
+ del acc, cons; gc.collect()
530
+ eranks = [v["target_erank"] for v in ledger.values() if "target_erank" in v]
531
+ if eranks:
532
+ print(f" consensus target erank: mean {np.mean(eranks):.1f} "
533
+ f"min {min(eranks):.1f} max {max(eranks):.1f} of 768")
534
+ print(" (v1's STUDENT read 23.6 β€” compare against this to tell 'student")
535
+ print(" collapsed' from 'student faithfully matched a low-rank target')")
536
+ bk.push(force=True, msg="stage2 target ledger")
537
+ return want
538
+
539
+
540
+ # ══════════════════════════════════════════════════════════════════
541
+ # STUDENT
542
+ # ══════════════════════════════════════════════════════════════════
543
+
544
+ class CaptionEncoder(nn.Module):
545
+ """Standalone caption encoder. No experts at inference. No bank."""
546
+
547
+ def __init__(self, vocab_size=30522, max_len=8192, d_model=512, n_heads=8,
548
+ n_layers=12, d_ff=2048, output_dim=768, dropout=0.1,
549
+ pad_token_id=0, pooling="mean"):
550
+ super().__init__()
551
+ self.pad_token_id, self.pooling = pad_token_id, pooling
552
+ self.token_emb = nn.Embedding(vocab_size, d_model, padding_idx=pad_token_id)
553
+ self.pos_emb = nn.Embedding(max_len, d_model)
554
+ self.emb_norm = nn.LayerNorm(d_model)
555
+ self.emb_drop = nn.Dropout(dropout)
556
+ layer = nn.TransformerEncoderLayer(
557
+ d_model=d_model, nhead=n_heads, dim_feedforward=d_ff, dropout=dropout,
558
+ activation="gelu", batch_first=True, norm_first=True)
559
+ self.encoder = nn.TransformerEncoder(layer, num_layers=n_layers,
560
+ enable_nested_tensor=False)
561
+ self.output_proj = nn.Sequential(
562
+ nn.Linear(d_model, d_model), nn.GELU(), nn.LayerNorm(d_model),
563
+ nn.Linear(d_model, output_dim))
564
+
565
+ def forward(self, input_ids, attention_mask=None):
566
+ L = input_ids.shape[1]
567
+ pos = torch.arange(L, device=input_ids.device).unsqueeze(0)
568
+ x = self.emb_drop(self.emb_norm(self.token_emb(input_ids) + self.pos_emb(pos)))
569
+ kpm = (~attention_mask.bool()) if attention_mask is not None \
570
+ else (input_ids == self.pad_token_id)
571
+ x = self.encoder(x, src_key_padding_mask=kpm)
572
+ if self.pooling == "cls":
573
+ pooled = x[:, 0]
574
+ else:
575
+ m = (attention_mask.unsqueeze(-1).float() if attention_mask is not None
576
+ else (~kpm).unsqueeze(-1).float())
577
+ pooled = (x * m).sum(1) / m.sum(1).clamp(min=1)
578
+ return F.normalize(self.output_proj(pooled), dim=-1)
579
+
580
+
581
+ # ══════════════════════════════════════════════════════════════════
582
+ # LOSS / GAUGES
583
+ # ══════════════════════════════════════════════════════════════════
584
+
585
+ def infonce(a, b, temperature=0.07):
586
+ logits = (a @ b.T) / temperature
587
+ lab = torch.arange(logits.shape[0], device=logits.device)
588
+ loss = (F.cross_entropy(logits, lab) + F.cross_entropy(logits.T, lab)) / 2
589
+ with torch.no_grad():
590
+ acc = (logits.argmax(-1) == lab).float().mean().item()
591
+ return loss, acc
592
+
593
+
594
+ def cayley_menger_vol2(pts):
595
+ pts = pts.float()
596
+ d = pts.unsqueeze(-2) - pts.unsqueeze(-3)
597
+ d2 = (d * d).sum(-1)
598
+ B, V, _ = d2.shape
599
+ cm = torch.zeros(B, V + 1, V + 1, device=d2.device, dtype=torch.float32)
600
+ cm[:, 0, 1:] = 1; cm[:, 1:, 0] = 1; cm[:, 1:, 1:] = d2
601
+ f = math.factorial(V - 1)
602
+ return ((-1.0) ** V) / ((2.0 ** (V - 1)) * f * f) * torch.linalg.det(cm)
603
+
604
+
605
+ def cv_loss(emb, target=0.084, n_samples=16):
606
+ B = emb.shape[0]
607
+ if B < 5:
608
+ return torch.zeros((), device=emb.device)
609
+ s = torch.stack([torch.sqrt(F.relu(cayley_menger_vol2(
610
+ emb[torch.randperm(B, device=emb.device)[:5]].unsqueeze(0))[0]) + 1e-12)
611
+ for _ in range(n_samples)])
612
+ return (s.std() / (s.mean() + 1e-8) - target).abs()
613
+
614
+
615
+ @torch.no_grad()
616
+ def cv_metric(emb, n=200):
617
+ v = [float(torch.sqrt(F.relu(cayley_menger_vol2(
618
+ emb[torch.randperm(emb.shape[0], device=emb.device)[:5]].unsqueeze(0))[0])
619
+ + 1e-12).item()) for _ in range(n)]
620
+ a = np.array([x for x in v if x > 0])
621
+ return float(a.std() / (a.mean() + 1e-8)) if len(a) >= 10 else 0.0
622
+
623
+
624
+ @torch.no_grad()
625
+ def frame_fit_gauge(E: torch.Tensor, T: torch.Tensor, n_pairs: int = 2500) -> Dict[str, float]:
626
+ """
627
+ Standing rider: judge relational objectives with a frame fit or they read as false
628
+ floors. MSE anchors the frame here and the consensus aligns to a REFERENCE MEMBER,
629
+ so a rotation should buy ~nothing. If it buys a lot, the frame is NOT pinned and
630
+ this model needs a shipped rotation after all. Held-out split, fp64.
631
+ """
632
+ N = E.shape[0]
633
+ k = min(n_pairs, N // 2)
634
+ if k < 64:
635
+ return {"skipped": True}
636
+ perm = torch.randperm(N, generator=torch.Generator().manual_seed(0))
637
+ fit_i, hold_i = perm[:k], perm[k:]
638
+ U, _, Vt = torch.linalg.svd(E[fit_i].double().T @ T[fit_i].double(), full_matrices=False)
639
+ Er = F.normalize((E.double() @ (U @ Vt)).float(), dim=-1)
640
+ m = min(2000, len(hold_i))
641
+ hi = hold_i[:m]
642
+ sim = Er[hi] @ T[hi].T
643
+ return {"r1_after_rotation": (sim.argmax(1) == torch.arange(m)).float().mean().item(),
644
+ "cos_after_rotation": F.cosine_similarity(Er[hi], T[hi], dim=-1).mean().item(),
645
+ "n_heldout": int(m)}
646
+
647
+
648
+ # ══════════════════════════════════════════════════════════════════
649
+ # DATA
650
+ # ══════════════════════════════════════════════════════════════════
651
+
652
+ class RamStore:
653
+ """
654
+ Everything resident in system RAM: ragged uint16 tokens + fp16 targets.
655
+
656
+ On the Pro+ box this is 48.8 GB of 176.9 β€” so the training loop does ZERO disk
657
+ I/O and needs no DataLoader workers. Ragged storage (flat token buffer + offsets)
658
+ keeps dynamic padding available at ~5.6 GB instead of the 14 GB a fixed 256-token
659
+ matrix would cost, and captions average ~100 tokens against a 256 ceiling.
660
+ """
661
+
662
+ def __init__(self, cfg, chunks: List[int], tokenizer, tag=""):
663
+ self.cfg, self.tok = cfg, tokenizer
664
+ self.pad = tokenizer.pad_token_id
665
+ flat, offs, tgts, total = [], [0], [], 0
666
+ for c in chunks:
667
+ caps = load_captions_chunk(cfg, c)
668
+ t = torch.load(f"{paths(cfg)['targets']}/consensus_{c:03d}.pt",
669
+ weights_only=True, map_location="cpu")
670
+ n = min(len(caps), t.shape[0])
671
+ caps, t = caps[:n], t[:n]
672
+ for i in range(0, n, 20000):
673
+ enc = tokenizer(caps[i:i + 20000], max_length=cfg.max_tokens,
674
+ truncation=True, padding=False)["input_ids"]
675
+ for ids in enc:
676
+ flat.append(np.asarray(ids, dtype=np.uint16))
677
+ total += len(ids)
678
+ offs.append(total)
679
+ tgts.append(t)
680
+ print(f" chunk {c:03d}: {n:,} rows | flat tokens {total/1e6:.1f}M")
681
+ del caps, t; gc.collect()
682
+ self.flat = np.concatenate(flat) if flat else np.zeros(0, np.uint16)
683
+ del flat; gc.collect()
684
+ self.offs = np.asarray(offs, dtype=np.int64)
685
+ self.tgt = torch.cat(tgts)
686
+ del tgts; gc.collect()
687
+ self.n = len(self.offs) - 1
688
+ gb = (self.flat.nbytes + self.offs.nbytes + self.tgt.numel() * 2) / 1e9
689
+ mean_len = total / max(self.n, 1)
690
+ print(f" RamStore{tag}: {self.n:,} rows | {gb:.1f} GB RAM | "
691
+ f"mean {mean_len:.0f} tokens (ceiling {cfg.max_tokens})")
692
+
693
+ def __len__(self):
694
+ return self.n
695
+
696
+ def batch(self, idx: np.ndarray):
697
+ """Gather a batch with DYNAMIC padding to the batch max."""
698
+ seqs = [self.flat[self.offs[i]:self.offs[i + 1]] for i in idx]
699
+ L = max(len(s) for s in seqs)
700
+ ids = np.full((len(seqs), L), self.pad, dtype=np.int64)
701
+ am = np.zeros((len(seqs), L), dtype=np.int64)
702
+ for r, s in enumerate(seqs):
703
+ ids[r, :len(s)] = s
704
+ am[r, :len(s)] = 1
705
+ return (torch.from_numpy(ids), torch.from_numpy(am),
706
+ self.tgt[torch.from_numpy(idx)])
707
+
708
+
709
+ class ChunkPairs(torch.utils.data.Dataset):
710
+ """Disk-streaming fallback when ram_resident=False."""
711
+ def __init__(self, cfg, chunk, tokenizer):
712
+ self.caps = load_captions_chunk(cfg, chunk)
713
+ self.tgt = torch.load(f"{paths(cfg)['targets']}/consensus_{chunk:03d}.pt",
714
+ weights_only=True, map_location="cpu")
715
+ n = min(len(self.caps), self.tgt.shape[0])
716
+ self.caps, self.tgt = self.caps[:n], self.tgt[:n]
717
+ self.tok, self.max_tokens = tokenizer, cfg.max_tokens
718
+
719
+ def __len__(self):
720
+ return len(self.caps)
721
+
722
+ def __getitem__(self, i):
723
+ return self.caps[i], self.tgt[i]
724
+
725
+ def collate(self, batch):
726
+ texts, tg = zip(*batch)
727
+ enc = self.tok(list(texts), max_length=self.max_tokens, padding=True,
728
+ truncation=True, return_tensors="pt") # DYNAMIC
729
+ return enc["input_ids"], enc["attention_mask"], torch.stack(tg)
730
+
731
+
732
+ @torch.no_grad()
733
+ def evaluate(student, source, cap=5000, batch=512) -> Dict[str, float]:
734
+ student.eval()
735
+ E, T = [], []
736
+ if isinstance(source, RamStore):
737
+ for i in range(0, min(cap, len(source)), batch):
738
+ ids, am, tg = source.batch(np.arange(i, min(i + batch, len(source))))
739
+ E.append(student(ids.to(DEVICE), am.to(DEVICE)).float().cpu())
740
+ T.append(tg.float())
741
+ else:
742
+ for ids, am, tg in source:
743
+ E.append(student(ids.to(DEVICE), am.to(DEVICE)).float().cpu())
744
+ T.append(tg.float())
745
+ if sum(x.shape[0] for x in E) >= cap:
746
+ break
747
+ E, T = torch.cat(E), F.normalize(torch.cat(T), dim=-1)
748
+ n = min(2000, E.shape[0])
749
+ sim = E[:n] @ T[:n].T
750
+ ss = E[:n] @ E[:n].T
751
+ ss.fill_diagonal_(0)
752
+ out = {"mimicry_r1": (sim.argmax(1) == torch.arange(n)).float().mean().item(),
753
+ "cos_to_target": F.cosine_similarity(E, T, dim=-1).mean().item(),
754
+ "self_cos": ss.mean().item(),
755
+ "erank": effective_rank(E),
756
+ "cv": cv_metric(E[:2000].to(DEVICE)),
757
+ "n": int(E.shape[0])}
758
+ out.update({f"frame_{k}": v for k, v in frame_fit_gauge(E, T).items()})
759
+ student.train()
760
+ return out
761
+
762
+
763
+ # ══════════════════════════════════════════════════════════════════
764
+ # STAGE 3 β€” TRAIN (cull-proof)
765
+ # ══════════════════════════════════════════════════════════════════
766
+
767
+ def save_state(cfg, path, student, opt, sched, scaler, step, epoch, chunk_i, order, best):
768
+ torch.save({"model": student.state_dict(), "opt": opt.state_dict(),
769
+ "sched": sched.state_dict(), "scaler": scaler.state_dict(),
770
+ "step": step, "epoch": epoch, "chunk_i": chunk_i, "order": order,
771
+ "best": best, "config": asdict(cfg),
772
+ "rng": {"torch": torch.get_rng_state(), "np": np.random.get_state(),
773
+ "py": random.getstate()}}, path)
774
+
775
+
776
+ def stage3_train(cfg, chunks: List[int], bk: "Backup"):
777
+ from transformers import AutoTokenizer
778
+ line("STAGE 3 β€” TRAIN")
779
+ P = paths(cfg)
780
+ torch.manual_seed(cfg.seed); np.random.seed(cfg.seed); random.seed(cfg.seed)
781
+ tok = AutoTokenizer.from_pretrained(cfg.ref_hf_name)
782
+ json.dump(asdict(cfg), open(f"{P['config']}/config.json", "w"), indent=2, default=str)
783
+
784
+ student = CaptionEncoder(
785
+ vocab_size=tok.vocab_size, max_len=cfg.max_len, d_model=cfg.d_model,
786
+ n_heads=cfg.n_heads, n_layers=cfg.n_layers, d_ff=cfg.d_ff,
787
+ output_dim=cfg.output_dim, dropout=cfg.dropout,
788
+ pad_token_id=tok.pad_token_id, pooling=cfg.pooling).to(DEVICE)
789
+ n_par = sum(p.numel() for p in student.parameters())
790
+ train_chunks = [c for c in chunks if c not in cfg.holdout_chunks]
791
+ rows = len(train_chunks) * src0(cfg)["chunk_rows"]
792
+ spe = rows // cfg.batch_size
793
+ total = spe * cfg.epochs
794
+ print(f" {cfg.run_name}: {n_par:,} params ({n_par/109_482_240:.2f}x bert-base)")
795
+ print(f" {cfg.n_layers}L {cfg.d_model}d {cfg.n_heads}h ff{cfg.d_ff} pool={cfg.pooling}")
796
+ print(f" {len(train_chunks)} chunks β‰ˆ {rows:,} rows | {spe:,} steps/ep x "
797
+ f"{cfg.epochs} = {total:,} steps @ batch {cfg.batch_size}")
798
+ print(f" loss = {cfg.nce_weight}*InfoNCE(T={cfg.nce_temperature}) + "
799
+ f"{cfg.mse_weight}*MSE + {cfg.cv_weight}*CV [champion consensus_nce_mse]")
800
+
801
+ opt = torch.optim.Adam(student.parameters(), lr=cfg.lr) # pure Adam, no wd
802
+ sched = torch.optim.lr_scheduler.SequentialLR(
803
+ opt, [torch.optim.lr_scheduler.LinearLR(opt, 0.01, 1.0, cfg.warmup_steps),
804
+ torch.optim.lr_scheduler.CosineAnnealingLR(
805
+ opt, T_max=max(total - cfg.warmup_steps, 1), eta_min=cfg.min_lr)],
806
+ milestones=[cfg.warmup_steps])
807
+ scaler = torch.amp.GradScaler(enabled=cfg.amp and DEVICE == "cuda")
808
+ tb = SummaryWriter(log_dir=f"{P['tb']}/{cfg.run_name}")
809
+ tb.add_text("config", f"```json\n{json.dumps(asdict(cfg), indent=2, default=str)}\n```")
810
+ if os.path.exists(f"{P['maps']}/fit_report.json"):
811
+ tb.add_text("alignment/fit_report",
812
+ f"```json\n{open(f'{P['maps']}/fit_report.json').read()}\n```")
813
+
814
+ step, ep0, chunk_i0, best = 0, 0, 0, -1.0
815
+ order = None
816
+ sp = f"{P['ckpt']}/state.pt"
817
+ if cfg.resume:
818
+ if not os.path.exists(sp):
819
+ bk.pull_latest()
820
+ alt = f"{P['root']}/checkpoints/state.pt"
821
+ if os.path.exists(alt) and alt != sp:
822
+ shutil.copy(alt, sp)
823
+ if os.path.exists(sp):
824
+ st = torch.load(sp, weights_only=False, map_location=DEVICE)
825
+ student.load_state_dict(st["model"]); opt.load_state_dict(st["opt"])
826
+ sched.load_state_dict(st["sched"]); scaler.load_state_dict(st["scaler"])
827
+ step, ep0, chunk_i0, best = st["step"], st["epoch"], st["chunk_i"], st["best"]
828
+ order = st.get("order")
829
+ try:
830
+ torch.set_rng_state(st["rng"]["torch"].cpu())
831
+ np.random.set_state(st["rng"]["np"]); random.setstate(st["rng"]["py"])
832
+ except Exception:
833
+ pass
834
+ print(f" RESUMED at step {step:,} epoch {ep0+1} chunk_i {chunk_i0}")
835
+
836
+ print(" building val store...")
837
+ if cfg.ram_resident:
838
+ val_src = RamStore(cfg, [cfg.holdout_chunks[-1]], tok, tag=" [val]")
839
+ else:
840
+ vds = ChunkPairs(cfg, cfg.holdout_chunks[-1], tok)
841
+ val_src = torch.utils.data.DataLoader(
842
+ vds, batch_size=cfg.batch_size, shuffle=False,
843
+ num_workers=cfg.num_workers, collate_fn=vds.collate)
844
+
845
+ if cfg.ram_resident:
846
+ print(" building train store (one pass, then zero disk I/O)...")
847
+ train_src = RamStore(cfg, train_chunks, tok, tag=" [train]")
848
+ N = len(train_src)
849
+ spe = N // cfg.batch_size
850
+ total = spe * cfg.epochs
851
+ print(f" {N:,} rows resident | {spe:,} steps/ep x {cfg.epochs} = {total:,} steps")
852
+
853
+ t0 = last_ck = time.time()
854
+ for ep in range(ep0, cfg.epochs):
855
+ if cfg.ram_resident:
856
+ # index permutation over the resident store; chunk_i doubles as batch index
857
+ perm = np.random.default_rng(cfg.seed + ep).permutation(N)
858
+ for ci in range(chunk_i0 if ep == ep0 else 0, spe):
859
+ ids, am, tg = train_src.batch(
860
+ perm[ci * cfg.batch_size:(ci + 1) * cfg.batch_size])
861
+ ids = ids.to(DEVICE, non_blocking=True)
862
+ am = am.to(DEVICE, non_blocking=True)
863
+ tgt = F.normalize(tg.to(DEVICE, non_blocking=True).float(), dim=-1)
864
+ with torch.amp.autocast("cuda", enabled=cfg.amp and DEVICE == "cuda"):
865
+ emb = student(ids, am)
866
+ emb = emb.float()
867
+ l_nce, acc = infonce(emb, tgt, cfg.nce_temperature)
868
+ l_mse = F.mse_loss(emb, tgt)
869
+ loss = cfg.nce_weight * l_nce + cfg.mse_weight * l_mse
870
+ l_cv = torch.zeros((), device=emb.device)
871
+ if cfg.cv_weight > 0:
872
+ l_cv = cv_loss(emb, cfg.cv_target)
873
+ loss = loss + cfg.cv_weight * l_cv
874
+ scaler.scale(loss).backward()
875
+ scaler.unscale_(opt)
876
+ gn = torch.nn.utils.clip_grad_norm_(student.parameters(), cfg.grad_clip)
877
+ scaler.step(opt); scaler.update()
878
+ opt.zero_grad(set_to_none=True); sched.step()
879
+ step += 1
880
+
881
+ if step % cfg.log_every == 0:
882
+ lr = opt.param_groups[0]["lr"]
883
+ tb.add_scalar("train/loss", loss.item(), step)
884
+ tb.add_scalar("train/nce", l_nce.item(), step)
885
+ tb.add_scalar("train/mse", l_mse.item(), step)
886
+ tb.add_scalar("train/cv", float(l_cv), step)
887
+ tb.add_scalar("train/batch_acc", acc, step)
888
+ tb.add_scalar("train/lr", lr, step)
889
+ tb.add_scalar("train/grad_norm", float(gn), step)
890
+ tb.add_scalar("train/tokens_per_seq", ids.shape[1], step)
891
+ print(f" e{ep+1} {step:>7,}/{total:,} loss {loss.item():.4f} "
892
+ f"nce {l_nce.item():.4f} mse {l_mse.item():.5f} acc {acc:.3f} "
893
+ f"lr {lr:.2e} L{ids.shape[1]} {(time.time()-t0)/60:.0f}m")
894
+
895
+ if step % cfg.eval_every == 0:
896
+ m = evaluate(student, val_src)
897
+ for k, v in m.items():
898
+ if isinstance(v, (int, float)):
899
+ tb.add_scalar(f"val/{k}", v, step)
900
+ for nm, p in student.named_parameters():
901
+ if p.grad is not None and ("output_proj" in nm or "token_emb" in nm):
902
+ tb.add_histogram(f"grad/{nm}", p.grad, step)
903
+ tb.add_histogram(f"weight/{nm}", p, step)
904
+ print(f" VAL r1 {m['mimicry_r1']:.4f} cos {m['cos_to_target']:.4f} "
905
+ f"self_cos {m['self_cos']:+.4f} erank {m['erank']:.1f} "
906
+ f"cv {m['cv']:.4f} | frame r1 "
907
+ f"{m.get('frame_r1_after_rotation', float('nan')):.4f}")
908
+ if m["cos_to_target"] > best:
909
+ best = m["cos_to_target"]
910
+ save_state(cfg, f"{P['ckpt']}/best_state.pt", student, opt,
911
+ sched, scaler, step, ep, ci, order, best)
912
+ torch.save(student.state_dict(), f"{P['ckpt']}/best_model.pt")
913
+
914
+ if (time.time() - last_ck) / 60 >= cfg.ckpt_every_min:
915
+ save_state(cfg, sp, student, opt, sched, scaler, step, ep, ci, order, best)
916
+ torch.save(student.state_dict(), f"{P['ckpt']}/model_s{step}.pt")
917
+ ck = sorted([f for f in os.listdir(P["ckpt"]) if f.startswith("model_s")],
918
+ key=lambda f: int(f.split("_s")[1].split(".")[0]))
919
+ for old in ck[:-cfg.keep_local_ckpts]:
920
+ os.remove(os.path.join(P["ckpt"], old))
921
+ tb.flush(); bk.push(msg=f"step {step}")
922
+ last_ck = time.time()
923
+ else:
924
+ if order is None or ep != ep0:
925
+ order = train_chunks[:]; random.shuffle(order)
926
+ for ci in range(chunk_i0 if ep == ep0 else 0, len(order)):
927
+ c = order[ci]
928
+ ds = ChunkPairs(cfg, c, tok)
929
+ dl = torch.utils.data.DataLoader(
930
+ ds, batch_size=cfg.batch_size, shuffle=True, drop_last=True,
931
+ num_workers=cfg.num_workers, collate_fn=ds.collate,
932
+ pin_memory=(DEVICE == "cuda"))
933
+ for ids, am, tg in dl:
934
+ ids = ids.to(DEVICE, non_blocking=True)
935
+ am = am.to(DEVICE, non_blocking=True)
936
+ tgt = F.normalize(tg.to(DEVICE, non_blocking=True).float(), dim=-1)
937
+ with torch.amp.autocast("cuda", enabled=cfg.amp and DEVICE == "cuda"):
938
+ emb = student(ids, am)
939
+ emb = emb.float()
940
+ l_nce, acc = infonce(emb, tgt, cfg.nce_temperature)
941
+ l_mse = F.mse_loss(emb, tgt)
942
+ loss = cfg.nce_weight * l_nce + cfg.mse_weight * l_mse
943
+ l_cv = torch.zeros((), device=emb.device)
944
+ if cfg.cv_weight > 0:
945
+ l_cv = cv_loss(emb, cfg.cv_target)
946
+ loss = loss + cfg.cv_weight * l_cv
947
+ scaler.scale(loss).backward()
948
+ scaler.unscale_(opt)
949
+ gn = torch.nn.utils.clip_grad_norm_(student.parameters(), cfg.grad_clip)
950
+ scaler.step(opt); scaler.update()
951
+ opt.zero_grad(set_to_none=True); sched.step()
952
+ step += 1
953
+
954
+ if step % cfg.log_every == 0:
955
+ lr = opt.param_groups[0]["lr"]
956
+ tb.add_scalar("train/loss", loss.item(), step)
957
+ tb.add_scalar("train/nce", l_nce.item(), step)
958
+ tb.add_scalar("train/mse", l_mse.item(), step)
959
+ tb.add_scalar("train/cv", float(l_cv), step)
960
+ tb.add_scalar("train/batch_acc", acc, step)
961
+ tb.add_scalar("train/lr", lr, step)
962
+ tb.add_scalar("train/grad_norm", float(gn), step)
963
+ tb.add_scalar("train/tokens_per_seq", ids.shape[1], step)
964
+ print(f" e{ep+1} {step:>7,}/{total:,} loss {loss.item():.4f} "
965
+ f"nce {l_nce.item():.4f} mse {l_mse.item():.5f} acc {acc:.3f} "
966
+ f"lr {lr:.2e} L{ids.shape[1]} {(time.time()-t0)/60:.0f}m")
967
+
968
+ if step % cfg.eval_every == 0:
969
+ m = evaluate(student, val_src)
970
+ for k, v in m.items():
971
+ if isinstance(v, (int, float)):
972
+ tb.add_scalar(f"val/{k}", v, step)
973
+ for nm, p in student.named_parameters():
974
+ if p.grad is not None and ("output_proj" in nm or "token_emb" in nm):
975
+ tb.add_histogram(f"grad/{nm}", p.grad, step)
976
+ tb.add_histogram(f"weight/{nm}", p, step)
977
+ print(f" VAL r1 {m['mimicry_r1']:.4f} cos {m['cos_to_target']:.4f} "
978
+ f"self_cos {m['self_cos']:+.4f} erank {m['erank']:.1f} "
979
+ f"cv {m['cv']:.4f} | frame r1 "
980
+ f"{m.get('frame_r1_after_rotation', float('nan')):.4f}")
981
+ if m["cos_to_target"] > best:
982
+ best = m["cos_to_target"]
983
+ save_state(cfg, f"{P['ckpt']}/best_state.pt", student, opt,
984
+ sched, scaler, step, ep, ci, order, best)
985
+ torch.save(student.state_dict(), f"{P['ckpt']}/best_model.pt")
986
+
987
+ if (time.time() - last_ck) / 60 >= cfg.ckpt_every_min:
988
+ save_state(cfg, sp, student, opt, sched, scaler, step, ep, ci, order, best)
989
+ torch.save(student.state_dict(), f"{P['ckpt']}/model_s{step}.pt")
990
+ ck = sorted([f for f in os.listdir(P["ckpt"]) if f.startswith("model_s")],
991
+ key=lambda f: int(f.split("_s")[1].split(".")[0]))
992
+ for old in ck[:-cfg.keep_local_ckpts]:
993
+ os.remove(os.path.join(P["ckpt"], old))
994
+ tb.flush(); bk.push(msg=f"step {step}")
995
+ last_ck = time.time()
996
+ del ds, dl; gc.collect()
997
+ chunk_i0 = 0
998
+
999
+ save_state(cfg, sp, student, opt, sched, scaler, step, cfg.epochs, 0, order, best)
1000
+ torch.save(student.state_dict(), f"{P['ckpt']}/final_model.pt")
1001
+ tok.save_pretrained(f"{P['ckpt']}/tokenizer")
1002
+ m = evaluate(student, val_src)
1003
+ line("FINAL")
1004
+ print(f" mimicry R@1 (student->consensus, NOT capability): {m['mimicry_r1']:.4f}")
1005
+ print(f" cos to target : {m['cos_to_target']:.4f}")
1006
+ print(f" self_cos : {m['self_cos']:+.4f} <- isotropy; teachers .81-.98")
1007
+ print(f" effective rank: {m['erank']:.1f}/{cfg.output_dim}")
1008
+ print(f" CV : {m['cv']:.4f}")
1009
+ print(f" frame-fit R@1 : {m.get('frame_r1_after_rotation', float('nan')):.4f} "
1010
+ f"(should be ~mimicry: reference-member alignment pins the frame)")
1011
+ print(" CAPABILITY is decided by STS-B / SICK vs the five teachers, not here.")
1012
+ json.dump({"config": asdict(cfg), "final": m}, open(f"{P['ckpt']}/metrics.json", "w"),
1013
+ indent=2, default=str)
1014
+ tb.flush(); tb.close(); bk.push(force=True, msg="final")
1015
+ return student
1016
+
1017
+
1018
+ # ══════════════════════════════════════════════════════════════════
1019
+ # RUN
1020
+ # ══════════════════════════════════════════════════════════════════
1021
+
1022
+ def run(cfg: BaseConfig = CFG):
1023
+ print("=" * 78)
1024
+ print(f"{cfg.run_name.upper()} β€” CONSENSUS DISTILLATION, CC12M SCALE")
1025
+ print("=" * 78)
1026
+ paths(cfg)
1027
+ print(f"device={DEVICE} work_dir={cfg.work_dir}")
1028
+ if DEVICE == "cuda":
1029
+ print(f"gpu={torch.cuda.get_device_name()} "
1030
+ f"vram={torch.cuda.get_device_properties(0).total_memory/1e9:.0f}GB")
1031
+ miss = src0(cfg).get("missing", {})
1032
+ print(f"chunks: {len(usable_chunks(cfg))} train + {len(cfg.holdout_chunks)} holdout "
1033
+ f"| excluded for missing experts: {miss}")
1034
+ if not cfg.require_all_experts:
1035
+ print(" !! require_all_experts=False -> 4-expert consensus on some chunks.")
1036
+ print(" !! The target definition then differs BETWEEN chunks. Discouraged.")
1037
+ bk = Backup(cfg)
1038
+
1039
+ if cfg.run_stage0:
1040
+ cfg.caption_field = stage0_parity(cfg)
1041
+ elif cfg.caption_field is None:
1042
+ raise RuntimeError("caption_field is None and stage 0 is disabled.")
1043
+
1044
+ maps = stage1_fit(cfg, bk) if cfg.run_stage1 else torch.load(
1045
+ f"{paths(cfg)['maps']}/alignment_maps.pt", weights_only=False)
1046
+ chunks = stage2_targets(cfg, maps, bk) if cfg.run_stage2 else sorted(
1047
+ set(usable_chunks(cfg)) | set(cfg.holdout_chunks))
1048
+ if cfg.run_stage3:
1049
+ return stage3_train(cfg, chunks, bk)
1050
+
1051
+
1052
+ if "get_ipython" in globals() or __name__ == "__main__":
1053
+ STUDENT = run(CFG)