"""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__" )