Implement Phase 1: full data pipeline, model, training, and config support

- Data pipeline: loader (parquet→numpy), transforms (log, local-frame
  Rodrigues rotation, Normalizer), StepsDataset with event-ID-based split
- Model: SinusoidalEmbedding, ConditionEncoder, ResBlock, DenoisingMLP
- Schedule: cosine DDPM and conditional flow matching loss (Lipman 2022)
- Samplers: flow (Euler ODE), DDPM ancestral, DDIM deterministic
- Training loop: AdamW + cosine LR, grad clipping, best-val checkpoint
- Validation: per-dimension marginal summary (normalised space)
- CLI: TOML config support with CLI-overrides; hyperparam-encoded output
  directory; config.toml with git hash saved into each run's checkpoint dir
- 21 unit tests covering transforms, network, flow/DDPM losses, dataset splits

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-17 10:48:03 +02:00
parent c3bf3abebf
commit 9277d79dff
16 changed files with 1693 additions and 10 deletions
+161 -1
View File
@@ -1 +1,161 @@
# CLI entry point: parse args, build dataset, instantiate model, call train loop.
import argparse
import subprocess
import tomllib
from pathlib import Path
import torch
from torch.utils.data import DataLoader
from giant.data.loader import load_steps, build_index_maps
from giant.data.transforms import build_features
from giant.data.dataset import train_val_split
from giant.model.network import DenoisingMLP
from giant.train import train as run_training
def _auto_device() -> torch.device:
if torch.cuda.is_available():
return torch.device("cuda")
if torch.backends.mps.is_available():
return torch.device("mps")
return torch.device("cpu")
def _git_hash() -> str:
try:
return subprocess.check_output(
["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL
).decode().strip()
except Exception:
return "unknown"
def _load_config(path: str) -> dict:
with open(path, "rb") as f:
return tomllib.load(f)
def _save_config(cfg: dict, out_dir: Path) -> None:
lines = [f"# git: {_git_hash()}", ""]
for section, values in cfg.items():
lines.append(f"[{section}]")
for k, v in values.items():
lines.append(f"{k:<12} = {repr(v) if isinstance(v, str) else v}")
lines.append("")
(out_dir / "config.toml").write_text("\n".join(lines))
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 steps parquet file")
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("--val-fraction", type=float)
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)
args = parser.parse_args()
# Defaults, overridden by config file, then by explicit CLI flags.
cfg: dict = {
"train": {
"mode": "flow", "epochs": 100, "batch_size": 4096, "lr": 3e-4,
"val_fraction": 0.1, "num_workers": 4,
},
"model": {
"hidden_dim": 256, "n_blocks": 6, "emb_dim": 16,
},
}
if args.config:
file_cfg = _load_config(args.config)
for section in ("train", "model"):
cfg[section].update(file_cfg.get(section, {}))
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,
}.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,
}.items() if v is not None}
cfg["train"].update(cli_train)
cfg["model"].update(cli_model)
t, m = cfg["train"], cfg["model"]
device = torch.device(args.device) if args.device else _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}")
print("loading data …")
data = load_steps(args.data)
pdg_map, mat_map = build_index_maps(data)
print(f" {len(data['event_id']):,} steps | {len(pdg_map)} PDG codes | {len(mat_map)} materials")
print("building features …")
cond_cont, cond_cat, target, cond_norm, tgt_norm = build_features(
data, pdg_map, mat_map, fit=True
)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=t["val_fraction"]
)
print(f" train: {len(train_ds):,} val: {len(val_ds):,}")
pin = device.type == "cuda"
train_loader = DataLoader(
train_ds, batch_size=t["batch_size"], shuffle=True,
num_workers=t["num_workers"], pin_memory=pin,
)
val_loader = DataLoader(
val_ds, batch_size=t["batch_size"], shuffle=False,
num_workers=t["num_workers"], pin_memory=pin,
)
model = DenoisingMLP(
pdg_vocab=len(pdg_map),
mat_vocab=len(mat_map),
hidden_dim=m["hidden_dim"],
n_blocks=m["n_blocks"],
emb_dim=m["emb_dim"],
)
n_params = sum(p.numel() for p in model.parameters())
print(f"model: {n_params:,} parameters")
out_dir.mkdir(parents=True, exist_ok=True)
_save_config(cfg, out_dir)
run_training(
model=model,
train_loader=train_loader,
val_loader=val_loader,
mode=t["mode"],
epochs=t["epochs"],
lr=t["lr"],
device=device,
out_dir=out_dir,
normalizer_dict={"cond": cond_norm.to_dict(), "target": tgt_norm.to_dict()},
pdg_map={str(k): v for k, v in pdg_map.items()},
mat_map={str(k): v for k, v in mat_map.items()},
)
if __name__ == "__main__":
main()