Dedup training pipeline, add seeding/resume and per-epoch metrics logging

cli.py and scripts/train.py duplicated ~140 lines of training setup and had
drifted (scripts/train.py forgot to save model_config, breaking predict on
those checkpoints). Extract shared logic into giant/constants.py (X_DIM,
target names), giant/config.py (device/git/TOML/seeding helpers, run
metadata), and giant/pipeline.py (the actual training-job orchestration),
so both entry points become thin CLI wrappers around the same code path.

Also adds --seed/--resume support (checkpoints now carry optimizer/scheduler
state, epoch, and best_val_loss), a richer [meta] section in the saved
config.toml (git hash, seed, versions, timestamp, invocation, dataset
stats), and a metrics.csv (train/val loss, lr, epoch time) written every
epoch and append-safe across resumes.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-18 13:32:37 +02:00
parent f1a82b5853
commit 43634ef77a
9 changed files with 340 additions and 349 deletions
+15 -159
View File
@@ -1,5 +1,3 @@
import subprocess
import tomllib
from enum import Enum
from pathlib import Path
from typing import Optional
@@ -7,18 +5,17 @@ from typing import Optional
import numpy as np
import torch
import typer
from torch.utils.data import DataLoader
from typing_extensions import Annotated
import pyarrow as pa
import pyarrow.parquet as pq
from giant import config as gconfig
from giant.constants import LOCAL_TARGET_NAMES
from giant.data.loader import (
find_parquet_files,
load_event_ids,
iter_file_chunks,
iter_cond_chunks,
build_index_maps_from_files,
)
from giant.data.transforms import (
build_features,
@@ -26,28 +23,14 @@ from giant.data.transforms import (
inv_local_frame_rotation,
inv_log_transform,
reconstruct_post_pos,
_WelfordAccumulator,
Normalizer,
)
from giant.data.dataset import make_event_split, StreamingStepsDataset
from giant.model.network import DenoisingMLP
from giant.pipeline import run_train_job
from giant.sample import sample_flow
from giant.train import train as run_training
app = typer.Typer(no_args_is_help=True)
_LOCAL_TARGET_NAMES = [
"log_step_length",
"log_delta_e",
"log_edep",
"post_dx",
"post_dy",
"post_dz",
"travel_dx",
"travel_dy",
"travel_dz",
]
@app.callback()
def _main() -> None:
@@ -64,38 +47,6 @@ class Coord(str, Enum):
local = "local"
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_toml(path: Path) -> 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))
@app.command()
def train(
data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")],
@@ -108,42 +59,26 @@ def train(
n_blocks: Annotated[Optional[int], typer.Option()] = None,
emb_dim: Annotated[Optional[int], typer.Option()] = None,
val_fraction: Annotated[Optional[float], typer.Option()] = None,
seed: Annotated[Optional[int], typer.Option(help="Random seed for reproducibility")] = None,
shuffle_buffer: Annotated[int, typer.Option(help="Rows held in RAM per worker for shuffling")] = 65536,
out: Annotated[Optional[Path], typer.Option(help="Checkpoint dir (default: auto from hyperparams)")] = None,
device: Annotated[Optional[str], typer.Option(help="cpu | cuda | mps (default: auto)")] = None,
num_workers: Annotated[Optional[int], typer.Option()] = None,
resume: Annotated[Optional[Path], typer.Option(help="Checkpoint .pt to resume training from")] = None,
) -> None:
"""Train the GIANT surrogate model."""
# Defaults → config file → 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 config is not None:
file_cfg = _load_toml(config)
for section in ("train", "model"):
cfg[section].update(file_cfg.get(section, {}))
cli_train = {k: v for k, v in {
"mode": mode.value if mode is not None else None,
"epochs": epochs, "batch_size": batch_size, "lr": lr,
"val_fraction": val_fraction, "num_workers": num_workers,
"val_fraction": val_fraction, "num_workers": num_workers, "seed": seed,
}.items() if v is not None}
cli_model = {k: v for k, v in {
"hidden_dim": hidden_dim, "n_blocks": n_blocks, "emb_dim": emb_dim,
}.items() if v is not None}
cfg["train"].update(cli_train)
cfg["model"].update(cli_model)
cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, cli_train, cli_model)
t, m = cfg["train"], cfg["model"]
_device = torch.device(device) if device else _auto_device()
_device = torch.device(device) if device else gconfig.auto_device()
out_dir = out or Path(
f"checkpoints/{t['mode']}"
f"_h{m['hidden_dim']}"
@@ -156,89 +91,10 @@ def train(
typer.echo(f"device: {_device}")
typer.echo(f"out_dir: {out_dir}")
files = find_parquet_files(data)
typer.echo(f"found {len(files)} parquet file(s)")
typer.echo("scanning event IDs …")
all_event_ids = np.concatenate([load_event_ids(f) for f in files])
train_events, val_events = make_event_split(all_event_ids, val_fraction=t["val_fraction"])
events_arr = np.array(sorted(train_events))
n_train_steps = int(np.isin(all_event_ids, events_arr).sum())
typer.echo(
f" {len(all_event_ids):,} steps | "
f"{len(train_events)} train events (~{n_train_steps:,} steps) | "
f"{len(val_events)} val events"
)
typer.echo("building vocabulary maps …")
pdg_map, mat_map = build_index_maps_from_files(files)
typer.echo(f" {len(pdg_map)} PDG codes | {len(mat_map)} materials")
typer.echo("fitting normalizer (streaming) …")
cond_acc = _WelfordAccumulator(9)
tgt_acc = _WelfordAccumulator(9)
for path in files:
for chunk in iter_file_chunks(path):
mask = np.isin(chunk["event_id"], events_arr)
if not mask.any():
continue
chunk_tr = {k: v[mask] for k, v in chunk.items()}
cond_cont, _, target, _, _ = build_features(chunk_tr, pdg_map, mat_map)
cond_acc.update(cond_cont)
tgt_acc.update(target)
cond_norm = cond_acc.to_normalizer()
tgt_norm = tgt_acc.to_normalizer()
train_ds = StreamingStepsDataset(
files=files, split_events=train_events,
pdg_map=pdg_map, mat_map=mat_map,
cond_normalizer=cond_norm, target_normalizer=tgt_norm,
batch_size=t["batch_size"],
shuffle_buffer=shuffle_buffer, shuffle=True,
)
val_ds = StreamingStepsDataset(
files=files, split_events=val_events,
pdg_map=pdg_map, mat_map=mat_map,
cond_normalizer=cond_norm, target_normalizer=tgt_norm,
batch_size=t["batch_size"],
shuffle=False,
)
# Dataset yields whole batches already, so batch_size=None tells DataLoader
# to pass them through instead of re-collating row-by-row in Python.
pin = _device.type == "cuda"
train_loader = DataLoader(
train_ds, batch_size=None,
num_workers=t["num_workers"], pin_memory=pin,
)
val_loader = DataLoader(
val_ds, batch_size=None,
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"],
)
typer.echo(f"model: {sum(p.numel() for p in model.parameters()):,} parameters")
out_dir.mkdir(parents=True, exist_ok=True)
_save_config(cfg, out_dir)
model_config = {
"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"],
}
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()},
model_config=model_config,
run_train_job(
data=data, cfg=cfg, out_dir=out_dir, device=_device,
shuffle_buffer=shuffle_buffer, num_workers=t["num_workers"],
resume=resume, echo=typer.echo,
)
@@ -258,7 +114,7 @@ def predict(
device: Annotated[Optional[str], typer.Option(help="cpu | cuda | mps (default: auto)")] = None,
) -> None:
"""Run trained model on a parquet file and save predictions."""
_device = torch.device(device) if device else _auto_device()
_device = torch.device(device) if device else gconfig.auto_device()
typer.echo(f"device: {_device}")
# --- Load checkpoint ---
@@ -329,8 +185,8 @@ def predict(
"material": chunk["material"],
"layer_id": chunk["layer_id"],
"n_sec": chunk["n_sec"],
**{f"pred_{name}": raw[:, j] for j, name in enumerate(_LOCAL_TARGET_NAMES)},
**{f"true_{name}": target_raw[:, j] for j, name in enumerate(_LOCAL_TARGET_NAMES)},
**{f"pred_{name}": raw[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)},
**{f"true_{name}": target_raw[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)},
})
else:
step_length = inv_log_transform(raw[:, 0])
+106
View File
@@ -0,0 +1,106 @@
import random
import subprocess
import sys
import tomllib
from datetime import datetime, timezone
from pathlib import Path
import numpy as np
import torch
DEFAULT_CONFIG: dict = {
"train": {
"mode": "flow", "epochs": 100, "batch_size": 4096, "lr": 3e-4,
"val_fraction": 0.1, "num_workers": 4, "seed": 0,
},
"model": {
"hidden_dim": 256, "n_blocks": 6, "emb_dim": 16,
},
}
def git_hash() -> str:
try:
return subprocess.check_output(
["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL
).decode().strip()
except Exception:
return "unknown"
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 load_toml(path: Path) -> dict:
with open(path, "rb") as f:
return tomllib.load(f)
def merge_cli_overrides(
defaults: dict,
config_path: Path | None,
train_overrides: dict,
model_overrides: dict,
) -> dict:
"""Resolve config as defaults -> TOML file -> explicit CLI flags."""
cfg = {"train": dict(defaults["train"]), "model": dict(defaults["model"])}
if config_path is not None:
file_cfg = load_toml(config_path)
for section in ("train", "model"):
cfg[section].update(file_cfg.get(section, {}))
cfg["train"].update(train_overrides)
cfg["model"].update(model_overrides)
return cfg
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def save_config(cfg: dict, out_dir: Path, meta: dict) -> None:
lines = []
for section, values in cfg.items():
lines.append(f"[{section}]")
for k, v in values.items():
lines.append(f"{k:<14} = {repr(v) if isinstance(v, str) else v}")
lines.append("")
lines.append("[meta]")
for k, v in meta.items():
lines.append(f"{k:<14} = {repr(v) if isinstance(v, str) else v}")
(out_dir / "config.toml").write_text("\n".join(lines))
def build_run_meta(
data: Path,
seed: int,
n_pdg_codes: int,
n_materials: int,
n_train_events: int,
n_val_events: int,
n_train_steps: int,
) -> dict:
return {
"git_hash": git_hash(),
"seed": seed,
"timestamp_utc": datetime.now(timezone.utc).isoformat(timespec="seconds"),
"python_version": sys.version.split()[0],
"torch_version": torch.__version__,
"command": " ".join(sys.argv),
"data_path": str(data),
"n_pdg_codes": n_pdg_codes,
"n_materials": n_materials,
"n_train_events": n_train_events,
"n_val_events": n_val_events,
"n_train_steps": n_train_steps,
}
+13
View File
@@ -0,0 +1,13 @@
X_DIM = 9
LOCAL_TARGET_NAMES = [
"log_step_length",
"log_delta_e",
"log_edep",
"post_dx",
"post_dy",
"post_dz",
"travel_dx",
"travel_dy",
"travel_dz",
]
+3 -1
View File
@@ -3,6 +3,8 @@ import math
import torch
import torch.nn as nn
from giant.constants import X_DIM
class SinusoidalEmbedding(nn.Module):
def __init__(self, dim: int) -> None:
@@ -73,7 +75,7 @@ class DenoisingMLP(nn.Module):
emb_dim: int = 16,
time_dim: int = 64,
cond_out_dim: int = 128,
x_dim: int = 9,
x_dim: int = X_DIM,
) -> None:
super().__init__()
self.time_emb = SinusoidalEmbedding(time_dim)
+129
View File
@@ -0,0 +1,129 @@
from pathlib import Path
import numpy as np
import torch
from torch.utils.data import DataLoader
from giant import config
from giant.constants import X_DIM
from giant.data.loader import (
find_parquet_files,
load_event_ids,
iter_file_chunks,
build_index_maps_from_files,
)
from giant.data.transforms import build_features, _WelfordAccumulator
from giant.data.dataset import make_event_split, StreamingStepsDataset
from giant.model.network import DenoisingMLP
from giant.train import train as run_training
def run_train_job(
data: Path,
cfg: dict,
out_dir: Path,
device: torch.device,
shuffle_buffer: int,
num_workers: int,
resume: Path | None = None,
echo=print,
) -> None:
t, m = cfg["train"], cfg["model"]
config.seed_everything(t["seed"])
out_dir = Path(out_dir)
files = find_parquet_files(data)
echo(f"found {len(files)} parquet file(s)")
echo("scanning event IDs …")
all_event_ids = np.concatenate([load_event_ids(f) for f in files])
train_events, val_events = make_event_split(all_event_ids, val_fraction=t["val_fraction"])
events_arr = np.array(sorted(train_events))
n_train_steps = int(np.isin(all_event_ids, events_arr).sum())
echo(
f" {len(all_event_ids):,} steps | "
f"{len(train_events)} train events (~{n_train_steps:,} steps) | "
f"{len(val_events)} val events"
)
echo("building vocabulary maps …")
pdg_map, mat_map = build_index_maps_from_files(files)
echo(f" {len(pdg_map)} PDG codes | {len(mat_map)} materials")
echo("fitting normalizer (streaming) …")
cond_acc = _WelfordAccumulator(X_DIM)
tgt_acc = _WelfordAccumulator(X_DIM)
for path in files:
for chunk in iter_file_chunks(path):
mask = np.isin(chunk["event_id"], events_arr)
if not mask.any():
continue
chunk_tr = {k: v[mask] for k, v in chunk.items()}
cond_cont, _, target, _, _ = build_features(chunk_tr, pdg_map, mat_map)
cond_acc.update(cond_cont)
tgt_acc.update(target)
cond_norm = cond_acc.to_normalizer()
tgt_norm = tgt_acc.to_normalizer()
train_ds = StreamingStepsDataset(
files=files, split_events=train_events,
pdg_map=pdg_map, mat_map=mat_map,
cond_normalizer=cond_norm, target_normalizer=tgt_norm,
batch_size=t["batch_size"],
shuffle_buffer=shuffle_buffer, shuffle=True,
)
val_ds = StreamingStepsDataset(
files=files, split_events=val_events,
pdg_map=pdg_map, mat_map=mat_map,
cond_normalizer=cond_norm, target_normalizer=tgt_norm,
batch_size=t["batch_size"],
shuffle=False,
)
# Dataset yields whole batches already, so batch_size=None tells DataLoader
# to pass them through instead of re-collating row-by-row in Python.
pin = device.type == "cuda"
train_loader = DataLoader(
train_ds, batch_size=None,
num_workers=num_workers, pin_memory=pin,
)
val_loader = DataLoader(
val_ds, batch_size=None,
num_workers=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"],
)
echo(f"model: {sum(p.numel() for p in model.parameters()):,} parameters")
out_dir.mkdir(parents=True, exist_ok=True)
meta = config.build_run_meta(
data=data,
seed=t["seed"],
n_pdg_codes=len(pdg_map),
n_materials=len(mat_map),
n_train_events=len(train_events),
n_val_events=len(val_events),
n_train_steps=n_train_steps,
)
config.save_config(cfg, out_dir, meta)
model_config = {
"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"],
}
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()},
model_config=model_config,
resume_path=resume,
)
+5 -3
View File
@@ -1,5 +1,7 @@
import torch
from giant.constants import X_DIM
@torch.no_grad()
def sample_flow(
@@ -12,7 +14,7 @@ def sample_flow(
model.eval()
B = cond_cont.size(0)
device = cond_cont.device
x = torch.randn(B, 9, device=device)
x = torch.randn(B, X_DIM, device=device)
dt = 1.0 / steps
for i in range(steps):
t = torch.full((B,), i * dt, device=device)
@@ -32,7 +34,7 @@ def sample_ddpm(
model.eval()
B = cond_cont.size(0)
device = cond_cont.device
x = torch.randn(B, 9, device=device)
x = torch.randn(B, X_DIM, device=device)
T = schedule.T
for i in reversed(range(T)):
t_norm = torch.full((B,), i / T, device=device)
@@ -63,7 +65,7 @@ def sample_ddim(
device = cond_cont.device
T = schedule.T
timesteps = torch.linspace(T - 1, 0, steps, dtype=torch.long, device=device)
x = torch.randn(B, 9, device=device)
x = torch.randn(B, X_DIM, device=device)
for step_idx, ts in enumerate(timesteps):
t_idx = int(ts.item())
t_norm = torch.full((B,), t_idx / T, device=device)
+55 -13
View File
@@ -1,3 +1,5 @@
import csv
import time
from pathlib import Path
import torch
@@ -6,6 +8,8 @@ from torch.utils.data import DataLoader
from giant.model.schedule import CosineSchedule, flow_matching_loss
_METRICS_FIELDS = ["epoch", "train_loss", "val_loss", "lr", "epoch_time_s"]
def train(
model: torch.nn.Module,
@@ -20,6 +24,7 @@ def train(
pdg_map: dict | None = None,
mat_map: dict | None = None,
model_config: dict | None = None,
resume_path: str | Path | None = None,
) -> None:
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
@@ -29,9 +34,27 @@ def train(
lr_sched = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
ddpm_schedule = CosineSchedule().to(device) if mode == "ddpm" else None
best_val_loss = float("inf")
for epoch in range(1, epochs + 1):
start_epoch = 1
best_val_loss = float("inf")
if resume_path is not None:
ckpt = torch.load(resume_path, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
lr_sched.load_state_dict(ckpt["lr_sched"])
start_epoch = ckpt.get("epoch", 0) + 1
best_val_loss = ckpt.get("best_val_loss", float("inf"))
metrics_path = out_dir / "metrics.csv"
write_header = not (resume_path is not None and metrics_path.exists())
metrics_file = open(metrics_path, "a", newline="")
metrics_writer = csv.DictWriter(metrics_file, fieldnames=_METRICS_FIELDS)
if write_header:
metrics_writer.writeheader()
for epoch in range(start_epoch, epochs + 1):
epoch_start = time.monotonic()
current_lr = optimizer.param_groups[0]["lr"]
model.train()
train_loss_sum = 0.0
train_n = 0
@@ -70,20 +93,39 @@ def train(
val_loss_sum += loss.item() * x1.size(0)
val_n += x1.size(0)
val_loss = val_loss_sum / max(val_n, 1)
epoch_time = time.monotonic() - epoch_start
print(f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f}")
print(
f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f} "
f"lr {current_lr:.2e} {epoch_time:.1f}s"
)
metrics_writer.writerow({
"epoch": epoch, "train_loss": train_loss, "val_loss": val_loss,
"lr": current_lr, "epoch_time_s": epoch_time,
})
metrics_file.flush()
ckpt: dict = {
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"lr_sched": lr_sched.state_dict(),
"epoch": epoch,
"best_val_loss": best_val_loss,
}
if normalizer_dict is not None:
ckpt["normalizer"] = normalizer_dict
if pdg_map is not None:
ckpt["pdg_map"] = pdg_map
if mat_map is not None:
ckpt["mat_map"] = mat_map
if model_config is not None:
ckpt["model_config"] = model_config
if val_loss < best_val_loss:
best_val_loss = val_loss
ckpt: dict = {"model": model.state_dict()}
if normalizer_dict is not None:
ckpt["normalizer"] = normalizer_dict
if pdg_map is not None:
ckpt["pdg_map"] = pdg_map
if mat_map is not None:
ckpt["mat_map"] = mat_map
if model_config is not None:
ckpt["model_config"] = model_config
ckpt["best_val_loss"] = best_val_loss
torch.save(ckpt, out_dir / "best.pt")
torch.save({"model": model.state_dict()}, out_dir / "last.pt")
torch.save(ckpt, out_dir / "last.pt")
metrics_file.close()
+2 -13
View File
@@ -2,20 +2,9 @@ import numpy as np
import torch
from torch.utils.data import DataLoader
from giant.constants import LOCAL_TARGET_NAMES
from giant.sample import sample_flow, sample_ddpm, sample_ddim
_TARGET_NAMES = [
"log_step_length",
"log_delta_e",
"log_edep",
"post_dx",
"post_dy",
"post_dz",
"travel_dx",
"travel_dy",
"travel_dz",
]
def validate_marginals(
model: torch.nn.Module,
@@ -56,7 +45,7 @@ def validate_marginals(
header = f"{'Dim':<20} {'real_mean':>10} {'gen_mean':>10} {'real_std':>10} {'gen_std':>10}"
print(f"\n{header}")
print("-" * len(header))
for j, name in enumerate(_TARGET_NAMES):
for j, name in enumerate(LOCAL_TARGET_NAMES):
r, g = real[:, j], generated[:, j]
print(f"{name:<20} {r.mean():>10.4f} {g.mean():>10.4f} {r.std():>10.4f} {g.std():>10.4f}")
+12 -160
View File
@@ -1,54 +1,10 @@
import argparse
import subprocess
import tomllib
from pathlib import Path
import numpy as np
import torch
from torch.utils.data import DataLoader
from giant.data.loader import (
find_parquet_files,
load_event_ids,
iter_file_chunks,
build_index_maps_from_files,
)
from giant.data.transforms import build_features, _WelfordAccumulator
from giant.data.dataset import make_event_split, StreamingStepsDataset
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))
from giant import config as gconfig
from giant.pipeline import run_train_job
def main() -> None:
@@ -63,42 +19,28 @@ def main() -> None:
parser.add_argument("--n-blocks", type=int)
parser.add_argument("--emb-dim", type=int)
parser.add_argument("--val-fraction", type=float)
parser.add_argument("--seed", type=int, help="Random seed for reproducibility")
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()
# Defaults → config file → 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,
"seed": args.seed,
}.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)
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 _auto_device()
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']}"
@@ -111,100 +53,10 @@ def main() -> None:
print(f"device: {device}")
print(f"out_dir: {out_dir}")
# --- Discover files ---
files = find_parquet_files(args.data)
print(f"found {len(files)} parquet file(s)")
# --- Scan event IDs (single column, cheap) ---
print("scanning event IDs …")
all_event_ids = np.concatenate([load_event_ids(f) for f in files])
train_events, val_events = make_event_split(all_event_ids, val_fraction=t["val_fraction"])
n_train_steps = (np.isin(all_event_ids, np.array(sorted(train_events)))).sum()
print(f" {len(all_event_ids):,} steps | "
f"{len(train_events)} train events (~{n_train_steps:,} steps) | "
f"{len(val_events)} val events")
# --- Scan PDG / material vocabularies (2 columns, cheap) ---
print("building vocabulary maps …")
pdg_map, mat_map = build_index_maps_from_files(files)
print(f" {len(pdg_map)} PDG codes | {len(mat_map)} materials")
# --- Streaming normalizer fit over training data ---
print("fitting normalizer (streaming) …")
events_arr = np.array(sorted(train_events))
cond_acc = _WelfordAccumulator(9)
tgt_acc = _WelfordAccumulator(9)
for path in files:
for chunk in iter_file_chunks(path):
mask = np.isin(chunk["event_id"], events_arr)
if not mask.any():
continue
chunk_tr = {k: v[mask] for k, v in chunk.items()}
cond_cont, _, target, _, _ = build_features(chunk_tr, pdg_map, mat_map)
cond_acc.update(cond_cont)
tgt_acc.update(target)
cond_norm = cond_acc.to_normalizer()
tgt_norm = tgt_acc.to_normalizer()
# --- Streaming datasets ---
train_ds = StreamingStepsDataset(
files=files,
split_events=train_events,
pdg_map=pdg_map,
mat_map=mat_map,
cond_normalizer=cond_norm,
target_normalizer=tgt_norm,
batch_size=t["batch_size"],
shuffle_buffer=args.shuffle_buffer,
shuffle=True,
)
val_ds = StreamingStepsDataset(
files=files,
split_events=val_events,
pdg_map=pdg_map,
mat_map=mat_map,
cond_normalizer=cond_norm,
target_normalizer=tgt_norm,
batch_size=t["batch_size"],
shuffle=False,
)
# Dataset yields whole batches already, so batch_size=None tells DataLoader
# to pass them through instead of re-collating row-by-row in Python.
pin = device.type == "cuda"
train_loader = DataLoader(
train_ds, batch_size=None,
num_workers=t["num_workers"], pin_memory=pin,
)
val_loader = DataLoader(
val_ds, batch_size=None,
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"],
)
print(f"model: {sum(p.numel() for p in model.parameters()):,} 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()},
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,
)