spec-b300 / source /scripts /cluster /ngram_train_continuous.py
khazic's picture
Archive three-epoch run: logs and provenance part 2
932bc69 verified
Raw History Blame Contribute Delete
6.45 kB
"""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__"
)