Download kernels/p4_stop_probe.py from Cion-lab/ounce100m-code: direct link, hf CLI and curl.
- Browser
- Download file 18.1 kB
-
https://huggingface.co/Cion-lab/ounce100m-code/resolve/main/kernels/p4_stop_probe.py
- Command line
-
hf download hf://Cion-lab/ounce100m-code/kernels/p4_stop_probe.py
-
curl -L -o p4_stop_probe.py https://huggingface.co/Cion-lab/ounce100m-code/resolve/main/kernels/p4_stop_probe.py
18.1 kB
| # P4 probe: does --stop-after-steps really produce a resumable, Hub-verified segment boundary? | |
| # | |
| # Why this exists and what it buys with ~0.4 GPU-hours: `--stop-after-steps` is the mechanism every Phase 4 | |
| # session ends on, and Gate 3 never exercised it -- P3's legs used `--max-steps`, so each leg *was* a whole | |
| # run. The launcher's plan, the trainer's stop, HubPush's push-verify-pointer-prune order, the resume scan | |
| # that refuses a stale pointer, and the final-push-plus-terminal-pointer path are therefore untested | |
| # together. E-035/4 is precisely a defect in that untested seam, and the rehearsal kernel cannot reach it | |
| # because it needs two T4s. This does, against the published mix and the frozen 22L geometry, in one | |
| # session: leg 1 trains to a mid-run stop step, leg 2 wipes the disk, resumes from the Hub and runs to the | |
| # horizon. | |
| # | |
| # The probe writes to its own checkpoint repo, never to Cion-lab/ounce100m-ckpt: the first real session has | |
| # to find that repo absent, which is the RepoMissing branch E-034 was about. | |
| # | |
| # It carries a second question, the one the user pushed back on (D-017): gradient checkpointing costs 32 % | |
| # of the throughput (9,696 -> 12,792 tok/s, 29.7 h -> 22.5 h) and gives up the memory headroom, sitting at | |
| # 12.25 GB of ~14.56. So this cell runs the *real* geometry at --accum 32, micro 4, WITHOUT checkpointing, | |
| # for 180 steps, through three push/verify/prune cycles and one forced cold resume and the end-of-run | |
| # validation pass -- the three things a 30-step throughput cell cannot show: allocator drift across a few | |
| # hundred steps, the save path's host/GPU copies while the card is 84 % full, and the eval forward pass. | |
| # Peak memory is asserted, not eyeballed. If it holds, the run adopts it as D-018 with a finer push cadence | |
| # as the bounded blast radius; if it does not, D-017 stands and this is the measurement that says so. | |
| import hashlib, json, os, re, shutil, signal, subprocess, sys, threading, time | |
| os.chdir("/kaggle/working") | |
| sys.path.insert(0, "/kaggle/working") | |
| REV = "55b8fc47dd7799bf3fc08b7943421f578bac4a2c" | |
| WANT = { | |
| "ounce100m_credentials.py": ("ounce100m_credentials.py", | |
| "6525f62f03f2d73650a1eb4f70fcb52d1194caad4ca88b2d8bd8fd54f88339b6"), | |
| "shard_dataset.py": ("train/shard_dataset.py", | |
| "f35653bf4c8f2cfe7bb0c2c7835e308505fb5a84b4eb002d0f82f2cada768ca6"), | |
| "hubckpt.py": ("train/hubckpt.py", | |
| "d4b50ed0928c678c94ce6764612a4655f2959c0a7b91d6fd588e8efddff298b2"), | |
| "train_ounce100m.py": ("train/train_ounce100m.py", | |
| "dcb0ac199c0616575423ddea0d6b76d2257c06089d3f73aae0e0359de958ff03"), | |
| } | |
| BASE = "https://huggingface.co/Cion-lab/ounce100m-code/resolve/" + REV | |
| for p, (rp, want) in sorted(WANT.items()): | |
| assert subprocess.run(["curl", "-sfL", f"{BASE}/{rp}", "-o", p]).returncode == 0, ("fetch", rp) | |
| got = hashlib.sha256(open(p, "rb").read()).hexdigest() | |
| assert got == want, ("SHA MISMATCH", rp, got[:16], want[:16]) | |
| print("OK", p, got[:12], flush=True) | |
| import ounce100m_credentials as C | |
| print("creds:", json.dumps(C.install(verify=True)), flush=True) | |
| import hubckpt | |
| from huggingface_hub import HfApi | |
| MIX = "Cion-lab/ounce100m-mix-v1" | |
| # A scratch repo per attempt (dated), never the real run's repo: v1 left `final` at step 20 and v2 was | |
| # correctly refused by the stale-stop guard, so rather than clearing state between attempts each one gets an | |
| # empty repo. That also makes every attempt walk the RepoMissing resume branch (E-034) that session 1 hits. | |
| PROBE = os.environ.get("P4_PROBE_REPO") or ( | |
| "Cion-lab/ounce100m-ckpt-probe-" + time.strftime("%m%d-%H%M", time.gmtime())) | |
| ROOT, RUN = "/kaggle/working/mixroot", "/kaggle/working/run" | |
| # Two modes, one file, so the assertions are literally the same code in both. `smoke` is the user's | |
| # suggestion and it is the right order: 20 steps costs ~12 minutes and answers "does the training code run | |
| # at all, and does one checkpoint survive the push/verify/pointer/prune cycle" -- which is exactly what the | |
| # first run of this probe failed at, in 8.6 seconds, on a malformed torchrun command line (E-037). `soak` | |
| # is the 180-step memory question, and it is only worth 1.6 GPU-hours once the mechanics are known to work. | |
| MODE = os.environ.get("P4_PROBE_MODE", "soak") | |
| if MODE not in ("smoke", "soak"): | |
| # A typo here would otherwise run the 1.6-hour soak when a 13-minute smoke was asked for. | |
| raise SystemExit("P4_PROBE_MODE must be smoke or soak, got %r" % MODE) | |
| TLOG = "/kaggle/working/tlogs" | |
| # torchrun takes its first positional as the SCRIPT, not a command: passing sys.executable made | |
| # it compile the Python binary (E-037). --redirects is a per-rank bitmask into --log-dir (1=stderr, | |
| # 2=stdout, 3=both) and --tee repeats the same streams to this process, so the loss curve §5 requires | |
| # watching stays live *and* each rank's traceback is on disk for diag() below. | |
| TORCHRUN = ["torchrun", "--nproc_per_node=2", "--redirects", "3", "--tee", "3", "--log-dir", TLOG] | |
| TPS = 262144 # the run's real step shape, in tokens | |
| # `--tee` prefixes every forwarded line with the worker name, and this file parses those lines: v5's | |
| # harvest() found no `RUN_JSON ` because the trainer's rank-0 output arrived as | |
| # "[default0]: RUN_JSON {...}", so every step/token/param check reported false on a leg that had actually | |
| # passed, and the failure looked like a trainer bug rather than a probe bug. Strip it once, on the way in. | |
| TEE = re.compile(r"^\[(?:default|rank|worker)\d*\]:\s*") | |
| if MODE == "smoke": | |
| STEPS, PUSH_EVERY, STOP1, VAL = 20, 10, 10, 200000 | |
| T_LEG1, T_LEG2, T_FRESH = 1200, 900, 1500 | |
| else: | |
| STEPS, PUSH_EVERY, STOP1, VAL = 180, 60, 120, 2000000 | |
| # Sized from the model rather than guessed: 120 steps at ~20.5 s + build + three pushes + eval is | |
| # ~3,000 s, and 60 steps + cold pull + push + eval is ~1,700 s. The three ceilings must also fit under | |
| # the notebook's own session timeout (10,800 s for the soak) with room for the model build, because a | |
| # stage that outlives the container reports nothing at all (review point B5). | |
| T_LEG1, T_LEG2, T_FRESH = 3900, 2400, 1200 | |
| TOKENS = STEPS * TPS | |
| GATE_PEAK = (MODE != "smoke") # 20 steps says nothing about allocator drift | |
| T0 = time.time() | |
| def run(argv, label, timeout): | |
| print("=== " + label, flush=True) | |
| t0 = time.time() | |
| e = dict(os.environ); e["PYTHONPATH"] = "/kaggle/working" | |
| e["PYTHONUNBUFFERED"] = "1" # the -u that torchrun cannot carry | |
| e["NCCL_DEBUG"] = "WARN" # a rank that dies in a collective says so here and nowhere else | |
| e["TORCH_CPP_LOG_LEVEL"] = "WARNING" | |
| p = subprocess.Popen(argv, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True, | |
| env=e, bufsize=1, start_new_session=True) | |
| killed = [] | |
| def _kill(): | |
| killed.append(True) | |
| try: | |
| os.killpg(os.getpgid(p.pid), signal.SIGTERM) | |
| except Exception: | |
| p.kill() | |
| def _hard(): | |
| killed.append(True) | |
| try: | |
| os.killpg(os.getpgid(p.pid), signal.SIGKILL) | |
| except Exception: | |
| p.kill() | |
| timer = threading.Timer(timeout, _kill) | |
| timer.daemon = True | |
| timer.start() | |
| hard = threading.Timer(timeout + 90, _hard) | |
| hard.daemon = True | |
| hard.start() | |
| keep, lines = [], [] | |
| try: | |
| for line in p.stdout: | |
| line = TEE.sub("", line.rstrip("\n"), count=1) | |
| lines.append(line) | |
| if line.startswith(("CKPT ", "RUN_JSON ", "resume from", "auto-resume", "precision:", | |
| "params:", "mix:", "checkpoint hub target", "segment boundary", | |
| "validation skipped", "latest.json", "TRAIN DONE", "validation loss", | |
| "[rank ", "past the ", "peak_stats:", "Traceback", "Error")): | |
| print(" KEY>", line[:300], flush=True) | |
| keep.append(line) | |
| del keep[:-400] | |
| finally: | |
| timer.cancel() | |
| hard.cancel() | |
| rc = p.wait() | |
| if killed: | |
| print(" TIMEOUT after %d s" % timeout, flush=True) | |
| rc = -9 | |
| if rc != 0: | |
| # Probe v6 died on the early-stop leg and everything that could have said *where* was in `lines` | |
| # and never printed: the last-40-line tail was all torchrun's summary, and the failing rank's own | |
| # traceback came a hundred lines earlier. Grep the whole capture, then show the tail. | |
| pat = ("[rank", "Traceback", "Error", "error", "Exception", "assert", "exitcode", "Signal", | |
| "SystemExit", "refusing", "skipped", "CUDA", "NCCL", "out of memory", 'File "', "line ", | |
| "raise ", "FileNotFoundError", "RuntimeError") | |
| hits = [l for l in lines if any(q in l for q in pat)] | |
| print(" FAILURE LINES (%d of %d):" % (len(hits), len(lines)), flush=True) | |
| for l in hits[-45:]: | |
| print(" !", l[:300], flush=True) | |
| print(" TAIL:\n" + "\n".join(keep)[-2500:], flush=True) | |
| print("%s_RC %s seconds %.1f elapsed %.0f" % (label, rc, time.time() - t0, time.time() - T0), | |
| flush=True) | |
| return rc, "\n".join(lines) | |
| rc, out = run([sys.executable, "-c", | |
| "import sys; sys.path.insert(0, '/kaggle/working')\n" | |
| "from huggingface_hub import snapshot_download\n" | |
| "p = snapshot_download(repo_id='%s', repo_type='dataset',\n" | |
| " local_dir='/kaggle/working/mixroot', max_workers=4)\n" | |
| "print('mix at', p)\n" % MIX], "FETCH_MIX", T_FRESH) | |
| if rc != 0: | |
| raise SystemExit("VERDICT P4PROBE_STOP could not fetch the published mix") | |
| man = json.load(open(os.path.join(ROOT, "manifest.json"))) | |
| print("mix", man["n_shards"], "shards", format(int(man["total_tokens"]), ","), "tokens", | |
| "| probe repo", PROBE, flush=True) | |
| common = ["train_ounce100m.py", "--root", ROOT, "--out", RUN, | |
| "--hub-repo", PROBE, "--prune", "--seq-len", "1024", "--attn", "eager", | |
| # The flag under test. Passing both --grad-ckpt and --no-grad-ckpt would leave it to argparse's | |
| # last-wins ordering, which is not a thing to be ambiguous about in a probe of this recipe. | |
| "--no-grad-ckpt", | |
| "--micro-batch", "4", "--accum", "32", | |
| "--tokens", str(TOKENS), "--lr", "6e-4", | |
| "--push-every-steps", str(PUSH_EVERY), "--val-tokens", str(VAL), "--log-every", "5", | |
| "--resume", "auto"] | |
| def diag(tag): | |
| """torchrun's ChildFailedError said `error_file: <N/A>` and printed no traceback, which is useless for | |
| a rank-1-only failure. --redirects writes each rank's own stdout/stderr into the log dir, so print | |
| those on a failed leg: the exception is in there, and guessing at it costs a session.""" | |
| # v4's diag found no *.log at all, so list the tree as well as tailing it: if torchrun named the | |
| # files something else, this says so instead of printing nothing and leaving me guessing. | |
| rc, out = run(["bash", "-c", | |
| 'ls -R "%s" 2>&1 | head -40; ' | |
| 'for f in $(find "%s" -type f 2>/dev/null | head -8); do echo "==== $f"; ' | |
| 'tail -70 "$f"; done' % (TLOG, TLOG)], "DIAG_" + tag, 180) | |
| # run() only echoes lines it recognises, which made v4's and v6's diag print *nothing* while the | |
| # rank logs sat right there on disk. Say what was found, even if it is a traceback shape we do not | |
| # have a filter word for. | |
| print(" DIAG %s (%d chars, rc %s):\n%s" % (tag, len(out or ""), rc, (out or "")[-4000:]), | |
| flush=True) | |
| return out | |
| print("PROBE_REPO", PROBE, flush=True) | |
| shutil.rmtree(RUN, ignore_errors=True) | |
| rc1, o1 = run(TORCHRUN + common + ["--stop-after-steps", str(STOP1)], "LEG1_STOP_EARLY", T_LEG1) | |
| if rc1 != 0: | |
| diag("leg1") | |
| api = HfApi(token=os.environ["HF_TOKEN"]) | |
| try: | |
| listed = sorted(hubckpt.hub_listing(PROBE, "dataset", token=os.environ["HF_TOKEN"])) | |
| except Exception as e: | |
| listed = ["<listing failed: %s>" % type(e).__name__] | |
| pushed = sorted({k.split("/")[1] for k in listed if k.startswith("ckpt/checkpoint-") | |
| and len(k.split("/")) > 1}) | |
| ptr = hubckpt.latest_pointer(PROBE, token=os.environ["HF_TOKEN"]) | |
| print("LEG1 pushed:", pushed, "pointer:", {k: ptr.get(k) for k in ("step", "path_in_repo", "error")}, | |
| flush=True) | |
| shutil.rmtree(RUN, ignore_errors=True) # force the cold-resume path (E-029's shape) | |
| rc2, o2 = run(TORCHRUN + common + ["--stop-after-steps", "0"], "LEG2_TO_HORIZON", T_LEG2) | |
| if rc2 != 0: | |
| diag("leg2") | |
| try: | |
| listed2 = sorted(hubckpt.hub_listing(PROBE, "dataset", token=os.environ["HF_TOKEN"])) | |
| except Exception as e: | |
| listed2 = ["<listing failed: %s>" % type(e).__name__] | |
| ptr2 = hubckpt.latest_pointer(PROBE, token=os.environ["HF_TOKEN"]) | |
| def harvest(text, tag): | |
| for line in (text or "").splitlines(): | |
| if line.startswith(tag): | |
| try: | |
| return json.loads(line[len(tag):].strip()) | |
| except Exception: | |
| return {"unparsed": line[:200]} | |
| return {} | |
| rj1, rj2 = harvest(o1, "RUN_JSON "), harvest(o2, "RUN_JSON ") | |
| res = { | |
| "leg1_rc": rc1, "leg2_rc": rc2, | |
| "leg1": {"final_step": rj1.get("final_step"), "segment_stop": rj1.get("segment_stop"), | |
| "tokens_consumed": rj1.get("tokens_consumed"), "tok_per_s": rj1.get("tok_per_s"), | |
| "params": rj1.get("params"), "peak_gpu_gb": rj1.get("peak_gpu_gb"), | |
| "tok_per_s": rj1.get("tok_per_s"), "final_loss": rj1.get("final_loss"), | |
| "peak_alloc_gb": rj1.get("peak_gpu_gb"), "peak_alloc_gb_max_rank": | |
| rj1.get("peak_gpu_gb_max_rank"), "peak_reserved_gb_max_rank": | |
| rj1.get("peak_reserved_gb_max_rank"), "val_error": rj1.get("val_error"), | |
| "pushed": pushed, "pointer": {k: ptr.get(k) for k in ("step", "path_in_repo")}}, | |
| "leg2": {"final_step": rj2.get("final_step"), "segment_stop": rj2.get("segment_stop"), | |
| "tokens_consumed": rj2.get("tokens_consumed"), "val_ppl": rj2.get("val_ppl"), | |
| "peak_gpu_gb": rj2.get("peak_gpu_gb"), "peak_gpu_gb_max_rank": | |
| rj2.get("peak_gpu_gb_max_rank"), "peak_reserved_gb_max_rank": | |
| rj2.get("peak_reserved_gb_max_rank"), "val_error": rj2.get("val_error"), | |
| "tok_per_s": rj2.get("tok_per_s"), | |
| "final_loss": rj2.get("final_loss"), | |
| "pointer_after": {k: ptr2.get(k) for k in ("step", "path_in_repo")}, | |
| "has_final": any(k.startswith("final/") for k in listed2)}, | |
| "expect": {"steps_planned": STEPS, "push_every": PUSH_EVERY, "stop1": STOP1, | |
| "tokens_at_stop1": STOP1 * TPS, "tokens_at_horizon": STEPS * TPS, "grad_ckpt": False}, | |
| } | |
| res["checks"] = { | |
| "leg1_stopped_at_the_stop_step": rj1.get("final_step") == STOP1, | |
| "leg1_reported_a_segment": rj1.get("segment_stop") is True, | |
| # Integer compare: sorted() over names would order "checkpoint-120" before "checkpoint-60". | |
| "leg1_pushed_every_interval": sorted(int(x.split("-")[1]) for x in pushed) == | |
| list(range(PUSH_EVERY, STOP1 + 1, PUSH_EVERY)), | |
| "leg1_pointer_at_stop": ptr.get("step") == STOP1, | |
| "leg2_resumed_and_finished": rj2.get("final_step") == STEPS, | |
| "leg2_was_not_a_segment": rj2.get("segment_stop") is False, | |
| "leg2_pushed_final": any(k.startswith("final/") for k in listed2), | |
| "all_steps_present": sorted(int(x.split("-")[1]) for x in | |
| {k.split("/")[1] for k in listed2 if k.startswith("ckpt/checkpoint-")}) | |
| == list(range(PUSH_EVERY, STEPS + 1, PUSH_EVERY)), | |
| "leg2_pointer_terminal": ptr2.get("step") == STEPS and ptr2.get("path_in_repo") == "final", | |
| "tokens_match_the_arithmetic": rj1.get("tokens_consumed") == STOP1 * TPS | |
| and rj2.get("tokens_consumed") == STEPS * TPS, | |
| "params_are_the_frozen_model": rj1.get("params") == 106194240, | |
| "checkpointing_really_off": rj1.get("grad_ckpt") is False, | |
| # What gets gated is leg 1, because it is the only leg whose number means "training". `max_memory_reserved` | |
| # is the high-water mark since the process started, and leg 2 runs the validation pass before sampling | |
| # it, so leg 2's figure is training-plus-eval -- on smoke that came out 13.66 GB reserved against 12.7 | |
| # allocated, while the segment that never evaluated is the state the run sits in for 3,814 steps. Gating | |
| # on the eval-inflated number would fail D-018 for a state the run never occupies. The other rank's peak | |
| # and the reserved bytes are what an OOM is actually about, hence `max_rank` rather than rank 0 alone. | |
| "peak_memory_during_training_below_13_6_gb": (not GATE_PEAK) or ( | |
| (rj1.get("peak_reserved_gb_max_rank") or rj1.get("peak_gpu_gb_max_rank") or 99) <= 13.6), | |
| # Leg 1 skips validation by design (E-040's fix; the crash it was written for turned out to be E-044, | |
| # but the skip still stands -- a mid-run PPL point is not worth an untested path in a billed session), | |
| # and leg 2 must actually run it, because that is where the report's PPL comes from. | |
| "leg1_skipped_validation_by_design": rj1.get("val_skipped") is True and rj1.get("val_ppl") is None, | |
| "leg2_validation_actually_ran": (rj2.get("val_ppl") is not None and rj2.get("val_error") is None | |
| and rj2.get("val_skipped") is False), | |
| "no_nan_and_loss_moved": (rj1.get("final_loss") or 1e9) < 11.0 | |
| and (rj2.get("final_loss") or 1e9) < 11.0, | |
| } | |
| res["PROBE_PASSED"] = all(res["checks"].values()) and rc1 == 0 and rc2 == 0 | |
| print("PROBE_MODE", MODE, "steps", STEPS, "push_every", PUSH_EVERY, | |
| "stop1", STOP1, "peak_gate", GATE_PEAK, flush=True) | |
| print("PROBE_JSON_BEGIN") | |
| print(json.dumps(res, indent=1, default=str)) | |
| print("PROBE_JSON_END") | |
| print("VERDICT P4PROBE", "PASS" if res["PROBE_PASSED"] else "FAIL", | |
| [k for k, v in res["checks"].items() if not v], "seconds", round(time.time() - T0, 1), flush=True) | |
| raise SystemExit(0 if res["PROBE_PASSED"] else 5) | |