Cion-lab commited on
Commit
e7fecad
·
verified ·
1 Parent(s): 4c7f492

p4_stop_probe: repair the soak branch's timeout tuple (an earlier edit dropped it, which would have been a NameError only in soak mode) and right-size the leg ceilings under the notebook timeout

Browse files
Files changed (1) hide show
  1. kernels/p4_stop_probe.py +279 -275
kernels/p4_stop_probe.py CHANGED
@@ -1,275 +1,279 @@
1
- # P4 probe: does --stop-after-steps really produce a resumable, Hub-verified segment boundary?
2
- #
3
- # Why this exists and what it buys with ~0.4 GPU-hours: `--stop-after-steps` is the mechanism every Phase 4
4
- # session ends on, and Gate 3 never exercised it -- P3's legs used `--max-steps`, so each leg *was* a whole
5
- # run. The launcher's plan, the trainer's stop, HubPush's push-verify-pointer-prune order, the resume scan
6
- # that refuses a stale pointer, and the final-push-plus-terminal-pointer path are therefore untested
7
- # together. E-035/4 is precisely a defect in that untested seam, and the rehearsal kernel cannot reach it
8
- # because it needs two T4s. This does, against the published mix and the frozen 22L geometry, in one
9
- # session: leg 1 trains to a mid-run stop step, leg 2 wipes the disk, resumes from the Hub and runs to the
10
- # horizon.
11
- #
12
- # The probe writes to its own checkpoint repo, never to Cion-lab/ounce100m-ckpt: the first real session has
13
- # to find that repo absent, which is the RepoMissing branch E-034 was about.
14
- #
15
- # It carries a second question, the one the user pushed back on (D-017): gradient checkpointing costs 32 %
16
- # of the throughput (9,696 -> 12,792 tok/s, 29.7 h -> 22.5 h) and gives up the memory headroom, sitting at
17
- # 12.25 GB of ~14.56. So this cell runs the *real* geometry at --accum 32, micro 4, WITHOUT checkpointing,
18
- # for 180 steps, through three push/verify/prune cycles and one forced cold resume and the end-of-run
19
- # validation pass -- the three things a 30-step throughput cell cannot show: allocator drift across a few
20
- # hundred steps, the save path's host/GPU copies while the card is 84 % full, and the eval forward pass.
21
- # Peak memory is asserted, not eyeballed. If it holds, the run adopts it as D-018 with a finer push cadence
22
- # as the bounded blast radius; if it does not, D-017 stands and this is the measurement that says so.
23
- import hashlib, json, os, shutil, signal, subprocess, sys, threading, time
24
-
25
- os.chdir("/kaggle/working")
26
- sys.path.insert(0, "/kaggle/working")
27
- REV = "850729d3ef34359c3b10826d898681634cf3274c"
28
- WANT = {
29
- "ounce100m_credentials.py": ("ounce100m_credentials.py",
30
- "6525f62f03f2d73650a1eb4f70fcb52d1194caad4ca88b2d8bd8fd54f88339b6"),
31
- "shard_dataset.py": ("train/shard_dataset.py",
32
- "f35653bf4c8f2cfe7bb0c2c7835e308505fb5a84b4eb002d0f82f2cada768ca6"),
33
- "hubckpt.py": ("train/hubckpt.py",
34
- "568c31b300906cb8d78a59ef890baa9731b18e064e3d69b0eaea4b45c546f23e"),
35
- "train_ounce100m.py": ("train/train_ounce100m.py",
36
- "98f1402ace0f76dfff89bc50bbdb8f49a270e165d711b8ceef71cd6030ba2665"),
37
- }
38
- BASE = "https://huggingface.co/Cion-lab/ounce100m-code/resolve/" + REV
39
- for p, (rp, want) in sorted(WANT.items()):
40
- assert subprocess.run(["curl", "-sfL", f"{BASE}/{rp}", "-o", p]).returncode == 0, ("fetch", rp)
41
- got = hashlib.sha256(open(p, "rb").read()).hexdigest()
42
- assert got == want, ("SHA MISMATCH", rp, got[:16], want[:16])
43
- print("OK", p, got[:12], flush=True)
44
-
45
- import ounce100m_credentials as C
46
- print("creds:", json.dumps(C.install(verify=True)), flush=True)
47
- import hubckpt
48
- from huggingface_hub import HfApi
49
-
50
- MIX = "Cion-lab/ounce100m-mix-v1"
51
- # A scratch repo per attempt (dated), never the real run's repo: v1 left `final` at step 20 and v2 was
52
- # correctly refused by the stale-stop guard, so rather than clearing state between attempts each one gets an
53
- # empty repo. That also makes every attempt walk the RepoMissing resume branch (E-034) that session 1 hits.
54
- PROBE = os.environ.get("P4_PROBE_REPO") or (
55
- "Cion-lab/ounce100m-ckpt-probe-" + time.strftime("%m%d-%H%M", time.gmtime()))
56
- ROOT, RUN = "/kaggle/working/mixroot", "/kaggle/working/run"
57
- # Two modes, one file, so the assertions are literally the same code in both. `smoke` is the user's
58
- # suggestion and it is the right order: 20 steps costs ~12 minutes and answers "does the training code run
59
- # at all, and does one checkpoint survive the push/verify/pointer/prune cycle" -- which is exactly what the
60
- # first run of this probe failed at, in 8.6 seconds, on a malformed torchrun command line (E-037). `soak`
61
- # is the 180-step memory question, and it is only worth 1.6 GPU-hours once the mechanics are known to work.
62
- MODE = os.environ.get("P4_PROBE_MODE", "soak")
63
- if MODE not in ("smoke", "soak"):
64
- # A typo here would otherwise run the 1.6-hour soak when a 13-minute smoke was asked for.
65
- raise SystemExit("P4_PROBE_MODE must be smoke or soak, got %r" % MODE)
66
- TLOG = "/kaggle/working/tlogs"
67
- # torchrun takes its first positional as the SCRIPT, not a command: passing sys.executable made
68
- # it compile the Python binary (E-037). --redirects is a bitmask into the --log-dir per rank:
69
- # 2 = redirect only stderr, so rank 1's traceback lands in a file while stdout (the loss curve,
70
- # which §5 requires monitoring on every wake) still streams live to this log. 3 would file both.
71
- TORCHRUN = ["torchrun", "--nproc_per_node=2", "--redirects", "3", "--tee", "3", "--log-dir", TLOG]
72
- TPS = 262144 # the run's real step shape, in tokens
73
- if MODE == "smoke":
74
- STEPS, PUSH_EVERY, STOP1, VAL = 20, 10, 10, 200000
75
- T_LEG1, T_LEG2, T_FRESH = 1200, 900, 1500
76
- else:
77
- STEPS, PUSH_EVERY, STOP1, VAL = 180, 60, 120, 2000000
78
- T_LEG1, T_LEG2, T_FRESH = 4800, 2700, 1500
79
- TOKENS = STEPS * TPS
80
- GATE_PEAK = (MODE != "smoke") # 20 steps says nothing about allocator drift
81
- T0 = time.time()
82
-
83
-
84
- def run(argv, label, timeout):
85
- print("=== " + label, flush=True)
86
- t0 = time.time()
87
- e = dict(os.environ); e["PYTHONPATH"] = "/kaggle/working"
88
- e["PYTHONUNBUFFERED"] = "1" # the -u that torchrun cannot carry
89
- e["NCCL_DEBUG"] = "WARN" # a rank that dies in a collective says so here and nowhere else
90
- e["TORCH_CPP_LOG_LEVEL"] = "WARNING"
91
- p = subprocess.Popen(argv, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
92
- env=e, bufsize=1, start_new_session=True)
93
- killed = []
94
-
95
- def _kill():
96
- killed.append(True)
97
- try:
98
- os.killpg(os.getpgid(p.pid), signal.SIGTERM)
99
- except Exception:
100
- p.kill()
101
-
102
- def _hard():
103
- killed.append(True)
104
- try:
105
- os.killpg(os.getpgid(p.pid), signal.SIGKILL)
106
- except Exception:
107
- p.kill()
108
-
109
- timer = threading.Timer(timeout, _kill)
110
- timer.daemon = True
111
- timer.start()
112
- hard = threading.Timer(timeout + 90, _hard)
113
- hard.daemon = True
114
- hard.start()
115
- keep, lines = [], []
116
- try:
117
- for line in p.stdout:
118
- line = line.rstrip("\n")
119
- lines.append(line)
120
- if line.startswith(("CKPT ", "RUN_JSON ", "resume from", "auto-resume", "precision:",
121
- "params:", "mix:", "checkpoint hub target", "segment boundary",
122
- "latest.json", "TRAIN DONE", "validation loss", "Traceback", "Error")):
123
- print(" KEY>", line[:300], flush=True)
124
- keep.append(line)
125
- del keep[:-40]
126
- finally:
127
- timer.cancel()
128
- hard.cancel()
129
- rc = p.wait()
130
- if killed:
131
- print(" TIMEOUT after %d s" % timeout, flush=True)
132
- rc = -9
133
- if rc != 0:
134
- print(" TAIL:\n" + "\n".join(keep)[-2500:], flush=True)
135
- print("%s_RC %s seconds %.1f elapsed %.0f" % (label, rc, time.time() - t0, time.time() - T0),
136
- flush=True)
137
- return rc, "\n".join(lines)
138
-
139
-
140
- rc, out = run([sys.executable, "-c",
141
- "import sys; sys.path.insert(0, '/kaggle/working')\n"
142
- "from huggingface_hub import snapshot_download\n"
143
- "p = snapshot_download(repo_id='%s', repo_type='dataset',\n"
144
- " local_dir='/kaggle/working/mixroot', max_workers=4)\n"
145
- "print('mix at', p)\n" % MIX], "FETCH_MIX", T_FRESH)
146
- if rc != 0:
147
- raise SystemExit("VERDICT P4PROBE_STOP could not fetch the published mix")
148
- man = json.load(open(os.path.join(ROOT, "manifest.json")))
149
- print("mix", man["n_shards"], "shards", format(int(man["total_tokens"]), ","), "tokens",
150
- "| probe repo", PROBE, flush=True)
151
-
152
- common = ["train_ounce100m.py", "--root", ROOT, "--out", RUN,
153
- "--hub-repo", PROBE, "--prune", "--seq-len", "1024", "--attn", "eager",
154
- # The flag under test. Passing both --grad-ckpt and --no-grad-ckpt would leave it to argparse's
155
- # last-wins ordering, which is not a thing to be ambiguous about in a probe of this recipe.
156
- "--no-grad-ckpt",
157
- "--micro-batch", "4", "--accum", "32",
158
- "--tokens", str(TOKENS), "--lr", "6e-4",
159
- "--push-every-steps", str(PUSH_EVERY), "--val-tokens", str(VAL), "--log-every", "5",
160
- "--resume", "auto"]
161
-
162
-
163
- def diag(tag):
164
- """torchrun's ChildFailedError said `error_file: <N/A>` and printed no traceback, which is useless for
165
- a rank-1-only failure. --redirects writes each rank's own stdout/stderr into the log dir, so print
166
- those on a failed leg: the exception is in there, and guessing at it costs a session."""
167
- # v4's diag found no *.log at all, so list the tree as well as tailing it: if torchrun named the
168
- # files something else, this says so instead of printing nothing and leaving me guessing.
169
- rc, out = run(["bash", "-c",
170
- 'ls -R "%s" 2>&1 | head -40; '
171
- 'for f in $(find "%s" -type f 2>/dev/null | head -8); do echo "==== $f"; '
172
- 'tail -70 "$f"; done' % (TLOG, TLOG)], "DIAG_" + tag, 180)
173
- return out
174
-
175
-
176
- print("PROBE_REPO", PROBE, flush=True)
177
- shutil.rmtree(RUN, ignore_errors=True)
178
- rc1, o1 = run(TORCHRUN + common + ["--stop-after-steps", str(STOP1)], "LEG1_STOP_EARLY", T_LEG1)
179
- if rc1 != 0:
180
- diag("leg1")
181
- api = HfApi(token=os.environ["HF_TOKEN"])
182
- try:
183
- listed = sorted(hubckpt.hub_listing(PROBE, "dataset", token=os.environ["HF_TOKEN"]))
184
- except Exception as e:
185
- listed = ["<listing failed: %s>" % type(e).__name__]
186
- pushed = sorted({k.split("/")[1] for k in listed if k.startswith("ckpt/checkpoint-")
187
- and len(k.split("/")) > 1})
188
- ptr = hubckpt.latest_pointer(PROBE, token=os.environ["HF_TOKEN"])
189
- print("LEG1 pushed:", pushed, "pointer:", {k: ptr.get(k) for k in ("step", "path_in_repo", "error")},
190
- flush=True)
191
-
192
- shutil.rmtree(RUN, ignore_errors=True) # force the cold-resume path (E-029's shape)
193
- rc2, o2 = run(TORCHRUN + common + ["--stop-after-steps", "0"], "LEG2_TO_HORIZON", T_LEG2)
194
- if rc2 != 0:
195
- diag("leg2")
196
- try:
197
- listed2 = sorted(hubckpt.hub_listing(PROBE, "dataset", token=os.environ["HF_TOKEN"]))
198
- except Exception as e:
199
- listed2 = ["<listing failed: %s>" % type(e).__name__]
200
- ptr2 = hubckpt.latest_pointer(PROBE, token=os.environ["HF_TOKEN"])
201
-
202
-
203
- def harvest(text, tag):
204
- for line in (text or "").splitlines():
205
- if line.startswith(tag):
206
- try:
207
- return json.loads(line[len(tag):].strip())
208
- except Exception:
209
- return {"unparsed": line[:200]}
210
- return {}
211
-
212
-
213
- rj1, rj2 = harvest(o1, "RUN_JSON "), harvest(o2, "RUN_JSON ")
214
- res = {
215
- "leg1_rc": rc1, "leg2_rc": rc2,
216
- "leg1": {"final_step": rj1.get("final_step"), "segment_stop": rj1.get("segment_stop"),
217
- "tokens_consumed": rj1.get("tokens_consumed"), "tok_per_s": rj1.get("tok_per_s"),
218
- "params": rj1.get("params"), "peak_gpu_gb": rj1.get("peak_gpu_gb"),
219
- "tok_per_s": rj1.get("tok_per_s"), "final_loss": rj1.get("final_loss"),
220
- "peak_alloc_gb": rj1.get("peak_gpu_gb"), "peak_alloc_gb_max_rank":
221
- rj1.get("peak_gpu_gb_max_rank"), "peak_reserved_gb_max_rank":
222
- rj1.get("peak_reserved_gb_max_rank"), "val_error": rj1.get("val_error"),
223
- "pushed": pushed, "pointer": {k: ptr.get(k) for k in ("step", "path_in_repo")}},
224
- "leg2": {"final_step": rj2.get("final_step"), "segment_stop": rj2.get("segment_stop"),
225
- "tokens_consumed": rj2.get("tokens_consumed"), "val_ppl": rj2.get("val_ppl"),
226
- "peak_gpu_gb": rj2.get("peak_gpu_gb"), "peak_gpu_gb_max_rank":
227
- rj2.get("peak_gpu_gb_max_rank"), "peak_reserved_gb_max_rank":
228
- rj2.get("peak_reserved_gb_max_rank"), "val_error": rj2.get("val_error"),
229
- "tok_per_s": rj2.get("tok_per_s"),
230
- "final_loss": rj2.get("final_loss"),
231
- "pointer_after": {k: ptr2.get(k) for k in ("step", "path_in_repo")},
232
- "has_final": any(k.startswith("final/") for k in listed2)},
233
- "expect": {"steps_planned": STEPS, "push_every": PUSH_EVERY, "stop1": STOP1,
234
- "tokens_at_stop1": STOP1 * TPS, "tokens_at_horizon": STEPS * TPS, "grad_ckpt": False},
235
- }
236
- res["checks"] = {
237
- "leg1_stopped_at_the_stop_step": rj1.get("final_step") == STOP1,
238
- "leg1_reported_a_segment": rj1.get("segment_stop") is True,
239
- # Integer compare: sorted() over names would order "checkpoint-120" before "checkpoint-60".
240
- "leg1_pushed_every_interval": sorted(int(x.split("-")[1]) for x in pushed) ==
241
- list(range(PUSH_EVERY, STOP1 + 1, PUSH_EVERY)),
242
- "leg1_pointer_at_stop": ptr.get("step") == STOP1,
243
- "leg2_resumed_and_finished": rj2.get("final_step") == STEPS,
244
- "leg2_was_not_a_segment": rj2.get("segment_stop") is False,
245
- "leg2_pushed_final": any(k.startswith("final/") for k in listed2),
246
- "all_steps_present": sorted(int(x.split("-")[1]) for x in
247
- {k.split("/")[1] for k in listed2 if k.startswith("ckpt/checkpoint-")})
248
- == list(range(PUSH_EVERY, STEPS + 1, PUSH_EVERY)),
249
- "leg2_pointer_terminal": ptr2.get("step") == STEPS and ptr2.get("path_in_repo") == "final",
250
- "tokens_match_the_arithmetic": rj1.get("tokens_consumed") == STOP1 * TPS
251
- and rj2.get("tokens_consumed") == STEPS * TPS,
252
- "params_are_the_frozen_model": rj1.get("params") == 106194240,
253
- "checkpointing_really_off": rj1.get("grad_ckpt") is False,
254
- # 13.6 GB of ~14.56 usable leaves the run room for a fragmentation spike and for the two CUDA contexts;
255
- # above that the answer to the user's question is "not on this hardware", not "keep going". Reported in
256
- # both modes, gated only in the soak: a 20-step peak is not evidence about 3,814 steps.
257
- # rank 0's allocated peak is not the constraint: reserved bytes, the CUDA contexts and the *other*
258
- # rank's peak are what an OOM is actually about, so the max across ranks is what gets gated.
259
- "peak_memory_below_13_6_gb": (not GATE_PEAK) or (
260
- (rj1.get("peak_reserved_gb_max_rank") or rj1.get("peak_gpu_gb_max_rank") or 99) <= 13.6
261
- and (rj2.get("peak_reserved_gb_max_rank") or rj2.get("peak_gpu_gb_max_rank") or 99) <= 13.6),
262
- "validation_actually_ran": (rj1.get("val_ppl") is not None and rj1.get("val_error") is None
263
- and rj2.get("val_ppl") is not None and rj2.get("val_error") is None),
264
- "no_nan_and_loss_moved": (rj1.get("final_loss") or 1e9) < 11.0
265
- and (rj2.get("final_loss") or 1e9) < 11.0,
266
- }
267
- res["PROBE_PASSED"] = all(res["checks"].values()) and rc1 == 0 and rc2 == 0
268
- print("PROBE_MODE", MODE, "steps", STEPS, "push_every", PUSH_EVERY,
269
- "stop1", STOP1, "peak_gate", GATE_PEAK, flush=True)
270
- print("PROBE_JSON_BEGIN")
271
- print(json.dumps(res, indent=1, default=str))
272
- print("PROBE_JSON_END")
273
- print("VERDICT P4PROBE", "PASS" if res["PROBE_PASSED"] else "FAIL",
274
- [k for k, v in res["checks"].items() if not v], "seconds", round(time.time() - T0, 1), flush=True)
275
- raise SystemExit(0 if res["PROBE_PASSED"] else 5)
 
 
 
 
 
1
+ # P4 probe: does --stop-after-steps really produce a resumable, Hub-verified segment boundary?
2
+ #
3
+ # Why this exists and what it buys with ~0.4 GPU-hours: `--stop-after-steps` is the mechanism every Phase 4
4
+ # session ends on, and Gate 3 never exercised it -- P3's legs used `--max-steps`, so each leg *was* a whole
5
+ # run. The launcher's plan, the trainer's stop, HubPush's push-verify-pointer-prune order, the resume scan
6
+ # that refuses a stale pointer, and the final-push-plus-terminal-pointer path are therefore untested
7
+ # together. E-035/4 is precisely a defect in that untested seam, and the rehearsal kernel cannot reach it
8
+ # because it needs two T4s. This does, against the published mix and the frozen 22L geometry, in one
9
+ # session: leg 1 trains to a mid-run stop step, leg 2 wipes the disk, resumes from the Hub and runs to the
10
+ # horizon.
11
+ #
12
+ # The probe writes to its own checkpoint repo, never to Cion-lab/ounce100m-ckpt: the first real session has
13
+ # to find that repo absent, which is the RepoMissing branch E-034 was about.
14
+ #
15
+ # It carries a second question, the one the user pushed back on (D-017): gradient checkpointing costs 32 %
16
+ # of the throughput (9,696 -> 12,792 tok/s, 29.7 h -> 22.5 h) and gives up the memory headroom, sitting at
17
+ # 12.25 GB of ~14.56. So this cell runs the *real* geometry at --accum 32, micro 4, WITHOUT checkpointing,
18
+ # for 180 steps, through three push/verify/prune cycles and one forced cold resume and the end-of-run
19
+ # validation pass -- the three things a 30-step throughput cell cannot show: allocator drift across a few
20
+ # hundred steps, the save path's host/GPU copies while the card is 84 % full, and the eval forward pass.
21
+ # Peak memory is asserted, not eyeballed. If it holds, the run adopts it as D-018 with a finer push cadence
22
+ # as the bounded blast radius; if it does not, D-017 stands and this is the measurement that says so.
23
+ import hashlib, json, os, shutil, signal, subprocess, sys, threading, time
24
+
25
+ os.chdir("/kaggle/working")
26
+ sys.path.insert(0, "/kaggle/working")
27
+ REV = "850729d3ef34359c3b10826d898681634cf3274c"
28
+ WANT = {
29
+ "ounce100m_credentials.py": ("ounce100m_credentials.py",
30
+ "6525f62f03f2d73650a1eb4f70fcb52d1194caad4ca88b2d8bd8fd54f88339b6"),
31
+ "shard_dataset.py": ("train/shard_dataset.py",
32
+ "f35653bf4c8f2cfe7bb0c2c7835e308505fb5a84b4eb002d0f82f2cada768ca6"),
33
+ "hubckpt.py": ("train/hubckpt.py",
34
+ "568c31b300906cb8d78a59ef890baa9731b18e064e3d69b0eaea4b45c546f23e"),
35
+ "train_ounce100m.py": ("train/train_ounce100m.py",
36
+ "98f1402ace0f76dfff89bc50bbdb8f49a270e165d711b8ceef71cd6030ba2665"),
37
+ }
38
+ BASE = "https://huggingface.co/Cion-lab/ounce100m-code/resolve/" + REV
39
+ for p, (rp, want) in sorted(WANT.items()):
40
+ assert subprocess.run(["curl", "-sfL", f"{BASE}/{rp}", "-o", p]).returncode == 0, ("fetch", rp)
41
+ got = hashlib.sha256(open(p, "rb").read()).hexdigest()
42
+ assert got == want, ("SHA MISMATCH", rp, got[:16], want[:16])
43
+ print("OK", p, got[:12], flush=True)
44
+
45
+ import ounce100m_credentials as C
46
+ print("creds:", json.dumps(C.install(verify=True)), flush=True)
47
+ import hubckpt
48
+ from huggingface_hub import HfApi
49
+
50
+ MIX = "Cion-lab/ounce100m-mix-v1"
51
+ # A scratch repo per attempt (dated), never the real run's repo: v1 left `final` at step 20 and v2 was
52
+ # correctly refused by the stale-stop guard, so rather than clearing state between attempts each one gets an
53
+ # empty repo. That also makes every attempt walk the RepoMissing resume branch (E-034) that session 1 hits.
54
+ PROBE = os.environ.get("P4_PROBE_REPO") or (
55
+ "Cion-lab/ounce100m-ckpt-probe-" + time.strftime("%m%d-%H%M", time.gmtime()))
56
+ ROOT, RUN = "/kaggle/working/mixroot", "/kaggle/working/run"
57
+ # Two modes, one file, so the assertions are literally the same code in both. `smoke` is the user's
58
+ # suggestion and it is the right order: 20 steps costs ~12 minutes and answers "does the training code run
59
+ # at all, and does one checkpoint survive the push/verify/pointer/prune cycle" -- which is exactly what the
60
+ # first run of this probe failed at, in 8.6 seconds, on a malformed torchrun command line (E-037). `soak`
61
+ # is the 180-step memory question, and it is only worth 1.6 GPU-hours once the mechanics are known to work.
62
+ MODE = os.environ.get("P4_PROBE_MODE", "soak")
63
+ if MODE not in ("smoke", "soak"):
64
+ # A typo here would otherwise run the 1.6-hour soak when a 13-minute smoke was asked for.
65
+ raise SystemExit("P4_PROBE_MODE must be smoke or soak, got %r" % MODE)
66
+ TLOG = "/kaggle/working/tlogs"
67
+ # torchrun takes its first positional as the SCRIPT, not a command: passing sys.executable made
68
+ # it compile the Python binary (E-037). --redirects is a per-rank bitmask into --log-dir (1=stderr,
69
+ # 2=stdout, 3=both) and --tee repeats the same streams to this process, so the loss curve §5 requires
70
+ # watching stays live *and* each rank's traceback is on disk for diag() below.
71
+ TORCHRUN = ["torchrun", "--nproc_per_node=2", "--redirects", "3", "--tee", "3", "--log-dir", TLOG]
72
+ TPS = 262144 # the run's real step shape, in tokens
73
+ if MODE == "smoke":
74
+ STEPS, PUSH_EVERY, STOP1, VAL = 20, 10, 10, 200000
75
+ T_LEG1, T_LEG2, T_FRESH = 1200, 900, 1500
76
+ else:
77
+ STEPS, PUSH_EVERY, STOP1, VAL = 180, 60, 120, 2000000
78
+ # Sized from the model rather than guessed: 120 steps at ~20.5 s + build + three pushes + eval is
79
+ # ~3,000 s, and 60 steps + cold pull + push + eval is ~1,700 s. The three ceilings must also fit under
80
+ # the notebook's own session timeout (10,800 s for the soak) with room for the model build, because a
81
+ # stage that outlives the container reports nothing at all (review point B5).
82
+ T_LEG1, T_LEG2, T_FRESH = 3900, 2400, 1200
83
+ TOKENS = STEPS * TPS
84
+ GATE_PEAK = (MODE != "smoke") # 20 steps says nothing about allocator drift
85
+ T0 = time.time()
86
+
87
+
88
+ def run(argv, label, timeout):
89
+ print("=== " + label, flush=True)
90
+ t0 = time.time()
91
+ e = dict(os.environ); e["PYTHONPATH"] = "/kaggle/working"
92
+ e["PYTHONUNBUFFERED"] = "1" # the -u that torchrun cannot carry
93
+ e["NCCL_DEBUG"] = "WARN" # a rank that dies in a collective says so here and nowhere else
94
+ e["TORCH_CPP_LOG_LEVEL"] = "WARNING"
95
+ p = subprocess.Popen(argv, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
96
+ env=e, bufsize=1, start_new_session=True)
97
+ killed = []
98
+
99
+ def _kill():
100
+ killed.append(True)
101
+ try:
102
+ os.killpg(os.getpgid(p.pid), signal.SIGTERM)
103
+ except Exception:
104
+ p.kill()
105
+
106
+ def _hard():
107
+ killed.append(True)
108
+ try:
109
+ os.killpg(os.getpgid(p.pid), signal.SIGKILL)
110
+ except Exception:
111
+ p.kill()
112
+
113
+ timer = threading.Timer(timeout, _kill)
114
+ timer.daemon = True
115
+ timer.start()
116
+ hard = threading.Timer(timeout + 90, _hard)
117
+ hard.daemon = True
118
+ hard.start()
119
+ keep, lines = [], []
120
+ try:
121
+ for line in p.stdout:
122
+ line = line.rstrip("\n")
123
+ lines.append(line)
124
+ if line.startswith(("CKPT ", "RUN_JSON ", "resume from", "auto-resume", "precision:",
125
+ "params:", "mix:", "checkpoint hub target", "segment boundary",
126
+ "latest.json", "TRAIN DONE", "validation loss", "Traceback", "Error")):
127
+ print(" KEY>", line[:300], flush=True)
128
+ keep.append(line)
129
+ del keep[:-40]
130
+ finally:
131
+ timer.cancel()
132
+ hard.cancel()
133
+ rc = p.wait()
134
+ if killed:
135
+ print(" TIMEOUT after %d s" % timeout, flush=True)
136
+ rc = -9
137
+ if rc != 0:
138
+ print(" TAIL:\n" + "\n".join(keep)[-2500:], flush=True)
139
+ print("%s_RC %s seconds %.1f elapsed %.0f" % (label, rc, time.time() - t0, time.time() - T0),
140
+ flush=True)
141
+ return rc, "\n".join(lines)
142
+
143
+
144
+ rc, out = run([sys.executable, "-c",
145
+ "import sys; sys.path.insert(0, '/kaggle/working')\n"
146
+ "from huggingface_hub import snapshot_download\n"
147
+ "p = snapshot_download(repo_id='%s', repo_type='dataset',\n"
148
+ " local_dir='/kaggle/working/mixroot', max_workers=4)\n"
149
+ "print('mix at', p)\n" % MIX], "FETCH_MIX", T_FRESH)
150
+ if rc != 0:
151
+ raise SystemExit("VERDICT P4PROBE_STOP could not fetch the published mix")
152
+ man = json.load(open(os.path.join(ROOT, "manifest.json")))
153
+ print("mix", man["n_shards"], "shards", format(int(man["total_tokens"]), ","), "tokens",
154
+ "| probe repo", PROBE, flush=True)
155
+
156
+ common = ["train_ounce100m.py", "--root", ROOT, "--out", RUN,
157
+ "--hub-repo", PROBE, "--prune", "--seq-len", "1024", "--attn", "eager",
158
+ # The flag under test. Passing both --grad-ckpt and --no-grad-ckpt would leave it to argparse's
159
+ # last-wins ordering, which is not a thing to be ambiguous about in a probe of this recipe.
160
+ "--no-grad-ckpt",
161
+ "--micro-batch", "4", "--accum", "32",
162
+ "--tokens", str(TOKENS), "--lr", "6e-4",
163
+ "--push-every-steps", str(PUSH_EVERY), "--val-tokens", str(VAL), "--log-every", "5",
164
+ "--resume", "auto"]
165
+
166
+
167
+ def diag(tag):
168
+ """torchrun's ChildFailedError said `error_file: <N/A>` and printed no traceback, which is useless for
169
+ a rank-1-only failure. --redirects writes each rank's own stdout/stderr into the log dir, so print
170
+ those on a failed leg: the exception is in there, and guessing at it costs a session."""
171
+ # v4's diag found no *.log at all, so list the tree as well as tailing it: if torchrun named the
172
+ # files something else, this says so instead of printing nothing and leaving me guessing.
173
+ rc, out = run(["bash", "-c",
174
+ 'ls -R "%s" 2>&1 | head -40; '
175
+ 'for f in $(find "%s" -type f 2>/dev/null | head -8); do echo "==== $f"; '
176
+ 'tail -70 "$f"; done' % (TLOG, TLOG)], "DIAG_" + tag, 180)
177
+ return out
178
+
179
+
180
+ print("PROBE_REPO", PROBE, flush=True)
181
+ shutil.rmtree(RUN, ignore_errors=True)
182
+ rc1, o1 = run(TORCHRUN + common + ["--stop-after-steps", str(STOP1)], "LEG1_STOP_EARLY", T_LEG1)
183
+ if rc1 != 0:
184
+ diag("leg1")
185
+ api = HfApi(token=os.environ["HF_TOKEN"])
186
+ try:
187
+ listed = sorted(hubckpt.hub_listing(PROBE, "dataset", token=os.environ["HF_TOKEN"]))
188
+ except Exception as e:
189
+ listed = ["<listing failed: %s>" % type(e).__name__]
190
+ pushed = sorted({k.split("/")[1] for k in listed if k.startswith("ckpt/checkpoint-")
191
+ and len(k.split("/")) > 1})
192
+ ptr = hubckpt.latest_pointer(PROBE, token=os.environ["HF_TOKEN"])
193
+ print("LEG1 pushed:", pushed, "pointer:", {k: ptr.get(k) for k in ("step", "path_in_repo", "error")},
194
+ flush=True)
195
+
196
+ shutil.rmtree(RUN, ignore_errors=True) # force the cold-resume path (E-029's shape)
197
+ rc2, o2 = run(TORCHRUN + common + ["--stop-after-steps", "0"], "LEG2_TO_HORIZON", T_LEG2)
198
+ if rc2 != 0:
199
+ diag("leg2")
200
+ try:
201
+ listed2 = sorted(hubckpt.hub_listing(PROBE, "dataset", token=os.environ["HF_TOKEN"]))
202
+ except Exception as e:
203
+ listed2 = ["<listing failed: %s>" % type(e).__name__]
204
+ ptr2 = hubckpt.latest_pointer(PROBE, token=os.environ["HF_TOKEN"])
205
+
206
+
207
+ def harvest(text, tag):
208
+ for line in (text or "").splitlines():
209
+ if line.startswith(tag):
210
+ try:
211
+ return json.loads(line[len(tag):].strip())
212
+ except Exception:
213
+ return {"unparsed": line[:200]}
214
+ return {}
215
+
216
+
217
+ rj1, rj2 = harvest(o1, "RUN_JSON "), harvest(o2, "RUN_JSON ")
218
+ res = {
219
+ "leg1_rc": rc1, "leg2_rc": rc2,
220
+ "leg1": {"final_step": rj1.get("final_step"), "segment_stop": rj1.get("segment_stop"),
221
+ "tokens_consumed": rj1.get("tokens_consumed"), "tok_per_s": rj1.get("tok_per_s"),
222
+ "params": rj1.get("params"), "peak_gpu_gb": rj1.get("peak_gpu_gb"),
223
+ "tok_per_s": rj1.get("tok_per_s"), "final_loss": rj1.get("final_loss"),
224
+ "peak_alloc_gb": rj1.get("peak_gpu_gb"), "peak_alloc_gb_max_rank":
225
+ rj1.get("peak_gpu_gb_max_rank"), "peak_reserved_gb_max_rank":
226
+ rj1.get("peak_reserved_gb_max_rank"), "val_error": rj1.get("val_error"),
227
+ "pushed": pushed, "pointer": {k: ptr.get(k) for k in ("step", "path_in_repo")}},
228
+ "leg2": {"final_step": rj2.get("final_step"), "segment_stop": rj2.get("segment_stop"),
229
+ "tokens_consumed": rj2.get("tokens_consumed"), "val_ppl": rj2.get("val_ppl"),
230
+ "peak_gpu_gb": rj2.get("peak_gpu_gb"), "peak_gpu_gb_max_rank":
231
+ rj2.get("peak_gpu_gb_max_rank"), "peak_reserved_gb_max_rank":
232
+ rj2.get("peak_reserved_gb_max_rank"), "val_error": rj2.get("val_error"),
233
+ "tok_per_s": rj2.get("tok_per_s"),
234
+ "final_loss": rj2.get("final_loss"),
235
+ "pointer_after": {k: ptr2.get(k) for k in ("step", "path_in_repo")},
236
+ "has_final": any(k.startswith("final/") for k in listed2)},
237
+ "expect": {"steps_planned": STEPS, "push_every": PUSH_EVERY, "stop1": STOP1,
238
+ "tokens_at_stop1": STOP1 * TPS, "tokens_at_horizon": STEPS * TPS, "grad_ckpt": False},
239
+ }
240
+ res["checks"] = {
241
+ "leg1_stopped_at_the_stop_step": rj1.get("final_step") == STOP1,
242
+ "leg1_reported_a_segment": rj1.get("segment_stop") is True,
243
+ # Integer compare: sorted() over names would order "checkpoint-120" before "checkpoint-60".
244
+ "leg1_pushed_every_interval": sorted(int(x.split("-")[1]) for x in pushed) ==
245
+ list(range(PUSH_EVERY, STOP1 + 1, PUSH_EVERY)),
246
+ "leg1_pointer_at_stop": ptr.get("step") == STOP1,
247
+ "leg2_resumed_and_finished": rj2.get("final_step") == STEPS,
248
+ "leg2_was_not_a_segment": rj2.get("segment_stop") is False,
249
+ "leg2_pushed_final": any(k.startswith("final/") for k in listed2),
250
+ "all_steps_present": sorted(int(x.split("-")[1]) for x in
251
+ {k.split("/")[1] for k in listed2 if k.startswith("ckpt/checkpoint-")})
252
+ == list(range(PUSH_EVERY, STEPS + 1, PUSH_EVERY)),
253
+ "leg2_pointer_terminal": ptr2.get("step") == STEPS and ptr2.get("path_in_repo") == "final",
254
+ "tokens_match_the_arithmetic": rj1.get("tokens_consumed") == STOP1 * TPS
255
+ and rj2.get("tokens_consumed") == STEPS * TPS,
256
+ "params_are_the_frozen_model": rj1.get("params") == 106194240,
257
+ "checkpointing_really_off": rj1.get("grad_ckpt") is False,
258
+ # 13.6 GB of ~14.56 usable leaves the run room for a fragmentation spike and for the two CUDA contexts;
259
+ # above that the answer to the user's question is "not on this hardware", not "keep going". Reported in
260
+ # both modes, gated only in the soak: a 20-step peak is not evidence about 3,814 steps.
261
+ # rank 0's allocated peak is not the constraint: reserved bytes, the CUDA contexts and the *other*
262
+ # rank's peak are what an OOM is actually about, so the max across ranks is what gets gated.
263
+ "peak_memory_below_13_6_gb": (not GATE_PEAK) or (
264
+ (rj1.get("peak_reserved_gb_max_rank") or rj1.get("peak_gpu_gb_max_rank") or 99) <= 13.6
265
+ and (rj2.get("peak_reserved_gb_max_rank") or rj2.get("peak_gpu_gb_max_rank") or 99) <= 13.6),
266
+ "validation_actually_ran": (rj1.get("val_ppl") is not None and rj1.get("val_error") is None
267
+ and rj2.get("val_ppl") is not None and rj2.get("val_error") is None),
268
+ "no_nan_and_loss_moved": (rj1.get("final_loss") or 1e9) < 11.0
269
+ and (rj2.get("final_loss") or 1e9) < 11.0,
270
+ }
271
+ res["PROBE_PASSED"] = all(res["checks"].values()) and rc1 == 0 and rc2 == 0
272
+ print("PROBE_MODE", MODE, "steps", STEPS, "push_every", PUSH_EVERY,
273
+ "stop1", STOP1, "peak_gate", GATE_PEAK, flush=True)
274
+ print("PROBE_JSON_BEGIN")
275
+ print(json.dumps(res, indent=1, default=str))
276
+ print("PROBE_JSON_END")
277
+ print("VERDICT P4PROBE", "PASS" if res["PROBE_PASSED"] else "FAIL",
278
+ [k for k, v in res["checks"].items() if not v], "seconds", round(time.time() - T0, 1), flush=True)
279
+ raise SystemExit(0 if res["PROBE_PASSED"] else 5)