| |
| from __future__ import annotations |
| import argparse |
| from pathlib import Path |
| import yaml |
| from sacflow.utils.config import load_yaml, save_yaml, set_by_path |
| from sacflow.utils.misc import seed_everything, ensure_dir |
| from sacflow.utils.distributed import init_distributed, cleanup, is_main_process, barrier |
| from sacflow.utils.wandb_utils import init_wandb, wandb_finish |
| from sacflow.engine.train_loop import run_training |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--config", required=True) |
| ap.add_argument("--name", default=None) |
| ap.add_argument("--resume", nargs="?", const="auto", default=None, help="Resume checkpoint path, or use --resume without a value for auto-detection.") |
| ap.add_argument("--opts", nargs="*", default=[], help="Override config values: key=value, e.g. train.epochs=220 optim.lr=1e-4") |
| args = ap.parse_args() |
| cfg = load_yaml(args.config) |
| for opt in args.opts: |
| if "=" not in opt: |
| raise ValueError(f"Invalid --opts entry {opt!r}; expected key=value") |
| key, value = opt.split("=", 1) |
| try: |
| parsed = yaml.safe_load(value) |
| except Exception: |
| parsed = value |
| set_by_path(cfg, key, parsed) |
| if args.resume is not None: |
| cfg.setdefault("train", {})["resume_checkpoint"] = args.resume |
| seed_everything(int(cfg.get("seed", 1337))) |
| device = init_distributed(cfg.get("distributed", {}).get("backend", "nccl")) |
| out_dir = ensure_dir(cfg["output_dir"]) |
| if is_main_process(): |
| save_yaml(cfg, out_dir / "resolved_config.yaml") |
| run = init_wandb(cfg, run_name=args.name or Path(cfg["output_dir"]).name) |
| try: |
| run_training(cfg, device, wandb_run=run) |
| finally: |
| barrier() |
| wandb_finish(run) |
| cleanup() |
|
|
| if __name__ == "__main__": |
| main() |
|
|