Download source/scripts/cluster/ngram_train_continuous.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 6.45 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/scripts/cluster/ngram_train_continuous.py
- Command line
-
hf download hf://khazic/spec-b300/source/scripts/cluster/ngram_train_continuous.py
-
curl -L -o ngram_train_continuous.py https://huggingface.co/khazic/spec-b300/resolve/main/source/scripts/cluster/ngram_train_continuous.py
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__" | |
| ) | |