xuan-luo/temp / utils /train_runner.py
xuan-luo's picture
download
raw
6.3 kB
#!/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.