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

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
Files changed (1) hide show
  1. 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
- # Validation first, and on every rank: `evaluate()` is collective, so gating it would deadlock.
 
 
 
 
 
 
 
 
616
  ppl = None
617
  val_err = None
618
- try:
619
- ev = tr.evaluate()
620
- ppl = math.exp(min(20.0, ev["eval_loss"]))
621
- say(f"validation loss={ev['eval_loss']:.4f} ppl={ppl:.2f} over {n_val} held-out windows")
622
- except Exception as e:
 
 
 
 
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
- val_err = f"{type(e).__name__}: {str(e)[:300]}"
628
- print(f"[rank {rank}] validation eval failed: {val_err}", flush=True)
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, "val_windows": n_val,
 
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 = bool(hub_cb and hub_cb.stopped_at and step < steps_planned)
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}, sort_keys=True))
 
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():