Cion-lab commited on
Commit
758d461
·
verified ·
1 Parent(s): 66e6f6e

Upload train/preflight.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. train/preflight.py +75 -25
train/preflight.py CHANGED
@@ -311,6 +311,21 @@ def torchrun(args, extra, timeout=7200):
311
  + extra, timeout=timeout, label="torchrun " + " ".join(extra[:6]))
312
 
313
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
314
  def p2_throughput(args, R):
315
  """T1: 20L vs 22L at two sequence lengths, on the real data, for enough steps that the number is
316
  steady-state. The main run's whole schedule is division by this number. 22L/576 is the frozen shape;
@@ -358,41 +373,76 @@ def p2_throughput(args, R):
358
 
359
 
360
  def p3_cold_resume(args, R):
361
- """Short run -> wipe every local trace -> resume with --resume auto, which must recover the exact
362
- sample position from the Hub. Loss continuity is the pass condition: a restart that silently
363
- re-initialised would jump the loss."""
 
 
 
 
 
 
 
 
 
 
364
  common = ["--seq-len", str(args.seq_len), "--tokens", str(args.tokens),
365
- "--accum", str(args.accum), "--micro-batch", "2",
366
  "--hub-repo", args.ckpt_repo, "--prune", "--log-every", "5"]
367
- a = torchrun(args, common + ["--max-steps", str(args.steps_a), "--out", WORK + "/p3"],
368
- timeout=9000)
369
- # erase the instance's memory of the run: local checkpoints, the run dir, and the downloaded mix?
370
- # No -- the mix is the dataset and a real interruption keeps it. Only the run state goes.
371
- wipe = sh(["bash", "-c", f"rm -rf {WORK}/p3 {WORK}/run; df -h {WORK} | tail -1"],
372
- label="P3 wipe local run state")
373
- b = torchrun(args, common + ["--max-steps", str(args.steps_a + args.steps_b),
374
- "--out", WORK + "/p3b", "--resume", "auto"], timeout=9000)
375
- R["P3"] = {"first": {"rc": a["rc"], "tail": a["out"][-2500:]},
376
- "wipe": wipe["out"][-400:],
377
- "resumed": {"rc": b["rc"], "tail": b["out"][-2500:]}}
378
- la = [l for l in a["out"].splitlines() if "'loss'" in l or "loss=" in l]
379
- lb = [l for l in b["out"].splitlines() if "'loss'" in l or "loss=" in l]
380
- R["P3_loss_last_before"] = la[-1][:200] if la else None
381
- R["P3_loss_first_after"] = lb[0][:200] if lb else None
382
- R["P3_pass"] = bool(a["rc"] == 0 and b["rc"] == 0 and "auto-resume: hub says step" in b["out"]
383
- and lb)
384
- print("VERDICT P3_pass=", R["P3_pass"], flush=True)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
385
 
386
 
387
  def main():
388
  ap = argparse.ArgumentParser()
389
  ap.add_argument("--stage", required=True,
390
  choices=["pargs", "p0", "p1", "p2", "p3", "summary", "all"])
391
- ap.add_argument("--seq-len", type=int, default=2048)
392
  ap.add_argument("--accum", type=int, default=32)
393
  ap.add_argument("--tokens", type=int, default=1_000_000_000)
394
- ap.add_argument("--steps-a", type=int, default=60)
395
- ap.add_argument("--steps-b", type=int, default=60)
396
  ap.add_argument("--ckpt-repo", default="Cion-lab/ounce100m-ckptbench-DELETEME")
397
  ap.add_argument("--out-json", default=WORK + "/preflight.json")
398
  args = ap.parse_args()
 
311
  + extra, timeout=timeout, label="torchrun " + " ".join(extra[:6]))
312
 
313
 
314
+ # The checkpoint rig is a timing fixture, not an artifact, so it is deleted on both ends of the run.
315
+ _DEL = """
316
+ import os, sys
317
+ sys.path.insert(0, "/kaggle/working")
318
+ import ounce100m_credentials
319
+ ounce100m_credentials.install()
320
+ from huggingface_hub import HfApi
321
+ try:
322
+ HfApi().delete_repo(repo_id="%s", repo_type="dataset", token=os.environ["HF_TOKEN"])
323
+ print("hub repo deleted", flush=True)
324
+ except Exception as e:
325
+ print("hub repo:", type(e).__name__, str(e)[:120], flush=True)
326
+ """
327
+
328
+
329
  def p2_throughput(args, R):
330
  """T1: 20L vs 22L at two sequence lengths, on the real data, for enough steps that the number is
331
  steady-state. The main run's whole schedule is division by this number. 22L/576 is the frozen shape;
 
373
 
374
 
375
  def p3_cold_resume(args, R):
376
+ """Three legs, each ending in one checkpoint, each started from the Hub with nothing local but the mix.
377
+
378
+ This is the test that matters most and the one that cannot be faked. Leg A trains and pushes. Legs B
379
+ and C each `--resume auto` after their local run directory has been deleted, so recovery is forced
380
+ through the Hub -- which is what every real interruption looks like. Two sequential resumes because §5
381
+ asks for more than one, and the cursor must advance monotonically and never re-read.
382
+ """
383
+ # The Hub side has to be wiped too: a latest.json left by an earlier attempt would make leg A a
384
+ # mid-run resume, and then the monotonic-cursor check below would pass for the wrong reason.
385
+ R["P3_hub_preclean"] = (sh(["python", "-c", _DEL % args.ckpt_repo],
386
+ label="P3 hub preclean")["out"] or "").strip()[-200:]
387
+ # Same geometry as the main run (D-011): 4 x 1024 x 32 accum x 2 cards = 262,144 tokens/step, so P3
388
+ # exercises the real step, the real checkpoint size and the real cursor, not a cheaper stand-in.
389
  common = ["--seq-len", str(args.seq_len), "--tokens", str(args.tokens),
390
+ "--accum", str(args.accum), "--micro-batch", "4",
391
  "--hub-repo", args.ckpt_repo, "--prune", "--log-every", "5"]
392
+ legs, cursors, losses = [], [], []
393
+ for i in range(3):
394
+ if i:
395
+ # the wipe IS the test: forget everything the last leg left on this instance except the dataset
396
+ w = sh(["bash", "-c", "rm -rf " + WORK + "/p3* " + WORK + "/run && df -h " + WORK
397
+ + " | tail -1"], label=f"P3 leg {chr(65 + i)} wipe")
398
+ R[f"P3_wipe_{chr(65 + i)}"] = (w["out"] or "").strip()[-200:]
399
+ steps = args.steps_a * (i + 1)
400
+ # One checkpoint per leg: pushing exactly at the leg's last step is what proves the cursor was
401
+ # written, verified on the Hub, and then read back cold by the next leg.
402
+ r = torchrun(args, common + ["--max-steps", str(steps), "--push-every-steps", str(steps),
403
+ "--resume", "auto",
404
+ "--out", WORK + f"/p3_{'abc'[i]}"], timeout=9000)
405
+ out = r["out"]
406
+ cur = None
407
+ for line in out.splitlines():
408
+ if line.startswith("resume from") and "cursor=" in line:
409
+ cur = line.split("cursor=", 1)[1][:220]
410
+ if line.startswith("CKPT "):
411
+ cursors.append(line[:260])
412
+ if "prior loss at resume:" in line:
413
+ losses.append(line[:120])
414
+ leg = {"rc": r["rc"], "seconds": r["seconds"], "max_steps": steps,
415
+ "auto_resume": ("auto-resume: hub says step" in out) or (i == 0),
416
+ "cursor_line": cur,
417
+ "tail": out[-1200:] if r["rc"] else None,
418
+ "err": (r["err"] or "")[-900:] if r["rc"] else None}
419
+ legs.append(leg)
420
+ print(f"VERDICT P3 leg {'abc'[i]} rc={r['rc']} auto_resume={leg['auto_resume']}", flush=True)
421
+ if r["rc"] != 0:
422
+ break
423
+ R["P3"] = {"legs": legs, "ckpt_lines": cursors, "prior_loss_lines": losses}
424
+ seen = [int(s.split("samples=")[1].split()[0].replace(",", "")) for s in cursors
425
+ if "samples=" in s]
426
+ R["P3_cursor_sequence"] = seen
427
+ ok = (len(legs) == 3 and all(l["rc"] == 0 for l in legs)
428
+ and all(l["auto_resume"] for l in legs[1:])
429
+ and len(cursors) >= 3 and len(seen) >= 3
430
+ and seen == sorted(seen) and len(set(seen)) == len(seen))
431
+ R["P3_pass"] = bool(ok)
432
+ print("VERDICT P3_pass=", R["P3_pass"], "cursors=", seen, flush=True)
433
+ R["P3_hub_postclean"] = (sh(["python", "-c", _DEL % args.ckpt_repo],
434
+ label="P3 hub postclean")["out"] or "").strip()[-200:]
435
 
436
 
437
  def main():
438
  ap = argparse.ArgumentParser()
439
  ap.add_argument("--stage", required=True,
440
  choices=["pargs", "p0", "p1", "p2", "p3", "summary", "all"])
441
+ ap.add_argument("--seq-len", type=int, default=1024) # D-011: the config the main run will use
442
  ap.add_argument("--accum", type=int, default=32)
443
  ap.add_argument("--tokens", type=int, default=1_000_000_000)
444
+ ap.add_argument("--steps-a", type=int, default=20,
445
+ help="P3 leg stride; legs run to 20, 40 and 60 cumulative steps")
446
  ap.add_argument("--ckpt-repo", default="Cion-lab/ounce100m-ckptbench-DELETEME")
447
  ap.add_argument("--out-json", default=WORK + "/preflight.json")
448
  args = ap.parse_args()