ounce100m-code / kernels /p4_stop_probe.py
Cion-lab's picture
p4_stop_probe: E-044 repin (rank-0 stop assertions; memory gate moved onto the training-only leg)
14e982b verified
Raw History Blame Contribute Delete
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)