train: Gate 3 preflight driver (reader, hub cycle, T1 throughput, cold resume)
Browse files- train/preflight.py +261 -0
train/preflight.py
ADDED
|
@@ -0,0 +1,261 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Gate 3 preflight driver: proves the pipeline, not the model, and writes the evidence itself.
|
| 2 |
+
|
| 3 |
+
# Every test here answers a question §5 Phase 3 lists, and each one emits a machine-readable verdict
|
| 4 |
+
# line so a future session can read the outcome without re-reading a log. GPU stages are few and short:
|
| 5 |
+
# the whole file is designed to fit inside the 6 GPU-hour lifetime test cap (memory/QUOTA.md), and it
|
| 6 |
+
# prints what it spent so the ledger can be updated from the log.
|
| 7 |
+
|
| 8 |
+
# Stages, in the cheap-first order they should run:
|
| 9 |
+
# P0 (CPU) reader: cursor math, resume slicing, determinism, val split really held out
|
| 10 |
+
# P1 (CPU) hub cycle: push_and_prune with a real-sized 100M checkpoint + optimizer state
|
| 11 |
+
# P2 (GPU) throughput: 20L vs 22L at seq 1024/2048 -- test T1, the number the main-run ETA is built on
|
| 12 |
+
# P3 (GPU) short train -> kill -> cold resume from the Hub on an empty disk -> loss continuity
|
| 13 |
+
# P4 (GPU) resume twice in sequence; the cursor must advance monotonically, never re-read
|
| 14 |
+
#
|
| 15 |
+
# P3/P4 are the tests that matter most and the ones that cannot be faked: a run that resumes from what it
|
| 16 |
+
# left on disk proves nothing, because every real interruption takes the disk with it (§3.13).
|
| 17 |
+
|
| 18 |
+
import argparse
|
| 19 |
+
import json
|
| 20 |
+
import os
|
| 21 |
+
import subprocess
|
| 22 |
+
import sys
|
| 23 |
+
import time
|
| 24 |
+
|
| 25 |
+
WORK = "/kaggle/working"
|
| 26 |
+
REV_DEFAULT = "" # filled from the launcher; only used for reporting
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def sh(argv, timeout=None, env=None, label=""):
|
| 30 |
+
print(f"=== {label or ' '.join(argv[:3])}", flush=True)
|
| 31 |
+
t0 = time.time()
|
| 32 |
+
p = subprocess.run(argv, cwd=WORK, capture_output=True, text=True, timeout=timeout,
|
| 33 |
+
env=dict(os.environ, **(env or {})))
|
| 34 |
+
for line in (p.stdout or "").splitlines():
|
| 35 |
+
print(" |", line[:240], flush=True)
|
| 36 |
+
if p.returncode != 0:
|
| 37 |
+
print(" STDERR:", (p.stderr or "")[-3000:], flush=True)
|
| 38 |
+
return {"rc": p.returncode, "out": (p.stdout or "")[-200000:],
|
| 39 |
+
"err": (p.stderr or "")[-3000:], "seconds": round(time.time() - t0, 1)}
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def last_json(text, begin, end):
|
| 43 |
+
if begin not in text or end not in text:
|
| 44 |
+
return None
|
| 45 |
+
body = text.rsplit(begin, 1)[1].split(end, 1)[0]
|
| 46 |
+
try:
|
| 47 |
+
return json.loads(body)
|
| 48 |
+
except Exception:
|
| 49 |
+
return None
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
# ------------------------------------------------------------------------ P0 reader (CPU, free)
|
| 53 |
+
def p0_reader(args, R):
|
| 54 |
+
"""The reader is where an unrecoverable main run would hide: if two runs at the same cursor read
|
| 55 |
+
different tokens, every resume silently trains on a subset. Checked on bytes, not on feelings."""
|
| 56 |
+
src = r'''
|
| 57 |
+
import json, os, sys, numpy as np
|
| 58 |
+
sys.path.insert(0, "/kaggle/working")
|
| 59 |
+
import shard_dataset as SD
|
| 60 |
+
man = json.load(open("/kaggle/working/mixroot/manifest.json"))
|
| 61 |
+
store = SD.PackedTokenStore("/kaggle/working/mixroot", man)
|
| 62 |
+
L = 256
|
| 63 |
+
full = SD.make_dataset(store, L, 1234)
|
| 64 |
+
n = len(full)
|
| 65 |
+
part = SD.make_dataset(store, L, 1234, start_sample=n - 40)
|
| 66 |
+
same = all(bool((full[n - 40 + j]["input_ids"] == part[j]["input_ids"]).all()) for j in range(40))
|
| 67 |
+
labels_ok = bool((full[0]["labels"][:-1] == full[0]["input_ids"][1:]).all()
|
| 68 |
+
and int(full[0]["labels"][-1]) == int(full[0]["input_ids"][0]))
|
| 69 |
+
det2 = SD.make_dataset(store, L, 1234)
|
| 70 |
+
deterministic = bool((full[7]["input_ids"] == det2[7]["input_ids"]).all())
|
| 71 |
+
diffseed = SD.make_dataset(store, L, 999)
|
| 72 |
+
seed_matters = not bool((full[7]["input_ids"] == diffseed[7]["input_ids"]).all())
|
| 73 |
+
vstore = SD.PackedTokenStore("/kaggle/working/mixroot", man,
|
| 74 |
+
files=[s["file"] for s in man["val_shards"]])
|
| 75 |
+
tv = set(); vv = set()
|
| 76 |
+
vds = SD.make_dataset(vstore, L, 0, shuffle=False, count=min(400, vstore.total_tokens // L))
|
| 77 |
+
for i in range(min(len(full), 4000)):
|
| 78 |
+
tv.add(full[i]["input_ids"][:64].numpy().tobytes())
|
| 79 |
+
for i in range(len(vds)):
|
| 80 |
+
vv.add(vds[i]["input_ids"][:64].numpy().tobytes())
|
| 81 |
+
print("READER_JSON_BEGIN")
|
| 82 |
+
print(json.dumps({
|
| 83 |
+
"total_tokens": store.total_tokens, "samples_at_L256": n,
|
| 84 |
+
"resume_matches_uninterrupted": same, "labels_are_next_token": labels_ok,
|
| 85 |
+
"same_seed_same_order": deterministic, "different_seed_different_order": seed_matters,
|
| 86 |
+
"val_tokens": vstore.total_tokens, "val_windows_checked": len(vds),
|
| 87 |
+
"train_val_prefix_collision": len(tv & vv),
|
| 88 |
+
"tokens_per_shard_min": min(s["tokens"] for s in man["shards"]),
|
| 89 |
+
"n_shards": len(man["shards"]),
|
| 90 |
+
}))
|
| 91 |
+
print("READER_JSON_END")
|
| 92 |
+
'''
|
| 93 |
+
r = sh([sys.executable, "-c", src], label="P0 reader")
|
| 94 |
+
R["P0"] = last_json(r["out"], "READER_JSON_BEGIN", "READER_JSON_END")
|
| 95 |
+
R["P0_rc"] = r["rc"]
|
| 96 |
+
j = R["P0"] or {}
|
| 97 |
+
R["P0_pass"] = bool(r["rc"] == 0 and j.get("resume_matches_uninterrupted")
|
| 98 |
+
and j.get("labels_are_next_token") and j.get("same_seed_same_order")
|
| 99 |
+
and j.get("different_seed_different_order")
|
| 100 |
+
and j.get("train_val_prefix_collision") == 0
|
| 101 |
+
and j.get("val_tokens", 0) > 0)
|
| 102 |
+
print("VERDICT P0_pass=", R["P0_pass"], flush=True)
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
# ------------------------------------------------------------------------ P1 hub cycle (CPU, free)
|
| 106 |
+
def p1_hub(args, R):
|
| 107 |
+
"""§3.13 at real size: a 106M-parameter checkpoint with optimizer state, pushed, verified from the
|
| 108 |
+
Hub by re-listing and re-hashing, then pruned. Cheap because the weights are random; the BYTES are
|
| 109 |
+
what is being timed."""
|
| 110 |
+
src = r'''
|
| 111 |
+
import json, os, sys, time, numpy as np
|
| 112 |
+
sys.path.insert(0, "/kaggle/working")
|
| 113 |
+
import ounce100m_credentials, hubckpt
|
| 114 |
+
ounce100m_credentials.install()
|
| 115 |
+
from huggingface_hub import HfApi
|
| 116 |
+
api = HfApi(); tok = os.environ["HF_TOKEN"]
|
| 117 |
+
repo = "Cion-lab/ounce100m-ckptbench-DELETEME"
|
| 118 |
+
try:
|
| 119 |
+
api.create_repo(repo_id=repo, repo_type="dataset", exist_ok=True, token=tok)
|
| 120 |
+
except Exception as e:
|
| 121 |
+
print("repo create:", type(e).__name__, str(e)[:120])
|
| 122 |
+
d = "/kaggle/working/ckptbench/checkpoint-1"
|
| 123 |
+
os.makedirs(d, exist_ok=True)
|
| 124 |
+
# 106,194,240 params x 4 B x 4 arrays (fp16-saved weights + fp32 master + Adam m and v) is the real
|
| 125 |
+
# checkpoint footprint, and it is the size D-007 measured at 1.6 GB. Timing a 425 MB mock would flatter
|
| 126 |
+
# the main run.
|
| 127 |
+
n = 106194240
|
| 128 |
+
t0 = time.time()
|
| 129 |
+
for name in ("weights.fp32", "master.fp32", "adam_m.fp32", "adam_v.fp32"):
|
| 130 |
+
a = np.lib.format.open_memmap(os.path.join(d, "state." + name.replace(".", "_") + ".npy"),
|
| 131 |
+
dtype=np.float32, mode="w+", shape=(n,))
|
| 132 |
+
a[:] = np.float32(0.0)
|
| 133 |
+
a.flush()
|
| 134 |
+
del a
|
| 135 |
+
json.dump({"step": 1, "samples_consumed": 381500}, open(os.path.join(d, "cursor.json"), "w"))
|
| 136 |
+
up_t0 = time.time()
|
| 137 |
+
try:
|
| 138 |
+
res = hubckpt.push_and_prune(repo, d, "ckpt/checkpoint-1", api, token=tok, prune=True)
|
| 139 |
+
ver = res["verify"]
|
| 140 |
+
# second, independent proof: pull it back into a clean directory and compare hashes
|
| 141 |
+
back = "/kaggle/working/ckptbench/restored"
|
| 142 |
+
got = hubckpt.download_checkpoint(repo, "ckpt/checkpoint-1", back, api, token=tok)
|
| 143 |
+
same = json.load(open(os.path.join(back, "cursor.json"))) == {"step": 1,
|
| 144 |
+
"samples_consumed": 381500}
|
| 145 |
+
# The repo is a timing rig, not an artifact: leaving a 1.7 GB public blob invites a future session
|
| 146 |
+
# to mistake it for a checkpoint. Everything measurable is in the log by this point.
|
| 147 |
+
api.delete_repo(repo_id=repo, repo_type="dataset", token=tok)
|
| 148 |
+
except Exception as e:
|
| 149 |
+
res, ver, got, same = {"error": f"{type(e).__name__}: {str(e)[:300]}"}, {}, {}, False
|
| 150 |
+
print("HUB_JSON_BEGIN")
|
| 151 |
+
print(json.dumps({"push_verify_seconds": round(time.time()-up_t0,1),
|
| 152 |
+
"cycle": {k: res.get(k) for k in ("verify","pruned","free_before_gb","free_after_gb")},
|
| 153 |
+
"downloaded": got, "readback_ok": bool(ver.get("ok")) and same,
|
| 154 |
+
"gen_seconds": round(time.time()-t0,1)}, default=str))
|
| 155 |
+
print("HUB_JSON_END")
|
| 156 |
+
'''
|
| 157 |
+
r = sh([sys.executable, "-c", src], timeout=5400, label="P1 hub cycle (425 MB mock checkpoint)")
|
| 158 |
+
R["P1"] = last_json(r["out"], "HUB_JSON_BEGIN", "HUB_JSON_END")
|
| 159 |
+
R["P1_rc"] = r["rc"]
|
| 160 |
+
j = R["P1"] or {}
|
| 161 |
+
cyc = ((j.get("cycle") or {}).get("verify") or {})
|
| 162 |
+
R["P1_pass"] = bool(r["rc"] == 0 and cyc.get("ok") and j.get("readback_ok")
|
| 163 |
+
and (j.get("cycle") or {}).get("pruned"))
|
| 164 |
+
print("VERDICT P1_pass=", R["P1_pass"], json.dumps(cyc)[:300], flush=True)
|
| 165 |
+
R["P1_cleanup"] = "delete Cion-lab/ounce100m-ckptbench-DELETEME when done reading it"
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
# ------------------------------------------------------------------------ P2..P4 (GPU)
|
| 169 |
+
def torchrun(args, extra, timeout=7200):
|
| 170 |
+
return sh(["torchrun", "--nproc_per_node=2", "train_ounce100m.py", "--root",
|
| 171 |
+
WORK + "/mixroot", "--seq-len", str(args.seq_len)] + extra,
|
| 172 |
+
timeout=timeout, label="torchrun " + " ".join(extra[:6]))
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def p2_throughput(args, R):
|
| 176 |
+
"""T1: 20L vs 22L at two sequence lengths, on the real data, for enough steps that the number is
|
| 177 |
+
steady-state. The main run's whole schedule is division by this number."""
|
| 178 |
+
out = {}
|
| 179 |
+
for layers in (20, 22):
|
| 180 |
+
for seq in (1024, 2048):
|
| 181 |
+
micro = 8 if seq == 1024 else 2
|
| 182 |
+
r = torchrun(args, ["--layers", str(layers), "--hidden", "576", "--seq-len", str(seq),
|
| 183 |
+
"--micro-batch", str(micro), "--accum", "1", "--max-steps", "25",
|
| 184 |
+
"--out", WORK + f"/t1_{layers}_{seq}", "--log-every", "5"],
|
| 185 |
+
timeout=5400)
|
| 186 |
+
out[f"L{layers}_s{seq}"] = {"rc": r["rc"], "seconds": r["seconds"],
|
| 187 |
+
"tail": r["out"][-1200:]}
|
| 188 |
+
R["P2"] = out
|
| 189 |
+
R["P2_pass"] = all(v["rc"] == 0 for v in out.values()) and len(out) == 4
|
| 190 |
+
print("VERDICT P2_pass=", R["P2_pass"], flush=True)
|
| 191 |
+
|
| 192 |
+
|
| 193 |
+
def p3_cold_resume(args, R):
|
| 194 |
+
"""Short run -> wipe every local trace -> resume with --resume auto, which must recover the exact
|
| 195 |
+
sample position from the Hub. Loss continuity is the pass condition: a restart that silently
|
| 196 |
+
re-initialised would jump the loss."""
|
| 197 |
+
common = ["--tokens", str(args.tokens), "--accum", str(args.accum), "--micro-batch", "2",
|
| 198 |
+
"--hub-repo", args.ckpt_repo, "--prune", "--log-every", "5"]
|
| 199 |
+
a = torchrun(args, common + ["--max-steps", str(args.steps_a), "--out", WORK + "/p3"],
|
| 200 |
+
timeout=9000)
|
| 201 |
+
# erase the instance's memory of the run: local checkpoints, the run dir, and the downloaded mix?
|
| 202 |
+
# No -- the mix is the dataset and a real interruption keeps it. Only the run state goes.
|
| 203 |
+
wipe = sh(["bash", "-c", f"rm -rf {WORK}/p3 {WORK}/run; df -h {WORK} | tail -1"],
|
| 204 |
+
label="P3 wipe local run state")
|
| 205 |
+
b = torchrun(args, common + ["--max-steps", str(args.steps_a + args.steps_b),
|
| 206 |
+
"--out", WORK + "/p3b", "--resume", "auto"], timeout=9000)
|
| 207 |
+
R["P3"] = {"first": {"rc": a["rc"], "tail": a["out"][-2500:]},
|
| 208 |
+
"wipe": wipe["out"][-400:],
|
| 209 |
+
"resumed": {"rc": b["rc"], "tail": b["out"][-2500:]}}
|
| 210 |
+
la = [l for l in a["out"].splitlines() if "'loss'" in l or "loss=" in l]
|
| 211 |
+
lb = [l for l in b["out"].splitlines() if "'loss'" in l or "loss=" in l]
|
| 212 |
+
R["P3_loss_last_before"] = la[-1][:200] if la else None
|
| 213 |
+
R["P3_loss_first_after"] = lb[0][:200] if lb else None
|
| 214 |
+
R["P3_pass"] = bool(a["rc"] == 0 and b["rc"] == 0 and "auto-resume: hub says step" in b["out"]
|
| 215 |
+
and lb)
|
| 216 |
+
print("VERDICT P3_pass=", R["P3_pass"], flush=True)
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
def main():
|
| 220 |
+
ap = argparse.ArgumentParser()
|
| 221 |
+
ap.add_argument("--stage", required=True,
|
| 222 |
+
choices=["p0", "p1", "p2", "p3", "summary", "all"])
|
| 223 |
+
ap.add_argument("--seq-len", type=int, default=2048)
|
| 224 |
+
ap.add_argument("--accum", type=int, default=32)
|
| 225 |
+
ap.add_argument("--tokens", type=int, default=1_000_000_000)
|
| 226 |
+
ap.add_argument("--steps-a", type=int, default=60)
|
| 227 |
+
ap.add_argument("--steps-b", type=int, default=60)
|
| 228 |
+
ap.add_argument("--ckpt-repo", default="Cion-lab/ounce100m-ckptbench-DELETEME")
|
| 229 |
+
ap.add_argument("--out-json", default=WORK + "/preflight.json")
|
| 230 |
+
args = ap.parse_args()
|
| 231 |
+
|
| 232 |
+
R = {}
|
| 233 |
+
if os.path.exists(args.out_json):
|
| 234 |
+
try:
|
| 235 |
+
R = json.load(open(args.out_json))
|
| 236 |
+
except Exception:
|
| 237 |
+
print("existing preflight.json unreadable; starting fresh", flush=True)
|
| 238 |
+
todo = ["p0", "p1", "p2", "p3"] if args.stage == "all" else [args.stage]
|
| 239 |
+
t0 = time.time()
|
| 240 |
+
if "p0" in todo:
|
| 241 |
+
p0_reader(args, R)
|
| 242 |
+
if "p1" in todo:
|
| 243 |
+
p1_hub(args, R)
|
| 244 |
+
if "p2" in todo:
|
| 245 |
+
p2_throughput(args, R)
|
| 246 |
+
if "p3" in todo:
|
| 247 |
+
p3_cold_resume(args, R)
|
| 248 |
+
R["seconds_this_invocation"] = round(time.time() - t0, 1)
|
| 249 |
+
R["gpu_hours_this_invocation"] = round(R["seconds_this_invocation"] / 3600.0, 3)
|
| 250 |
+
R["PASSES"] = {k: R.get(k + "_pass") for k in ("P0", "P1", "P2", "P3")}
|
| 251 |
+
R["GATE_3_READY"] = all(v is True for v in R["PASSES"].values())
|
| 252 |
+
with open(args.out_json, "w") as f:
|
| 253 |
+
json.dump(R, f, indent=1, default=str)
|
| 254 |
+
print("PREFLIGHT_JSON")
|
| 255 |
+
print(json.dumps({"PASSES": R["PASSES"], "GATE_3_READY": R["GATE_3_READY"],
|
| 256 |
+
"gpu_hours_this_invocation": R["gpu_hours_this_invocation"]}))
|
| 257 |
+
print("/PREFLIGHT_JSON")
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
if __name__ == "__main__":
|
| 261 |
+
main()
|