File size: 18,073 Bytes
dd31402
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bc34e7d
dd31402
 
 
14e982b
dd31402
 
 
 
 
 
a4158df
dd31402
14e982b
dd31402
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bc34e7d
 
 
 
 
dd31402
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bc34e7d
dd31402
 
 
98922ef
 
dd31402
 
98922ef
dd31402
 
 
 
 
 
 
 
98922ef
 
 
 
 
 
 
 
 
 
dd31402
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
98922ef
 
 
 
 
dd31402
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
14e982b
 
 
 
 
 
 
 
 
 
 
dd31402
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
# 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)