train: a segment that stopped short of the horizon skips the validation pass instead of entering a collective the two ranks can reach a batch apart (probe 1 v1/v4 found rank 1 dying there with no traceback, while every horizon leg passed the same call)
Browse files- train/train_ounce100m.py +25 -11
train/train_ounce100m.py
CHANGED
|
@@ -612,20 +612,32 @@ def main():
|
|
| 612 |
if final_loss is None or not math.isfinite(final_loss):
|
| 613 |
raise SystemExit(f"final loss is not finite ({final_loss}) -- not exporting anything")
|
| 614 |
|
| 615 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 616 |
ppl = None
|
| 617 |
val_err = None
|
| 618 |
-
|
| 619 |
-
|
| 620 |
-
|
| 621 |
-
|
| 622 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 623 |
# `say()` is a no-op on rank 1, and that silence is what made a rank-1 evaluation crash unreadable
|
| 624 |
# in probe 1 v1: the handler ran on both ranks, only rank 0 spoke, and the divergence looked like an
|
| 625 |
# unexplained exit 1 with `error_file: <N/A>`. Report on every rank, and carry the failure into the
|
| 626 |
# record so a missing PPL can never be read as a passing validation pass.
|
| 627 |
-
|
| 628 |
-
|
| 629 |
|
| 630 |
# Sampled here, on every rank, before the rank-0-only export block: a collective reduction inside
|
| 631 |
# `if rank == 0:` would hang, and reading rank 0's allocator alone is what made the probe's memory
|
|
@@ -674,13 +686,14 @@ def main():
|
|
| 674 |
"attn": args.attn, "effective_attn": eff_attn, "grad_ckpt": bool(args.grad_ckpt),
|
| 675 |
"optim": args.optim, "lr": args.lr, "warmup_frac": args.warmup_frac,
|
| 676 |
"decay_frac": args.decay_frac, "seed": args.seed, "data_seed": args.data_seed,
|
| 677 |
-
"val_ppl": ppl, "val_loss_error": val_err, "
|
|
|
|
| 678 |
"peak_gpu_gb": mem["rank0_allocated_gb"],
|
| 679 |
"peak_reserved_gb": mem["rank0_reserved_gb"],
|
| 680 |
"peak_gpu_gb_max_rank": mem["max_allocated_gb"],
|
| 681 |
"peak_reserved_gb_max_rank": mem["max_reserved_gb"]},
|
| 682 |
open(os.path.join(exp, "run_summary.json"), "w"), indent=1, sort_keys=True)
|
| 683 |
-
segment =
|
| 684 |
if api is not None and not segment and rank == 0:
|
| 685 |
r = hubckpt.push_and_prune(args.hub_repo, exp, "final", api, repo_type=args.hub_repo_type,
|
| 686 |
token=os.environ.get("HF_TOKEN"), prune=args.prune)
|
|
@@ -722,7 +735,8 @@ def main():
|
|
| 722 |
# So a caller cannot report success for a validation pass that never
|
| 723 |
# happened: a null val_ppl with no error here would be a bug, and a null
|
| 724 |
# with an error is a measured failure.
|
| 725 |
-
"val_error": val_err, "val_windows": n_val},
|
|
|
|
| 726 |
|
| 727 |
|
| 728 |
def peak_stats():
|
|
|
|
| 612 |
if final_loss is None or not math.isfinite(final_loss):
|
| 613 |
raise SystemExit(f"final loss is not finite ({final_loss}) -- not exporting anything")
|
| 614 |
|
| 615 |
+
# A segment that stopped short of the horizon does not run the validation pass, on every rank.
|
| 616 |
+
#
|
| 617 |
+
# Probe 1 v1 and v4 showed the same split twice: the leg stopped early by `--stop-after-steps` died inside
|
| 618 |
+
# `evaluate()` on rank 1 with no Python traceback at all, while the leg that reached its horizon ran the
|
| 619 |
+
# same call and exported cleanly. `should_training_stop` is raised from `on_step_end`, so the two ranks
|
| 620 |
+
# can leave the inner loop one batch apart, and a collective entered in that state kills a rank without
|
| 621 |
+
# raising -- which is exactly the signature. The report's PPL comes from the horizon leg, and the
|
| 622 |
+
# mid-run legs keep their training loss, which is what §5 asks to watch on every wake.
|
| 623 |
+
forced_stop = bool(hub_cb and hub_cb.stopped_at and step < steps_planned)
|
| 624 |
ppl = None
|
| 625 |
val_err = None
|
| 626 |
+
if forced_stop:
|
| 627 |
+
say(f"validation skipped: segment ended at step {step} of {steps_planned} (evaluate() after a "
|
| 628 |
+
"forced stop diverges between ranks -- probe 1 v1/v4)")
|
| 629 |
+
else:
|
| 630 |
+
try:
|
| 631 |
+
ev = tr.evaluate()
|
| 632 |
+
ppl = math.exp(min(20.0, ev["eval_loss"]))
|
| 633 |
+
say(f"validation loss={ev['eval_loss']:.4f} ppl={ppl:.2f} over {n_val} held-out windows")
|
| 634 |
+
except Exception as e:
|
| 635 |
# `say()` is a no-op on rank 1, and that silence is what made a rank-1 evaluation crash unreadable
|
| 636 |
# in probe 1 v1: the handler ran on both ranks, only rank 0 spoke, and the divergence looked like an
|
| 637 |
# unexplained exit 1 with `error_file: <N/A>`. Report on every rank, and carry the failure into the
|
| 638 |
# record so a missing PPL can never be read as a passing validation pass.
|
| 639 |
+
val_err = f"{type(e).__name__}: {str(e)[:300]}"
|
| 640 |
+
print(f"[rank {rank}] validation eval failed: {val_err}", flush=True)
|
| 641 |
|
| 642 |
# Sampled here, on every rank, before the rank-0-only export block: a collective reduction inside
|
| 643 |
# `if rank == 0:` would hang, and reading rank 0's allocator alone is what made the probe's memory
|
|
|
|
| 686 |
"attn": args.attn, "effective_attn": eff_attn, "grad_ckpt": bool(args.grad_ckpt),
|
| 687 |
"optim": args.optim, "lr": args.lr, "warmup_frac": args.warmup_frac,
|
| 688 |
"decay_frac": args.decay_frac, "seed": args.seed, "data_seed": args.data_seed,
|
| 689 |
+
"val_ppl": ppl, "val_loss_error": val_err, "val_skipped": forced_stop,
|
| 690 |
+
"val_windows": n_val,
|
| 691 |
"peak_gpu_gb": mem["rank0_allocated_gb"],
|
| 692 |
"peak_reserved_gb": mem["rank0_reserved_gb"],
|
| 693 |
"peak_gpu_gb_max_rank": mem["max_allocated_gb"],
|
| 694 |
"peak_reserved_gb_max_rank": mem["max_reserved_gb"]},
|
| 695 |
open(os.path.join(exp, "run_summary.json"), "w"), indent=1, sort_keys=True)
|
| 696 |
+
segment = forced_stop # the same predicate, computed before the validation decision above
|
| 697 |
if api is not None and not segment and rank == 0:
|
| 698 |
r = hubckpt.push_and_prune(args.hub_repo, exp, "final", api, repo_type=args.hub_repo_type,
|
| 699 |
token=os.environ.get("HF_TOKEN"), prune=args.prune)
|
|
|
|
| 735 |
# So a caller cannot report success for a validation pass that never
|
| 736 |
# happened: a null val_ppl with no error here would be a bug, and a null
|
| 737 |
# with an error is a measured failure.
|
| 738 |
+
"val_error": val_err, "val_skipped": forced_stop, "val_windows": n_val},
|
| 739 |
+
sort_keys=True))
|
| 740 |
|
| 741 |
|
| 742 |
def peak_stats():
|