ounce100m-code / kernels /phase4_prep.py
Cion-lab's picture
phase4_prep: REV must be at or after the launcher it sha-pins -- point at c237b478
6a668e6 verified
Raw History Blame Contribute Delete
21.2 kB
# Phase 4 rehearsal on a CPU instance: every check a session can fail, run where failing costs no quota.
#
# The first GPU session of the run bills ~6.5 hours before it is provably healthy, and E-031/E-034 are both
# exactly the shape of bug that only shows up mid-run or at the very first launch. So this kernel does, on
# free CPU: fetches the pinned code and asserts its sha256s, resolves credentials, reads the checkpoint
# pointer, downloads the PUBLISHED mix from the Hub, hashes every shard against the manifest, rehearses the
# session planner across the whole schedule, and proves resume-equivalence of the reader at the boundaries
# the run actually crosses.
#
# It ends with the same PLAN_JSON the GPU session will print, so the numbers can be compared before and
# after rather than trusting that the code did not change in between.
import hashlib, json, os, subprocess, sys, time
os.chdir("/kaggle/working")
sys.path.insert(0, "/kaggle/working")
REV = os.environ.get("OUNCE100M_REV") or "c237b478076f68da43dabd3403daa5639286899e"
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"),
"phase4_session.py": ("kernels/phase4_session.py",
"5fb14bc3597d999a21772a66b6f409f01f722019dc78da313feddb2f9c769c48"),
}
BASE = "https://huggingface.co/Cion-lab/ounce100m-code/resolve/" + REV
for p, (rp, want) in sorted(WANT.items()):
r = subprocess.run(["curl", "-sfL", BASE + "/" + rp, "-o", p], capture_output=True, text=True)
assert r.returncode == 0, ("fetch failed", rp, (r.stderr or "")[:200])
got = hashlib.sha256(open(p, "rb").read()).hexdigest()
assert got == want, ("SHA MISMATCH", rp, got, want)
print("OK", p, got[:12], flush=True)
print("REV_OK", REV, flush=True)
import ounce100m_credentials as C
print("creds:", json.dumps(C.install(verify=True)), flush=True)
T0 = time.time()
def run(argv, label, timeout, env=None):
print("=== " + label, flush=True)
t0 = time.time()
e = dict(os.environ); e["PYTHONPATH"] = "/kaggle/working"
e.update(env or {})
import signal as _sig
import threading
p = subprocess.Popen(argv, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
env=e, bufsize=1, start_new_session=True)
killed = []
def _kill():
# Kill the group, not the child: a staged `python -c snapshot_download` orphaned by a timeout
# keeps writing into mixroot while the next stage reads it.
killed.append(True)
try:
os.killpg(os.getpgid(p.pid), _sig.SIGTERM)
except Exception:
p.kill()
timer = threading.Timer(timeout, _kill)
timer.daemon = True
timer.start()
out = []
try:
for line in p.stdout:
out.append(line.rstrip("\n"))
print(" |", line[:300], flush=True)
finally:
timer.cancel()
rc = p.wait()
if killed:
print(" TIMEOUT after %d s, killed" % timeout, flush=True)
rc = -9
print("%s_RC %s seconds %.1f elapsed %.0f" % (label, rc, time.time() - t0, time.time() - T0),
flush=True)
return rc, "\n".join(out)
# Stage 1 -- the launcher itself, in rehearsal mode. On a CPU box torchrun is never reached: PREP_ONLY
# stops right after the pointer read, the mix download and the planner. Everything upstream of the first
# optimiser step is exercised with the real files and the real Hub.
PLANENV = {"PLANNING_TOK_PER_S": os.environ.get("PLANNING_TOK_PER_S", "9358"),
"SESSION_GPU_HOURS": os.environ.get("SESSION_GPU_HOURS", "6.9"),
"QUOTA_LEFT_HOURS": os.environ.get("QUOTA_LEFT_HOURS", "26.6"),
"MIN_FREE_GB": "6"}
rc1, out1 = run([sys.executable, "phase4_session.py", REV], "PREP_LAUNCHER", 3000,
env=dict(PLANENV, PHASE4_PREP_ONLY="1"))
if rc1 != 0:
print("VERDICT PREP_STOP launcher rehearsal rc", rc1, flush=True)
raise SystemExit(2)
# The launcher prints the argument list it *would* hand to torchrun, so stage 3 can parse exactly those
# flags against the pinned trainer. Duplicating them here would test a copy rather than the real thing.
ARGV_LINE = next((l for l in (out1 or "").splitlines() if l.startswith("TRAIN_ARGV ")), None)
assert ARGV_LINE, ("PREP_STOP: the launcher printed no TRAIN_ARGV -- its torchrun command line is "
"untested, which is how E-037 reached a billed session")
json.dump(json.loads(ARGV_LINE[len("TRAIN_ARGV "):]), open("train_argv.json", "w"))
print("TRAIN_ARGV captured", len(json.loads(ARGV_LINE[len("TRAIN_ARGV "):])), "tokens", flush=True)
# Stage 2 -- the published mix, byte for byte, and the reader at the boundaries the run will actually
# cross. `snapshot_download` proves the repo is anonymously readable; this proves the bytes are the ones
# the manifest certifies and that a resume at step X reads the same data an uninterrupted run would have.
CHECK = """
import hashlib, json, os, sys
sys.path.insert(0, "/kaggle/working")
ROOT = "/kaggle/working/mixroot"
man = json.load(open(os.path.join(ROOT, "manifest.json")))
tok_s = int(man["total_tokens"])
bad, nb = [], 0
for grp in ("shards", "val_shards"):
for rec in man.get(grp, []):
p = os.path.join(ROOT, rec["file"])
if not os.path.exists(p):
bad.append((rec["file"], "absent")); continue
h = hashlib.sha256()
with open(p, "rb") as fh:
for chunk in iter(lambda: fh.read(1 << 22), b""):
h.update(chunk)
nb += 1
if h.hexdigest() != rec["sha256"]:
bad.append((rec["file"], "sha " + h.hexdigest()[:12] + " != " + rec["sha256"][:12]))
print("SHASHED", nb, "files", "total_tokens", format(tok_s, ","), "mismatches", len(bad), flush=True)
for b in bad[:10]:
print(" BAD", b, flush=True)
assert not bad, "the published mix does not match its own manifest"
import phase4_session as S # constants and plan() only -- main() is behind __name__
import shard_dataset as SD
SEQ, steps = S.SEQ_LEN, S.HORIZON_STEPS
seqs_per_step = S.MICRO_BATCH * S.WORLD * S.ACCUM # the launcher's own geometry, not retyped
tps = seqs_per_step * SEQ
assert tps == S.TOKENS_PER_STEP, f"prep disagrees with the launcher about tokens/step: {tps}"
assert steps == int(S.TOKENS / tps), f"horizon disagrees: {steps}"
store = SD.PackedTokenStore(ROOT)
n_samples = store.total_tokens // SEQ
fingerprint = SD.mix_fingerprint(ROOT, man)
# The manifest is the corpus definition, and three different things have to agree about it: the record
# list, the counts it claims, and the bytes the reader actually mapped. A rebuild that truncated the list
# while leaving total_tokens alone would otherwise certify 1.11 B tokens of a smaller corpus (review).
if len(man["shards"]) != int(man.get("n_shards") or -1) or len(store.shards) != len(man["shards"]):
raise AssertionError(f"shard records {len(man['shards'])} vs n_shards {man.get('n_shards')} vs "
f"files opened {len(store.shards)}")
summed = sum(int(s["tokens"]) for s in man["shards"])
if summed != tok_s or store.total_tokens != summed:
raise AssertionError(f"manifest does not add up: records sum {summed}, total_tokens says {tok_s}, "
f"bytes the reader mapped say {store.total_tokens}")
print("READER", format(store.total_tokens, ","), "tokens in", len(store.shards), "shards ->",
format(n_samples, ","), "windows of", SEQ, "|", steps, "steps x", tps, "tok =",
format(steps * tps, ","), "tokens | fingerprint", fingerprint, flush=True)
assert n_samples >= steps * seqs_per_step, "not enough windows for the horizon"
# The schedule the run will actually live through: drive the real planner from step 0 until it reports the
# horizon, and require every leg to start where the previous one stopped. This is the check that catches a
# planner which stops one interval short of the end -- the last four steps of the run, unreachable, while
# every later session fails its own budget test for want of a whole interval.
legs, start = [], 0
while start < steps and len(legs) < 40:
pl = S.plan(start, 15.0)
if not pl["planned_steps"]:
raise AssertionError(f"the planner refuses to start at step {start} with "
f"{pl['usable_hours']} usable hours -- the run can never finish")
end = steps if pl["stop_after_steps"] == 0 else pl["stop_after_steps"]
assert end > start, (start, end)
need = (steps - start) * seqs_per_step # the trainer's own guard, at this start step
avail = n_samples - start * seqs_per_step
assert avail >= need, (start, need, avail)
store.tokens((end * seqs_per_step - 1) * SEQ, SEQ) # last window this leg reads must be addressable
legs.append((start, end))
print("LEG %5d -> %5d pct %5.1f -> %5.1f need %9d avail %9d" % (
start, end, 100.0 * start / steps, 100.0 * end / steps, need, avail), flush=True)
start = end
print("SESSIONS", len(legs), "ends_at", legs[-1][1], "horizon", steps, "contiguous",
[l[0] for l in legs[1:]] == [l[1] for l in legs[:-1]], flush=True)
assert legs[-1][1] == steps and [l[0] for l in legs[1:]] == [l[1] for l in legs[:-1]], legs
# Resume-equivalence through the cursor the trainer would actually write, at boundaries the run crosses:
# the first checkpoint, a mid-session one, and the last interval before the horizon. Building the resumed
# dataset from load_cursor() rather than a hand-computed offset is the point -- the derivation from
# step to samples_consumed is part of what is being tested.
full = SD.make_dataset(store, SEQ, 20260919, start_sample=0)
order = SD.perm_sha(n_samples, 20260919)
assert full.order_sha == order, "two derivations of the order identity disagree"
for boundary in sorted({S.CKPT_EVERY, legs[len(legs) // 2][0], steps - S.CKPT_EVERY}):
d = "/tmp/cur_%d" % boundary
os.makedirs(d, exist_ok=True)
SD.save_cursor(d, SD.cursor_dict(SEQ, boundary * seqs_per_step, 20260919, fingerprint,
step=boundary, order_sha=order))
c = SD.load_cursor(d)
assert c["step"] * seqs_per_step == c["samples_consumed"], c
assert c["dataset_files_sha"] == fingerprint, c
assert c["shuffle_perm_sha"] == order, c
s0 = c["samples_consumed"]
resumed = SD.make_dataset(store, SEQ, c["shuffle_seed"], start_sample=s0)
assert resumed.order_sha == order, "a resume lost the visit order"
assert len(resumed) == len(full) - s0, (len(resumed), len(full), s0)
same = True
for j in (0, 1, len(resumed) // 2, len(resumed) - 1):
a, b = full[s0 + j], resumed[j]
same = same and a["input_ids"].tolist() == b["input_ids"].tolist() \
and a["labels"].tolist() == b["labels"].tolist()
print("RESUME_EQUIV step", boundary, "start_sample", s0, "len", len(resumed), "identical", same,
flush=True)
assert same, boundary
print("CHECK_OK windows", format(n_samples, ","), "shards", len(store.shards), flush=True)
"""
open("prepcheck.py", "w").write(CHECK)
rc2, _ = run([sys.executable, "prepcheck.py"], "CHECK_PUBLISHED_MIX", 2400, env=dict(PLANENV))
# Stage 3 -- the reader's cursor format round-trips, and the trainer imports and prints --help on this
# image. Both are cheap and both have bitten before (E-028: a build script that could not import its
# sibling module surfaced only after a GPU session had already been billed).
rc3, _ = run([sys.executable, "-c",
"import sys; sys.path.insert(0, '/kaggle/working')\n"
"import shard_dataset as SD, json, os\n"
"os.makedirs('/tmp/cur', exist_ok=True)\n"
"c = SD.cursor_dict(1024, 97536, 20260919, dataset_files_sha='ab' * 32, step=381)\n"
"SD.save_cursor('/tmp/cur', c)\n"
"back = SD.load_cursor('/tmp/cur')\n"
"print('CURSOR', json.dumps(back), 'roundtrip', back == c)\n"
"assert back == c\n"
"import subprocess\n"
"p = subprocess.run([sys.executable, 'train_ounce100m.py', '--help'],\n"
" capture_output=True, text=True)\n"
"print('TRAINER_HELP_RC', p.returncode, 'lines', len((p.stdout or '').splitlines()))\n"
"assert p.returncode == 0, (p.stdout or '')[-1500:] + (p.stderr or '')[-1500:]\n"
# The launcher's own flag list, parsed by the pinned trainer. argparse validates every
# option it has already consumed by the time it reaches a trailing --help, so rc 0 here
# means the real session command line is well-formed -- and rc 2 names the flag that is not.
"argv = json.load(open('/kaggle/working/train_argv.json'))\n"
"p2 = subprocess.run([sys.executable, 'train_ounce100m.py'] + argv[1:] + ['--help'],\n"
" capture_output=True, text=True)\n"
"out2 = (p2.stdout or '') + (p2.stderr or '')\n"
"print('LAUNCHER_ARGV_RC', p2.returncode, 'flags',\n"
" sum(1 for a in argv if a.startswith('--')))\n"
"assert p2.returncode == 0, out2[-1200:]\n"],
"CURSOR_AND_IMPORT", 900)
# Stage 4 -- checkpoint verification against real Hub bytes, on the theory that a gate which authorises
# deleting the only local copy of a checkpoint should be tested where failing is free. Probe v5 measured
# the old whole-object read-back at `push+verify=797.7s` for an 850 MB optimizer.pt (and 26 s for the same
# cycle once the CDN was warm), so the read-back is now two bounded Range requests; this proves the
# replacement still says yes to the 1.27 GB the probe really pushed and still says no three different ways.
# PREP_VERIFY_REPO exists because the probe repos are scratch: point it at `Cion-lab/ounce100m-ckpt` plus a
# `ckpt/checkpoint-N` once the run has one and this stage outlives them.
VERIFY = """
import json, os, sys, time
sys.path.insert(0, "/kaggle/working")
import hubckpt
from huggingface_hub import snapshot_download
REPO = os.environ.get("PREP_VERIFY_REPO") or "Cion-lab/ounce100m-ckpt-probe-0920-0747"
SUB = os.environ.get("PREP_VERIFY_SUB") or "ckpt/checkpoint-10"
d = "/kaggle/working/verify_subject"
t0 = time.time()
snapshot_download(repo_id=REPO, repo_type="dataset", allow_patterns=SUB + "/*", local_dir=d,
max_workers=4)
local = os.path.join(d, SUB)
print("GOT", len(os.listdir(local)), "files from", REPO, SUB, "in", round(time.time() - t0, 1), "s",
flush=True)
assert os.path.exists(os.path.join(local, "model.safetensors")), sorted(os.listdir(local))
t0 = time.time()
v = hubckpt.verify_checkpoint(REPO, local, SUB, repo_type="dataset")
secs = round(time.time() - t0, 1)
print("VERIFY_CLEAN ok", v["ok"], "readback_ok", v["readback_ok"], "seconds", secs,
json.dumps({k: v[k] for k in ("n_local", "n_hub", "missing", "mismatch", "extra", "readback")},
default=str)[:600], flush=True)
assert v["ok"] and v["readback_ok"], v
assert secs < 180, f"verification took {secs}s -- the bounded read-back was supposed to make this cheap"
rb = v["readback"]
assert rb["path"] == "optimizer.pt" and rb["bytes"] < (10 << 20) and len(rb["ranges"]) == 2, rb
# 1. A byte the Hub does not have, flipped INSIDE the read-back window, so the byte compare itself is what
# catches it (the oid compare would too, but only the compare proves the served offset is ours).
p = os.path.join(local, "optimizer.pt")
size = hubckpt.hub_listing(REPO, "dataset")[SUB + "/optimizer.pt"][0]
print("SUBJECT", size, "bytes; ranges", [[0, (4 << 20) - 1], [size // 2, size // 2 + (4 << 20) - 1]],
flush=True)
def tamper(offset):
with open(p, "r+b") as fh:
fh.seek(offset)
b = fh.read(1)
fh.seek(offset)
fh.write(bytes([b[0] ^ 0xFF]))
return b, offset
def restore(b, offset):
with open(p, "r+b") as fh:
fh.seek(offset)
fh.write(b)
b, off = tamper(12345)
v2 = hubckpt.verify_checkpoint(REPO, local, SUB, repo_type="dataset")
print("VERIFY_TAMPERED_IN_RANGE ok", v2["ok"], "readback",
json.dumps(v2["readback"].get("error", ""))[:200], flush=True)
assert not v2["ok"], "a flipped byte at 12,345 verified clean"
assert "differ from the local file" in str(v2["readback"]["error"]), v2["readback"]
restore(b, off)
b, off = tamper(500_000_001) # outside both windows: only the stored-oid compare can see this
v2b = hubckpt.verify_checkpoint(REPO, local, SUB, repo_type="dataset")
print("VERIFY_TAMPERED_OUTSIDE ok", v2b["ok"], "mismatch", json.dumps(v2b["mismatch"])[:200],
"readback_ok", v2b["readback_ok"], flush=True)
assert not v2b["ok"] and v2b["mismatch"] and v2b["readback_ok"], "the oid compare did not see it"
restore(b, off)
assert hubckpt.verify_checkpoint(REPO, local, SUB, repo_type="dataset")["ok"], "restoring did not help"
# 2. A path that is not on the Hub at all: nothing listed means nothing verified. `missing` is capped at
# ten entries because it is a log field, so the uncapped fact (`n_hub`) is what the assertion uses --
# rehearsing this stage caught that, not a defect in the verifier.
v3 = hubckpt.verify_checkpoint(REPO, local, "ckpt/checkpoint-99999999", repo_type="dataset")
print("VERIFY_ABSENT ok", v3["ok"], "n_hub", v3["n_hub"], "missing(shown)", len(v3["missing"]),
"of n_local", v3["n_local"], flush=True)
assert not v3["ok"] and v3["n_hub"] == 0 and v3["missing"], v3
# 3. Three dishonest servers, each caught by a different clause of the bounded read-back.
class _Resp:
def __init__(self, status, body, crange):
self.status, self._b, self._cr = status, body, crange
self.headers = {} if crange is None else {"Content-Range": crange}
def __enter__(self): return self
def __exit__(self, *a): return False
def read(self, n=-1): return self._b if n is None or n < 0 else self._b[:n]
def _fake(kind):
def f(req, *a, **k):
if "/resolve/" not in getattr(req, "full_url", str(req)):
return real(req, *a, **k) # hub_listing shares urlopen; only the download is faked
rng = req.get_header("Range") or ""
lo, _, hi = rng.split("=")[-1].partition("-")
n = int(hi) - int(lo) + 1 if lo.isdigit() and hi.isdigit() else 0
if kind == "ignores_range": # 200 with a short body and no Content-Range at all
return _Resp(200, b"x" * max(n - 1, 0), None)
if kind == "wrong_span": # 206, exactly the bytes asked for, from the wrong offset
return _Resp(206, b"y" * n, f"bytes 0-{n - 1}/{n}")
return _Resp(206, bytes(n), f"bytes {lo}-{hi}/{size}") # right span, someone else's bytes
return f
real = hubckpt.urllib.request.urlopen
for kind, needle in (("ignores_range", "expected 206"), ("wrong_span", "Content-Range"),
("wrong_bytes", "differ from the local file")):
hubckpt.urllib.request.urlopen = _fake(kind)
try:
vx = hubckpt.verify_checkpoint(REPO, local, SUB, repo_type="dataset")
finally:
hubckpt.urllib.request.urlopen = real
print("VERIFY_" + kind.upper(), "ok", vx["ok"], json.dumps(vx["readback"].get("error"))[:190],
flush=True)
assert not vx["ok"] and needle in str(vx["readback"]["error"]), (kind, vx["readback"])
# 4. The deletion guard itself (E-042): an unverified result must leave the directory alone.
d2 = "/tmp/prune_verified_probe"
os.makedirs(d2, exist_ok=True)
open(os.path.join(d2, "model.safetensors"), "wb").write(b"z" * 1024)
try:
hubckpt.prune_verified(d2, {"verify": {"ok": False, "missing": ["a"], "mismatch": [],
"extra": [], "readback_ok": False}}, min_free_gb=0.0)
raise AssertionError("prune_verified deleted an unverified checkpoint")
except RuntimeError as e:
print("PRUNE_BLOCKED", str(e)[:150], "dir_still_there", os.path.isdir(d2), flush=True)
assert os.path.isdir(d2), "the copy was deleted anyway"
hubckpt.prune_verified(d2, {"verify": {"ok": True, "missing": [], "mismatch": [], "extra": [],
"readback_ok": True}}, min_free_gb=0.0)
assert not os.path.isdir(d2), "a verified checkpoint was not pruned"
print("VERIFY_STAGE_OK clean, two tampers, absent, three dishonest servers, prune guard", flush=True)
"""
open("prepverify.py", "w").write(VERIFY)
rc4, _ = run([sys.executable, "prepverify.py"], "CHECK_VERIFY", 1800, env=dict(PLANENV))
print("VERDICT PHASE4_PREP rc_launcher", rc1, "rc_check", rc2, "rc_cursor", rc3, "rc_verify", rc4,
"seconds", round(time.time() - T0, 1), flush=True)
raise SystemExit(0 if (rc1 == rc2 == rc3 == rc4 == 0) else 4)