Reverendo commited on
Commit
ba18fe4
·
verified ·
1 Parent(s): 1263dc6

fix DDP hang: single checkpoint write, NCCL timeout

Browse files
Files changed (1) hide show
  1. scugnizz-llama.py +29 -8
scugnizz-llama.py CHANGED
@@ -5,7 +5,8 @@
5
  # dependencies = ["torch","datasets","transformers","huggingface_hub","numpy"]
6
  # ///
7
 
8
- import argparse, json, math, os, random, time, warnings
 
9
  from contextlib import nullcontext
10
  from dataclasses import dataclass
11
  from pathlib import Path
@@ -301,7 +302,7 @@ def init_dist():
301
  if "RANK" not in os.environ:
302
  dev = "cuda" if torch.cuda.is_available() else "cpu"
303
  return 0, 1, 0, dev
304
- dist.init_process_group(backend="nccl")
305
  rank = int(os.environ["RANK"])
306
  local_rank = int(os.environ["LOCAL_RANK"])
307
  world_size = int(os.environ["WORLD_SIZE"])
@@ -433,10 +434,19 @@ def save_ckpt(path, model, opt, step, best, rank, world_size, mirror_resume=Fals
433
  print("SAVED", path, f"step={step}", flush=True)
434
  if mirror_resume:
435
  resume = path.parent / RESUME_CKPT
436
- tmp_resume = resume.with_suffix(resume.suffix + ".tmp")
437
- torch.save(payload, tmp_resume)
438
- tmp_resume.replace(resume)
439
- print("SAVED", resume, f"step={step} (keep for optimizer resume)", flush=True)
 
 
 
 
 
 
 
 
 
440
  barrier(rank, world_size)
441
 
442
 
@@ -690,7 +700,7 @@ def main():
690
 
691
  log0(rank, "initializing data streams...", flush=True)
692
  train = Batcher(a, tok, dev, rank, log=is_main(rank))
693
- val = Batcher(a, tok, dev, 10_000 + rank, log=False)
694
  log0(rank, f"starting training loop from step {start}...", flush=True)
695
  t0 = time.time()
696
  roll = 0.0
@@ -730,17 +740,28 @@ def main():
730
  if step > 0 and a.weights_save_interval > 0 and step % a.weights_save_interval == 0:
731
  save_weights(weights_last, model, step, rank, world_size)
732
 
 
733
  if step > 0 and step % a.eval_interval == 0:
734
  if is_main(rank):
735
  vl = eval_model(raw_model(model), val, a, dev, dt)
736
  log0(rank, f"EVAL {step:08d} | val_loss {vl:.4f} | val_ppl {math.exp(min(20, vl)):.1f}", flush=True)
737
  if vl < best:
738
  best = vl
739
- save_ckpt(bestp, model, opt, step, best, rank, world_size)
 
 
 
 
 
 
 
 
740
  barrier(rank, world_size)
741
 
742
  if step > 0 and step % a.save_interval == 0:
743
  save_ckpt(last, model, opt, step, best, rank, world_size, mirror_resume=True)
 
 
744
  if is_main(rank) and a.push_every_save:
745
  upload_final(str(out), a.hub_repo_id, a.hub_path, f"pretrain checkpoint step {step}")
746
 
 
5
  # dependencies = ["torch","datasets","transformers","huggingface_hub","numpy"]
6
  # ///
7
 
8
+ import argparse, json, math, os, random, shutil, time, warnings
9
+ from datetime import timedelta
10
  from contextlib import nullcontext
11
  from dataclasses import dataclass
12
  from pathlib import Path
 
302
  if "RANK" not in os.environ:
303
  dev = "cuda" if torch.cuda.is_available() else "cpu"
304
  return 0, 1, 0, dev
305
+ dist.init_process_group(backend="nccl", timeout=timedelta(hours=2))
306
  rank = int(os.environ["RANK"])
307
  local_rank = int(os.environ["LOCAL_RANK"])
308
  world_size = int(os.environ["WORLD_SIZE"])
 
434
  print("SAVED", path, f"step={step}", flush=True)
435
  if mirror_resume:
436
  resume = path.parent / RESUME_CKPT
437
+ shutil.copy2(path, resume)
438
+ print("SAVED", resume, f"step={step} (copy for optimizer resume)", flush=True)
439
+ barrier(rank, world_size)
440
+
441
+
442
+ def copy_best_ckpt(out, rank, world_size):
443
+ if not is_main(rank):
444
+ barrier(rank, world_size)
445
+ return
446
+ last, best = out / "checkpoint_last.pt", out / "checkpoint_best.pt"
447
+ if last.exists():
448
+ shutil.copy2(last, best)
449
+ print("SAVED", best, "(copy of checkpoint_last)", flush=True)
450
  barrier(rank, world_size)
451
 
452
 
 
700
 
701
  log0(rank, "initializing data streams...", flush=True)
702
  train = Batcher(a, tok, dev, rank, log=is_main(rank))
703
+ val = Batcher(a, tok, dev, 10_000, log=False) if is_main(rank) else None
704
  log0(rank, f"starting training loop from step {start}...", flush=True)
705
  t0 = time.time()
706
  roll = 0.0
 
740
  if step > 0 and a.weights_save_interval > 0 and step % a.weights_save_interval == 0:
741
  save_weights(weights_last, model, step, rank, world_size)
742
 
743
+ val_improved = False
744
  if step > 0 and step % a.eval_interval == 0:
745
  if is_main(rank):
746
  vl = eval_model(raw_model(model), val, a, dev, dt)
747
  log0(rank, f"EVAL {step:08d} | val_loss {vl:.4f} | val_ppl {math.exp(min(20, vl)):.1f}", flush=True)
748
  if vl < best:
749
  best = vl
750
+ val_improved = True
751
+ best_pkg = [best]
752
+ if world_size > 1:
753
+ dist.broadcast_object_list(best_pkg, src=0)
754
+ best = float(best_pkg[0])
755
+ if world_size > 1:
756
+ imp_pkg = [val_improved]
757
+ dist.broadcast_object_list(imp_pkg, src=0)
758
+ val_improved = imp_pkg[0]
759
  barrier(rank, world_size)
760
 
761
  if step > 0 and step % a.save_interval == 0:
762
  save_ckpt(last, model, opt, step, best, rank, world_size, mirror_resume=True)
763
+ if val_improved:
764
+ copy_best_ckpt(out, rank, world_size)
765
  if is_main(rank) and a.push_every_save:
766
  upload_final(str(out), a.hub_repo_id, a.hub_path, f"pretrain checkpoint step {step}")
767