fix DDP hang: single checkpoint write, NCCL timeout
Browse files- 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 |
-
|
| 437 |
-
|
| 438 |
-
|
| 439 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
|