| #!/usr/bin/env python3 | |
| """Generic Megatron training driver shared by KVPath/GQA train.py wrappers. | |
| This is a KVPath-repo-local copy of DeltaKV's training/baseline/train.py driver | |
| logic (kept path-generic here instead of living under a "baseline" directory, | |
| since this repo has no baseline model of its own). Behavior is otherwise | |
| identical: build the Megatron launch command from a model config.toml, run it, | |
| prune old checkpoints, and optionally run the HF export step. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import copy | |
| import os | |
| import shlex | |
| import subprocess | |
| import sys | |
| import threading | |
| import time | |
| from pathlib import Path | |
| ROOT = Path(__file__).resolve().parents[1] | |
| if str(ROOT) not in sys.path: | |
| sys.path.insert(0, str(ROOT)) | |
| from utils.checkpoint import prune_checkpoints | |
| from utils.megatron_launcher import ( | |
| build_train_command, | |
| load_config, | |
| required_train_files, | |
| resolve_load_dir, | |
| validate_config, | |
| ) | |
| from utils.paths import resolve | |
| def _read_iteration(tracker: Path) -> int | None: | |
| if not tracker.exists(): | |
| return None | |
| text = tracker.read_text().strip() | |
| if not text or text == "release": | |
| return None | |
| return int(text) | |
| def _watch_checkpoints(checkpoint_root: Path, keep_last: int, stop_event: threading.Event, poll_s: float) -> None: | |
| tracker = checkpoint_root / "latest_checkpointed_iteration.txt" | |
| last_seen: int | None = None | |
| while not stop_event.is_set(): | |
| iteration = _read_iteration(tracker) | |
| if iteration is not None and iteration != last_seen: | |
| removed = prune_checkpoints(checkpoint_root, keep_last=keep_last) | |
| if removed: | |
| print(f"Pruned old checkpoints: {removed}", flush=True) | |
| last_seen = iteration | |
| stop_event.wait(poll_s) | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", type=Path, default=Path(__file__).with_name("config.toml")) | |
| parser.add_argument("--optimizer", choices=("adamw",)) | |
| parser.add_argument( | |
| "--data-fraction", | |
| type=float, | |
| default=1.0, | |
| help="Fraction of configured training tokens to consume, in (0, 1].", | |
| ) | |
| parser.add_argument("--dry-run", action="store_true") | |
| parser.add_argument( | |
| "--skip-export", | |
| action="store_true", | |
| help="Skip HuggingFace export after successful training.", | |
| ) | |
| parser.add_argument( | |
| "megatron_args", | |
| nargs=argparse.REMAINDER, | |
| help="Extra Megatron arguments after '--'.", | |
| ) | |
| args = parser.parse_args() | |
| config_path = args.config.resolve() | |
| config = copy.deepcopy(load_config(config_path)) | |
| if not 0.0 < args.data_fraction <= 1.0: | |
| parser.error("--data-fraction must be in (0, 1]") | |
| configured_train_iters = int(config["training"]["train_iters"]) | |
| effective_train_iters = max(1, round(configured_train_iters * args.data_fraction)) | |
| config["training"]["train_iters"] = effective_train_iters | |
| if args.data_fraction < 1.0: | |
| # Keep the LR schedule internally consistent when shrinking train_iters for a | |
| # smoke test: Megatron asserts lr_warmup_iters < lr_decay_iters (== train_iters), | |
| # and lr_wsd_decay_iters must not exceed train_iters either. | |
| schedule = config.get("schedule", {}) | |
| if "lr_warmup_iters" in schedule: | |
| schedule["lr_warmup_iters"] = min( | |
| int(schedule["lr_warmup_iters"]), max(0, effective_train_iters - 1) | |
| ) | |
| if "lr_wsd_decay_iters" in schedule: | |
| schedule["lr_wsd_decay_iters"] = min( | |
| int(schedule["lr_wsd_decay_iters"]), effective_train_iters | |
| ) | |
| effective_tokens = ( | |
| effective_train_iters | |
| * int(config["data"]["global_batch_size"]) | |
| * int(config["data"]["seq_length"]) | |
| ) | |
| print( | |
| f"Training data fraction: {args.data_fraction:g} " | |
| f"({effective_train_iters}/{configured_train_iters} steps, " | |
| f"{effective_tokens / 1e9:.6f}B tokens)", | |
| flush=True, | |
| ) | |
| optimizer = validate_config(config, args.optimizer) | |
| extra = args.megatron_args[1:] if args.megatron_args[:1] == ["--"] else args.megatron_args | |
| missing = [path for path in required_train_files(config) if not path.exists()] | |
| if missing and not args.dry_run: | |
| rendered = "\n ".join(str(path) for path in missing) | |
| raise FileNotFoundError( | |
| "Required files are missing. Run utils/prepare_data.py and utils/save_init_checkpoint.py:\n " | |
| + rendered | |
| ) | |
| load_dir = resolve_load_dir(config["paths"]) | |
| init_checkpoint = config["paths"].get("init_checkpoint") | |
| loading_initial_weights = ( | |
| load_dir is not None | |
| and init_checkpoint is not None | |
| and load_dir == resolve(str(init_checkpoint)) | |
| ) | |
| if loading_initial_weights: | |
| extra = list(extra) | |
| if "--finetune" not in extra: | |
| # The random-init checkpoint is a weight container, not a training | |
| # resume point. Reset iteration/LR schedule to zero for the real run. | |
| extra += ["--finetune"] | |
| command = build_train_command( | |
| config, | |
| optimizer, | |
| extra_megatron_args=extra, | |
| python_executable=sys.executable, | |
| load_dir=load_dir, | |
| ) | |
| print("Launching:\n" + shlex.join(command), flush=True) | |
| if args.dry_run: | |
| return | |
| checkpoint_root = resolve(config["paths"]["output_dir"]) | |
| keep_last = int(config.get("checkpoint", {}).get("keep_last", 2)) | |
| stop_event = threading.Event() | |
| watcher = threading.Thread( | |
| target=_watch_checkpoints, | |
| args=(checkpoint_root, keep_last, stop_event, 30.0), | |
| daemon=True, | |
| ) | |
| watcher.start() | |
| try: | |
| env = os.environ.copy() | |
| env["PYTHONPATH"] = os.pathsep.join( | |
| part for part in (str(ROOT), env.get("PYTHONPATH")) if part | |
| ) | |
| subprocess.run(command, cwd=ROOT, check=True, env=env) | |
| finally: | |
| stop_event.set() | |
| watcher.join(timeout=5) | |
| prune_checkpoints(checkpoint_root, keep_last=keep_last) | |
| if args.skip_export: | |
| return | |
| from utils.export_hf import export_checkpoint | |
| export_checkpoint(config_path) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 6.3 kB
- Xet hash:
- 8e39434d9f7077314aa043ba5dad78780d1a477b24d2625f7a9e654bde2c4873
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.