Files
giant/scripts/train.py
T
lars 8cebc4809d Apply ruff format and document lint/type tooling in CLAUDE.md
First repo-wide ruff format pass, plus a note in CLAUDE.md to run
ruff and ty periodically.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-06-18 17:40:27 +02:00

113 lines
3.3 KiB
Python

import argparse
from pathlib import Path
import torch
from giant import config as gconfig
from giant.pipeline import run_train_job
def main() -> None:
parser = argparse.ArgumentParser(description="Train GIANT surrogate model")
parser.add_argument("--config", default=None, help="Path to TOML config file")
parser.add_argument(
"--data", required=True, help="Path to parquet file or directory"
)
parser.add_argument("--mode", choices=["flow", "ddpm"])
parser.add_argument("--epochs", type=int)
parser.add_argument("--batch-size", type=int)
parser.add_argument("--lr", type=float)
parser.add_argument("--hidden-dim", type=int)
parser.add_argument("--n-blocks", type=int)
parser.add_argument("--emb-dim", type=int)
parser.add_argument(
"--dropout", type=float, help="Dropout probability in ResBlocks (default: 0.1)"
)
parser.add_argument("--val-fraction", type=float)
parser.add_argument("--seed", type=int, help="Random seed for reproducibility")
parser.add_argument(
"--validate-every",
type=int,
help="Run marginal+KL validation every N epochs (0 disables)",
)
parser.add_argument(
"--shuffle-buffer",
type=int,
default=65536,
help="Rows held in RAM for shuffling per worker (default: 65536)",
)
parser.add_argument(
"--out",
default=None,
help="Checkpoint output directory (default: auto from hyperparams)",
)
parser.add_argument(
"--device", default=None, help="cpu | cuda | mps (default: auto)"
)
parser.add_argument("--num-workers", type=int)
parser.add_argument(
"--resume", default=None, help="Checkpoint .pt to resume training from"
)
args = parser.parse_args()
cli_train = {
k: v
for k, v in {
"mode": args.mode,
"epochs": args.epochs,
"batch_size": args.batch_size,
"lr": args.lr,
"val_fraction": args.val_fraction,
"num_workers": args.num_workers,
"seed": args.seed,
"validate_every": args.validate_every,
}.items()
if v is not None
}
cli_model = {
k: v
for k, v in {
"hidden_dim": args.hidden_dim,
"n_blocks": args.n_blocks,
"emb_dim": args.emb_dim,
"dropout": args.dropout,
}.items()
if v is not None
}
config_path = Path(args.config) if args.config else None
cfg = gconfig.merge_cli_overrides(
gconfig.DEFAULT_CONFIG, config_path, cli_train, cli_model
)
t, m = cfg["train"], cfg["model"]
device = torch.device(args.device) if args.device else gconfig.auto_device()
out_dir = Path(
args.out
or (
f"checkpoints/{t['mode']}"
f"_h{m['hidden_dim']}"
f"_b{m['n_blocks']}"
f"_e{m['emb_dim']}"
f"_lr{t['lr']}"
f"_bs{t['batch_size']}"
)
)
print(f"device: {device}")
print(f"out_dir: {out_dir}")
run_train_job(
data=Path(args.data),
cfg=cfg,
out_dir=out_dir,
device=device,
shuffle_buffer=args.shuffle_buffer,
num_workers=t["num_workers"],
resume=Path(args.resume) if args.resume else None,
echo=print,
)
if __name__ == "__main__":
main()