Cion-lab commited on
Commit
e7ee32a
·
verified ·
1 Parent(s): 6ddf1d2

train: Gate 3 preflight driver (reader, hub cycle, T1 throughput, cold resume)

Browse files
Files changed (1) hide show
  1. 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()