File size: 6,449 Bytes
932bc69 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | """Full-run launcher: transactional checkpoints and durable progress metadata."""
import json
import logging
import os
import runpy
import shutil
import tempfile
import time
from collections.abc import Mapping
from pathlib import Path
import torch.distributed as dist
from ngram_checkpoint_transaction import (
MARKER,
latest_managed_epoch,
publish_checkpoint,
)
from ngram_epoch_accounting import full_epoch_steps, validate_completed_state
from speculators.train.checkpointer import BaseCheckpointer
from ngram_wandb_resume import update_resumable_config
from speculators.train.logger import WandbHandler, _flatten_dict
from speculators.train.trainer import Trainer
run_root = Path(os.environ["NGRAM_RUN_ROOT"])
log_root = Path(os.environ["LOG_DIR"])
original_save = Trainer.maybe_save_checkpoint
original_setup = Trainer.setup_optimizer
original_run = Trainer.run_training
original_wandb_setup = WandbHandler._setup
original_wandb_emit = WandbHandler.emit
original_previous_epoch = BaseCheckpointer._get_previous_epoch
def previous_epoch(self):
# Upstream ignores all symlinks. Permit only our validated numeric pointers;
# descriptive aliases and arbitrary external symlinks remain ignored.
return latest_managed_epoch(self.path, original_previous_epoch(self))
def wandb_setup(self):
# Segment-local paths change at renewal. Keep the same W&B run while the
# immutable source/config and per-segment provenance preserve all versions.
self.init_kwargs.setdefault("allow_val_change", True)
return original_wandb_setup(self)
def wandb_emit(self, record):
if getattr(record, "hparams", False) and isinstance(record.msg, Mapping):
if self._run is None:
self._run = self._setup()
update_resumable_config(self._run.config, _flatten_dict(record.msg))
return
return original_wandb_emit(self, record)
def write_json(path, value):
temporary = path.with_suffix(".json.tmp")
temporary.write_text(json.dumps(value, indent=2) + "\n")
os.replace(temporary, path)
class ProgressRecorder(logging.Handler):
def emit(self, record):
if isinstance(record.msg, dict) and "global_step" in record.msg:
write_json(
run_root / "progress.json",
{
"time": time.time(),
"segment": str(log_root),
"global_step": record.msg["global_step"],
"train": record.msg.get("train", {}),
"profile": record.msg.get("profile", {}),
},
)
def setup(self):
original_setup(self)
if self.rank == 0:
for name in ("train_command.txt", "speculators.patch", "run.yaml"):
path = self.checkpointer.path / name
if path.is_file():
shutil.copy2(path, log_root / name)
if (self.config.max_steps is not None
or self.config.num_epochs != int(os.environ["EPOCHS"])):
raise ValueError(
"Run must match the manifest epoch count, with no max_steps"
)
logging.getLogger("speculators.metrics").addHandler(ProgressRecorder())
write_json(
log_root / "training_start.json",
{
"epoch_steps": len(self.train_loader),
"epochs": self.config.num_epochs,
"total_steps": self.config.num_epochs * len(self.train_loader),
"resume_global_step": self.global_step,
"resume_local_step": self._resume_local_step,
"current_epoch": self.current_epoch,
"checkpoint": str(self.checkpointer.prev_path),
"scheduler_last_epochs": [s.last_epoch for s in self.schedulers],
},
)
def transactional_save(self, epoch, local_step=0):
if not isinstance(epoch, int):
return original_save(self, epoch, local_step)
if self.config.save_best or self.config.checkpoint_freq >= 1:
raise ValueError("Use periodic checkpoint_freq < 1 without save_best")
root = self.checkpointer.path
paths = [None]
if self.rank == 0:
stage = Path(tempfile.mkdtemp(prefix=".pending-", dir=root))
(stage / MARKER).touch()
for name in ("train_command.txt", "speculators.patch", "run.yaml"):
if (root / name).is_file():
shutil.copy2(root / name, stage / name)
paths[0] = str(stage)
if self.is_distributed:
dist.broadcast_object_list(paths, src=0)
stage = Path(paths[0])
self.checkpointer.path = stage
try:
original_save(self, epoch, local_step)
finally:
self.checkpointer.path = root
if self.is_distributed:
dist.barrier()
if self.rank == 0:
result = publish_checkpoint(root, stage, epoch, self.global_step)
result.update(time=time.time(), segment=str(log_root))
write_json(run_root / "checkpoint_commit.json", result)
print("ATOMIC_CHECKPOINT " + json.dumps(result), flush=True)
if self.is_distributed:
dist.barrier()
def run(self):
original_run(self)
if self.rank == 0:
final_epoch = self.config.num_epochs - 1
checkpoint = self.checkpointer.path / str(final_epoch)
state = json.loads(
(checkpoint / "training_state.json").read_text()
)
# Packed batch counts can differ by epoch, and a resumed sampler's cache
# contains only its remaining slice. Audit full epochs on a separate
# sampler; never infer completion by multiplying the final loader length.
epoch_steps = full_epoch_steps(self.train_loader, self.config.num_epochs)
validate_completed_state(state, self.config.num_epochs, epoch_steps)
write_json(
run_root / "training_complete.json",
{
"passed": True,
"time": time.time(),
"state": state,
"checkpoint": str(checkpoint),
"epoch_steps": epoch_steps,
"total_steps": sum(epoch_steps),
},
)
Trainer.setup_optimizer = setup
Trainer.maybe_save_checkpoint = transactional_save
Trainer.run_training = run
WandbHandler._setup = wandb_setup
WandbHandler.emit = wandb_emit
BaseCheckpointer._get_previous_epoch = previous_epoch
runpy.run_path(
str(Path(__file__).resolve().parents[1] / "train.py"), run_name="__main__"
)
|