| """SN3 cycle-63 salvo: FIRST full-objective-coverage CPT, on a 2x4 (batch x peak_lr) grid. |
| |
| WHY (read /var/lib/mining/sn-packages/sn3/learning/ first): |
| |
| Cycle 62 found the root cause of ~15 failed cycles: the trainable `.npy` shards on |
| `mining-datasets` covered only **9% of the eval objective weight** (dendrite 0.04 + |
| cosmopedia 0.05). `trainmix-v1` fixes that (10/10 corpora, exact validator weights, |
| eval-disjointness verified by byte-range intersection, 871 contaminated rows excised). |
| |
| But the cycle-62 ladder log carries a second, sharper signal that the coverage story |
| alone does not explain: at `peak_lr=1e-5, accum=16` (B = 32 seqs = 65k tok), after only |
| 8 updates **every corpus got worse, including `dendrite-synth-run` itself** (mu |
| -0.002923) — the very corpus being trained on. Training a corpus cannot hurt that |
| corpus if the step is signal-dominated. And damage scaled as lr^2 |
| (u=8: lr1e-5 -0.00132, lr4e-5 -0.01867 ~ 14x for a 4x lr). |
| |
| That is the signature of the noise term of SGD around a minimum: |
| |
| E[dLoss] ~ -lr * |g|^2 + (lr^2 / 2) * tr(H Sigma) / B |
| |
| i.e. the king sits where the *gradient-noise* term dominates at B=32. Fitting the |
| observed damage gives c/B|_{B=32} ~ 1.65e6 per update per lr^2, so the improvement |
| window is lr < 2*B*|g|^2 / tr(H Sigma) — it widens **linearly in B**. Lowering lr alone |
| shrinks the signal too (it is only linear in lr) and drives mu -> 0; raising B is the |
| only knob that moves the optimum *and* raises its height. |
| |
| So the single most informative salvo is not a 1-D lr sweep: it is the 2x4 grid |
| {accum 16, 64} x {1e-6, 3e-6, 1e-5, 3e-5}. Both marginals are readable, the predicted |
| optima for both batches (~5e-6 and ~2e-5 at s~17) are bracketed, and the cell |
| (accum 16, lr 1e-5) is a **direct control against cycle-62's own measurement** — |
| identical recipe, only the corpus coverage changed (0.09 -> 1.00). |
| |
| CONTRACTS |
| * base = REIGNING king (reign 179, hf:ce6efba4...). Never a dethroned king. |
| * train data = /data/trainmix-v1 ONLY. Never the eval set (/data/truemix-v1). |
| * u=0 wiring check vs an INDEPENDENT cached per-sequence reference; fail closed. |
| * gate-facing `mu_hat` is the FINAL eval point, not the argmax over the trajectory |
| (no peeking optimism); the full trajectory is kept in detail.json for analysis. |
| * writes lab-bus result.json so `python -m lab.collect|select` applies the real gate. |
| """ |
| from __future__ import annotations |
|
|
| import json |
| import os |
| import time |
|
|
| import modal |
|
|
| APP = "mining-trainmix-salvo" |
| SEQ_LEN = 2048 |
| VOCAB_SIZE = 248320 |
|
|
| TRAIN_MIX = "/data/trainmix-v1" |
| EVAL_MIX = "/data/truemix-v1" |
| KING = "/kings/king-647-cp0" |
| KING_DIGEST = "hf:ce6efba444047328046a511e21d3e37a84c329c5" |
| REF_PATH = "/data/truemix-results/rebase-newking-n300/per_sequence.json" |
| LAB_BUS = "/data/lab-results" |
|
|
| image = ( |
| modal.Image.from_registry("nvidia/cuda:12.8.1-devel-ubuntu24.04", add_python="3.12") |
| .env( |
| { |
| "MAX_JOBS": "8", |
| "TORCH_CUDA_ARCH_LIST": "9.0", |
| "CAUSAL_CONV1D_FORCE_BUILD": "TRUE", |
| "HF_HUB_ENABLE_HF_TRANSFER": "1", |
| "CUDA_HOME": "/usr/local/cuda", |
| } |
| ) |
| .apt_install("git", "build-essential") |
| .pip_install("torch==2.8.0", index_url="https://download.pytorch.org/whl/cu128") |
| .pip_install( |
| "transformers==5.8.0", "safetensors", "accelerate", "einops", "numpy", |
| "huggingface_hub", "hf_transfer", "ninja", "packaging", "setuptools", "wheel", |
| ) |
| .pip_install("flash-linear-attention") |
| .pip_install( |
| "https://github.com/Dao-AILab/causal-conv1d/releases/download/v1.6.2.post1/causal_conv1d-1.6.2.post1%2Bcu12torch2.8cxx11abiTRUE-cp312-cp312-linux_x86_64.whl" |
| ) |
| ) |
|
|
| app = modal.App(APP) |
| kings = modal.Volume.from_name("mining-kings") |
| ckpts = modal.Volume.from_name("mining-checkpoints") |
| datasets = modal.Volume.from_name("mining-datasets") |
| VOLS = {"/kings": kings, "/ckpts": ckpts, "/data": datasets} |
|
|
|
|
| def _log(m): |
| print("[%s] %s" % (time.strftime("%H:%M:%S"), m), flush=True) |
|
|
|
|
| @app.function(image=image, gpu="B200:1", volumes=VOLS, timeout=60 * 180) |
| def variant( |
| tag: str, |
| variant_id: str, |
| peak_lr: float, |
| accum: int, |
| max_updates: int = 40, |
| micro_bs: int = 2, |
| warmup: int = 4, |
| min_lr_frac: float = 1.0, |
| train_top_k: int = 8, |
| train_lm_head: bool = True, |
| eval_every: int = 10, |
| eval_n: int = 300, |
| seed: int = 777, |
| ) -> dict: |
| import math |
|
|
| import numpy as np |
| import torch |
| from transformers import AutoModelForCausalLM |
|
|
| t0 = time.time() |
| outdir = f"{LAB_BUS}/{tag}/variants/{variant_id}" |
| os.makedirs(outdir, exist_ok=True) |
|
|
| |
| eidx = json.load(open(f"{EVAL_MIX}/index.json")) |
| corpora = [(c["corpus"], float(c["weight"]), c["file"]) for c in eidx["corpora"]] |
| eval_tok = {} |
| for name, _w, fn in corpora: |
| arr = np.load(f"{EVAL_MIX}/{fn}", mmap_mode="r") |
| if arr.shape[0] < eval_n: |
| raise RuntimeError("eval corpus %s has %d < %d sequences" % (name, arr.shape[0], eval_n)) |
| eval_tok[name] = np.ascontiguousarray(arr[:eval_n]) |
|
|
| blob = json.load(open(REF_PATH)) |
| if KING not in blob: |
| raise RuntimeError("reference %s has no key %s (keys=%s)" % (REF_PATH, KING, list(blob))) |
| ref = blob[KING] |
| king_ref = {} |
| for n_, _w, _f in corpora: |
| v = np.array(ref[n_][:eval_n], dtype=np.float64) |
| if v.shape[0] != eval_tok[n_].shape[0]: |
| raise RuntimeError("ref/eval length mismatch %s: %d vs %d" |
| % (n_, v.shape[0], eval_tok[n_].shape[0])) |
| king_ref[n_] = v |
| _log("king ref %s key=%s (n=%d/corpus)" % (REF_PATH, KING, eval_n)) |
|
|
| |
| tidx = json.load(open(f"{TRAIN_MIX}/index.json")) |
| if abs(float(tidx["weight_total"]) - 1.0) > 1e-9: |
| raise RuntimeError("trainmix weight_total != 1.0: %r" % tidx["weight_total"]) |
| if "eval_disjointness" not in tidx: |
| raise RuntimeError("trainmix index has no eval_disjointness block - refusing to train") |
| pool = [] |
| for c in tidx["corpora"]: |
| p = f"{TRAIN_MIX}/{c['file']}" |
| a = np.load(p, mmap_mode="r") |
| if a.ndim != 2 or a.shape[1] != SEQ_LEN: |
| raise RuntimeError("bad train shard shape %s %s" % (p, a.shape)) |
| pool.append((c["corpus"], float(c["weight"]), a)) |
| _log(" train %-45s w=%.2f seqs=%d" % (c["corpus"], float(c["weight"]), a.shape[0])) |
| train_names = {n_ for n_, _w, _a in pool} |
| eval_names = {n_ for n_, _w, _f in corpora} |
| if train_names != eval_names: |
| raise RuntimeError("train/eval corpus set mismatch: %s" % (train_names ^ eval_names)) |
| w = np.array([x[1] for x in pool], dtype=np.float64) |
| w = w / w.sum() |
| cum = np.cumsum(w) |
| n_train_seqs = int(sum(int(a.shape[0]) for _n, _w2, a in pool)) |
| _log("train pool: %d corpora, %d sequences, coverage=%.2f" % (len(pool), n_train_seqs, w.sum())) |
|
|
| |
| rs_audit = np.random.default_rng(12345) |
| bad = 0 |
| for _ in range(2000): |
| ci = int(np.searchsorted(cum, rs_audit.random())) |
| arr = pool[min(ci, len(pool) - 1)][2] |
| si = int(rs_audit.integers(0, arr.shape[0])) |
| seq = np.asarray(arr[si]) |
| if seq.shape[0] != SEQ_LEN or (seq >= VOCAB_SIZE).any(): |
| bad += 1 |
| _log("OOB audit: %d/2000 sampled trainmix sequences unusable" % bad) |
| if bad > 20: |
| raise RuntimeError("trainmix OOB rate %d/2000 too high - fail closed" % bad) |
|
|
| oob = {"n": 0, "drawn": 0} |
| drawn_by_corpus = {n_: 0 for n_, _w2, _a in pool} |
|
|
| def make_draw(rs): |
| def draw(k): |
| out = np.zeros((k, SEQ_LEN), dtype=np.int64) |
| for j in range(k): |
| ci = min(int(np.searchsorted(cum, rs.random())), len(pool) - 1) |
| cname, _w2, arr = pool[ci] |
| si = int(rs.integers(0, arr.shape[0])) |
| seq = np.asarray(arr[si], dtype=np.int64) |
| oob["drawn"] += 1 |
| drawn_by_corpus[cname] += 1 |
| if seq.shape[0] != SEQ_LEN or (seq >= VOCAB_SIZE).any(): |
| oob["n"] += 1 |
| seq = np.zeros(SEQ_LEN, dtype=np.int64) |
| out[j] = seq |
| return out |
| return draw |
|
|
| |
| import sys as _sys |
|
|
| if KING not in _sys.path: |
| _sys.path.insert(0, KING) |
| _log("loading king %s" % KING) |
| model = AutoModelForCausalLM.from_pretrained( |
| KING, trust_remote_code=True, dtype=torch.bfloat16, low_cpu_mem_usage=True |
| ).cuda() |
| model.config.use_cache = False |
| try: |
| model.gradient_checkpointing_enable() |
| _log("gradient checkpointing on") |
| except Exception as e: |
| _log("no gradient checkpointing: %s" % e) |
|
|
| nl = model.config.num_hidden_layers |
| keep_from = nl - train_top_k |
| trainable = [] |
| for n_, p in model.named_parameters(): |
| want = False |
| if ".layers." in n_: |
| try: |
| li = int(n_.split(".layers.")[1].split(".")[0]) |
| except Exception: |
| li = -1 |
| want = li >= keep_from |
| elif "lm_head" in n_: |
| want = train_lm_head |
| elif n_.endswith("model.norm.weight") or n_.endswith("final_layernorm.weight"): |
| want = True |
| p.requires_grad_(want) |
| if want: |
| trainable.append((n_, p)) |
| ntp = sum(p.numel() for _n, p in trainable) |
| _log("trainable %d tensors / %.3fB params (layers >= %d of %d, lm_head=%s)" |
| % (len(trainable), ntp / 1e9, keep_from, nl, train_lm_head)) |
|
|
| def chunked_ce(logits, tgt, step=256): |
| b, tm1, _ = logits.shape |
| tot = None |
| for s in range(0, tm1, step): |
| sl = logits[:, s : s + step, :].float() |
| st = tgt[:, s : s + step] |
| ce = torch.nn.functional.cross_entropy( |
| sl.reshape(-1, sl.shape[-1]), st.reshape(-1), reduction="sum" |
| ) |
| tot = ce if tot is None else tot + ce |
| return tot / (b * tm1) |
|
|
| @torch.no_grad() |
| def eval_objective(): |
| model.eval() |
| per = {} |
| for name, _w2, _fn in corpora: |
| toks = eval_tok[name] |
| out = np.zeros(toks.shape[0], dtype=np.float64) |
| for i in range(0, toks.shape[0], 2): |
| ids = torch.from_numpy(toks[i : i + 2].astype(np.int64)).cuda() |
| lg = model(input_ids=ids, use_cache=False).logits[:, :-1, :] |
| tg = ids[:, 1:] |
| b, tm1, _ = lg.shape |
| acc = torch.zeros(b, dtype=torch.float64, device=lg.device) |
| for s in range(0, tm1, 256): |
| sl = lg[:, s : s + 256, :].float() |
| st = tg[:, s : s + 256] |
| ce = torch.nn.functional.cross_entropy( |
| sl.reshape(-1, sl.shape[-1]), st.reshape(-1), reduction="none" |
| ).view(b, -1) |
| acc += ce.sum(dim=1).double() |
| out[i : i + ids.shape[0]] = (acc / tm1).cpu().numpy() |
| del lg |
| per[name] = out |
| model.train() |
| mu, contrib, var = {}, {}, 0.0 |
| for name, wt, _fn in corpora: |
| d = king_ref[name] - per[name] |
| mu[name] = float(d.mean()) |
| contrib[name] = float(wt * mu[name]) |
| var += (wt ** 2) * float(d.var(ddof=1)) / len(d) |
| se = math.sqrt(var) |
| m = float(sum(contrib.values())) |
| return { |
| "mu_hat": m, "se": se, "lcb999": m - 3.09 * se, |
| "mu_per_corpus": mu, "weighted_contribution": contrib, |
| "cand_loss": {k: float(v.mean()) for k, v in per.items()}, |
| } |
|
|
| recipe = { |
| "peak_lr": peak_lr, "accum": accum, "micro_bs": micro_bs, |
| "effective_batch_seqs": micro_bs * accum, |
| "effective_batch_tokens": micro_bs * accum * SEQ_LEN, |
| "max_updates": max_updates, "warmup": warmup, "min_lr_frac": min_lr_frac, |
| "train_top_k": train_top_k, "train_lm_head": train_lm_head, |
| "seed": seed, "data_order_seed": seed, "eval_n": eval_n, |
| "train_data": TRAIN_MIX, "objective_coverage": 1.0, |
| "king": KING, "king_digest": KING_DIGEST, "precision": "bf16", |
| } |
|
|
| detail = { |
| "tag": tag, "variant_id": variant_id, "recipe": recipe, |
| "train_index": {k: tidx.get(k) for k in |
| ("name", "seed", "weight_total", "tokenizer", "eval_disjointness")}, |
| "n_train_sequences": n_train_seqs, |
| "eval_ref": REF_PATH, "eval_mix": EVAL_MIX, |
| } |
|
|
| def persist(status, gate_eval, base, hist, note=""): |
| used_real = bool(n_train_seqs > 0 and oob["n"] == 0 and bad <= 20) |
| promotable = bool(used_real and status == "ok") |
| bus = { |
| "variant_id": variant_id, |
| "tag": tag, |
| "mu_hat": gate_eval["mu_hat"] if gate_eval else 0.0, |
| "se": gate_eval["se"] if gate_eval else 1.0, |
| "lcb999": (gate_eval["mu_hat"] - 3.09 * gate_eval["se"]) if gate_eval else -1.0, |
| "king_digest": KING_DIGEST, |
| "used_real_data": used_real, |
| "real_data": used_real, |
| "promotable": promotable, |
| "oob_count": oob["n"], |
| "oob_drawn": oob["drawn"], |
| "effective_accumulation": accum, |
| "effective_peak_lr": peak_lr, |
| "recipe": recipe, |
| "baseline_u0": base, |
| "status": status, |
| "gate_point": "final_eval_point (no argmax peeking)", |
| "note": note, |
| } |
| with open(f"{outdir}/result.json", "w") as fh: |
| json.dump(bus, fh, indent=2) |
| with open(f"{outdir}/detail.json", "w") as fh: |
| json.dump({**detail, "baseline_u0": base, "history": hist, |
| "oob": dict(oob), "drawn_by_corpus": dict(drawn_by_corpus), |
| "status": status, "elapsed_s": round(time.time() - t0, 1)}, fh, indent=2) |
| datasets.commit() |
|
|
| |
| _log("=== u=0 baseline (must reproduce the pinned king vs independent reference) ===") |
| base = eval_objective() |
| _log("u=0 mu_hat=%+.6f se=%.6f" % (base["mu_hat"], base["se"])) |
| tol = max(5e-5, 4.0 * base["se"]) |
| if abs(base["mu_hat"]) > tol: |
| persist("wiring_failed", None, base, [], note="u=0 mu_hat %.6g exceeds tol %.6g" % (base["mu_hat"], tol)) |
| raise RuntimeError("baseline mu_hat=%.6f exceeds tol %.6f -> wrong king/reference pairing" |
| % (base["mu_hat"], tol)) |
| _log("baseline wiring OK (tol %.6g)" % tol) |
|
|
| hist = [{"update": 0, "tokens": 0, **base}] |
| persist("running", base, base, hist, note="u=0 only") |
|
|
| |
| masters = [p.detach().float().clone() for _n, p in trainable] |
| for m in masters: |
| m.requires_grad_(False) |
| opt = torch.optim.AdamW(masters, lr=peak_lr, betas=(0.9, 0.95), eps=1e-8, weight_decay=0.0) |
| rs = np.random.default_rng(seed) |
| draw = make_draw(rs) |
|
|
| def lr_at(u): |
| if u < warmup: |
| return peak_lr * (u + 1) / warmup |
| prog = (u - warmup) / max(1, max_updates - warmup) |
| return peak_lr * (min_lr_frac + (1 - min_lr_frac) * 0.5 * (1 + math.cos(math.pi * min(1.0, prog)))) |
|
|
| model.train() |
| last = base |
| for u in range(max_updates): |
| lr = lr_at(u) |
| for g in opt.param_groups: |
| g["lr"] = lr |
| opt.zero_grad(set_to_none=True) |
| for _n, p in trainable: |
| p.grad = None |
| tl, tu = 0.0, time.time() |
| for _a in range(accum): |
| ids = torch.from_numpy(draw(micro_bs)).cuda() |
| logits = model(input_ids=ids, use_cache=False).logits[:, :-1, :] |
| loss = chunked_ce(logits, ids[:, 1:]) / accum |
| loss.backward() |
| tl += float(loss.detach()) * accum |
| del logits, loss |
| gn = 0.0 |
| for (nm, p), m in zip(trainable, masters): |
| if p.grad is None: |
| continue |
| m.grad = p.grad.float() |
| gn += float(m.grad.pow(2).sum()) |
| gn = gn ** 0.5 |
| torch.nn.utils.clip_grad_norm_(masters, 1.0) |
| opt.step() |
| with torch.no_grad(): |
| for (nm, p), m in zip(trainable, masters): |
| p.data.copy_(m.to(p.dtype)) |
| for m in masters: |
| m.grad = None |
| if u % 5 == 0 or u == max_updates - 1: |
| _log(" u=%d loss=%.5f lr=%.3e gnorm=%.3f %.1fs/upd mem=%.1fGB" |
| % (u + 1, tl / accum, lr, gn, time.time() - tu, |
| torch.cuda.max_memory_allocated() / 1e9)) |
| if (u + 1) % eval_every == 0 or (u + 1) == max_updates: |
| e = eval_objective() |
| last = e |
| _log(" *** %s u=%d mu_hat=%+.6f se=%.6f lcb999=%+.6f" % |
| (variant_id, u + 1, e["mu_hat"], e["se"], e["mu_hat"] - 3.09 * e["se"])) |
| hist.append({"update": u + 1, "tokens": (u + 1) * micro_bs * accum * SEQ_LEN, **e}) |
| persist("running", e, base, hist, note="u=%d" % (u + 1)) |
|
|
| |
| persist("ok", last, base, hist, note="final eval point at u=%d" % max_updates) |
| best = max(hist, key=lambda h: h["mu_hat"]) |
| summary = { |
| "tag": tag, "variant_id": variant_id, "peak_lr": peak_lr, "accum": accum, |
| "u0_mu_hat": base["mu_hat"], |
| "final_mu_hat": last["mu_hat"], "final_se": last["se"], |
| "final_lcb999": last["mu_hat"] - 3.09 * last["se"], |
| "best_mu_hat": best["mu_hat"], "best_update": best["update"], |
| "oob": dict(oob), "drawn_by_corpus": dict(drawn_by_corpus), |
| "elapsed_s": round(time.time() - t0, 1), |
| } |
| _log("done " + json.dumps(summary)) |
| return summary |
|
|
|
|
| @app.local_entrypoint() |
| def main( |
| tag: str = "tmix63", |
| variant_id: str = "b16-lr1e-06", |
| peak_lr: float = 1e-6, |
| accum: int = 16, |
| max_updates: int = 40, |
| eval_every: int = 10, |
| eval_n: int = 300, |
| micro_bs: int = 2, |
| warmup: int = 4, |
| train_top_k: int = 8, |
| seed: int = 777, |
| ): |
| r = variant.remote( |
| tag=tag, variant_id=variant_id, peak_lr=peak_lr, accum=accum, |
| max_updates=max_updates, micro_bs=micro_bs, warmup=warmup, |
| train_top_k=train_top_k, eval_every=eval_every, eval_n=eval_n, seed=seed, |
| ) |
| print(json.dumps(r, indent=2)[:6000]) |
|
|