Files
giant/giant/cli.py
T
lars 92d38cbed4 Add --batch-size auto to predict, matching train
Estimates batch size from free GPU memory using the checkpoint's
hidden_dim/n_blocks, same as the train command.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-06-22 08:23:58 +02:00

464 lines
15 KiB
Python

from enum import Enum
from pathlib import Path
from typing import Optional
import numpy as np
import torch
import typer
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,
PREDICT_COORD_METADATA_KEY,
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
)
from giant.data.loader import (
find_parquet_files,
iter_file_chunks,
iter_cond_chunks,
)
from giant.data.transforms import (
build_features,
build_cond_features,
inv_local_frame_rotation,
inv_log_transform,
reconstruct_post_pos,
Normalizer,
)
from giant.model.network import DenoisingMLP
from giant.pipeline import run_train_job
from giant.sample import sample_flow
app = typer.Typer(no_args_is_help=True)
@app.callback()
def _main() -> None:
"""GIANT — Geant4 step-function surrogate."""
class Mode(str, Enum):
flow = "flow"
ddpm = "ddpm"
class Coord(str, Enum):
global_ = "global"
local = "local"
@app.command()
def train(
data: Annotated[
Path, typer.Argument(help="Parquet file or directory of parquet files")
],
config: Annotated[
Optional[Path],
typer.Option(
"--config", "-c", help="TOML config file (overridden by explicit flags)"
),
] = None,
mode: Annotated[
Optional[Mode],
typer.Option("--mode", "-m", help="Generative model: flow matching or DDPM"),
] = None,
epochs: Annotated[Optional[int], typer.Option("--epochs", "-e")] = None,
batch_size: Annotated[
Optional[str],
typer.Option(
"--batch-size",
"-b",
help="Integer, or 'auto' to estimate from free GPU memory "
"(cuda devices only)",
),
] = None,
lr: Annotated[Optional[float], typer.Option("--lr", "-l")] = None,
warmup_epochs: Annotated[
Optional[int], typer.Option("--warmup-epochs", "-w")
] = None,
hidden_dim: Annotated[Optional[int], typer.Option("--hidden-dim", "-H")] = None,
n_blocks: Annotated[Optional[int], typer.Option("--n-blocks", "-n")] = None,
emb_dim: Annotated[Optional[int], typer.Option("--emb-dim", "-E")] = None,
dropout: Annotated[
Optional[float],
typer.Option(
"--dropout", "-d", help="Dropout probability in ResBlocks (default: 0.1)"
),
] = None,
val_fraction: Annotated[
Optional[float], typer.Option("--val-fraction", "-f")
] = None,
seed: Annotated[
Optional[int],
typer.Option("--seed", "-s", help="Random seed for reproducibility"),
] = None,
validate_every: Annotated[
Optional[int],
typer.Option(
"--validate-every",
"-v",
help="Run marginal+KL validation every N epochs (0 disables)",
),
] = None,
validate_steps: Annotated[
Optional[int],
typer.Option(
"--validate-steps",
"-t",
help="Flow matching ODE steps used during marginal validation "
"(ignored in ddpm mode, which always runs the full schedule)",
),
] = None,
shuffle_buffer: Annotated[
int,
typer.Option(
"--shuffle-buffer", "-B", help="Rows held in RAM per worker for shuffling"
),
] = 65536,
out: Annotated[
Optional[Path],
typer.Option(
"--out", "-o", help="Checkpoint dir (default: auto from hyperparams)"
),
] = None,
device: Annotated[
Optional[str],
typer.Option("--device", "-D", help="cpu | cuda | mps (default: auto)"),
] = None,
num_workers: Annotated[Optional[int], typer.Option("--num-workers", "-j")] = None,
resume: Annotated[
Optional[Path],
typer.Option("--resume", "-r", help="Checkpoint .pt to resume training from"),
] = None,
) -> None:
"""Train the GIANT surrogate model."""
batch_size_auto = False
batch_size_value: Optional[int] = None
if batch_size is not None:
if batch_size.strip().lower() == "auto":
batch_size_auto = True
else:
try:
batch_size_value = int(batch_size)
except ValueError:
typer.echo(
f"error: --batch-size must be an integer or 'auto', "
f"got {batch_size!r}",
err=True,
)
raise typer.Exit(1)
cli_train = {
k: v
for k, v in {
"mode": mode.value if mode is not None else None,
"epochs": epochs,
"batch_size": batch_size_value,
"lr": lr,
"warmup_epochs": warmup_epochs,
"val_fraction": val_fraction,
"num_workers": num_workers,
"seed": seed,
"validate_every": validate_every,
"validate_steps": validate_steps,
}.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,
"dropout": dropout,
}.items()
if v is not None
}
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 gconfig.auto_device()
if batch_size_auto:
try:
t["batch_size"] = gconfig.estimate_batch_size(
m["hidden_dim"], m["n_blocks"], _device
)
except ValueError as exc:
typer.echo(f"error: {exc}", err=True)
raise typer.Exit(1)
typer.echo(
f"batch_size: {t['batch_size']} (auto-estimated from free GPU memory)"
)
out_dir = out or Path(
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']}"
)
typer.echo(f"device: {_device}")
typer.echo(f"out_dir: {out_dir}")
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,
)
@app.command()
def predict(
data: Annotated[
Path, typer.Argument(help="Parquet file or directory of parquet files")
],
checkpoint: Annotated[
Path,
typer.Option(
"--checkpoint",
"-c",
help="Path to checkpoint .pt file (best.pt or last.pt)",
),
],
coord: Annotated[
Coord,
typer.Option(
"--coord",
"-C",
help="global: full physical units, world frame (default). "
"local: raw 9D model output (denormalised only, local frame, "
"log-scaled scalars) alongside the matching ground-truth target "
"for the same input file — requires post-step columns.",
),
] = Coord.global_,
out: Annotated[
Optional[Path],
typer.Option(
"--out",
"-o",
help="Output parquet path (default: <data>_predicted[_local].parquet)",
),
] = None,
batch_size: Annotated[
str,
typer.Option(
"--batch-size",
"-b",
help="Inference batch size, or 'auto' to estimate from free GPU "
"memory (cuda devices only)",
),
] = "4096",
steps: Annotated[
int, typer.Option("--steps", "-s", help="Flow matching ODE steps")
] = 10,
device: Annotated[
Optional[str],
typer.Option("--device", "-d", help="cpu | cuda | mps (default: auto)"),
] = None,
) -> None:
"""Run trained model on a parquet file and save predictions."""
batch_size_auto = False
batch_size_value: Optional[int] = None
if batch_size.strip().lower() == "auto":
batch_size_auto = True
else:
try:
batch_size_value = int(batch_size)
except ValueError:
typer.echo(
f"error: --batch-size must be an integer or 'auto', "
f"got {batch_size!r}",
err=True,
)
raise typer.Exit(1)
_device = torch.device(device) if device else gconfig.auto_device()
typer.echo(f"device: {_device}")
# --- Load checkpoint ---
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
if "model_config" not in ckpt:
typer.echo(
"error: checkpoint has no model_config — retrain with the current code",
err=True,
)
raise typer.Exit(1)
model_cfg = ckpt["model_config"]
if batch_size_auto:
try:
batch_size_value = gconfig.estimate_batch_size(
model_cfg["hidden_dim"], model_cfg["n_blocks"], _device
)
except ValueError as exc:
typer.echo(f"error: {exc}", err=True)
raise typer.Exit(1)
typer.echo(
f"batch_size: {batch_size_value} (auto-estimated from free GPU memory)"
)
assert batch_size_value is not None
bs = batch_size_value
pdg_map = {int(k): v for k, v in ckpt["pdg_map"].items()}
mat_map = {str(k): v for k, v in ckpt["mat_map"].items()}
cond_norm = Normalizer.from_dict(ckpt["normalizer"]["cond"])
tgt_norm = Normalizer.from_dict(ckpt["normalizer"]["target"])
model = DenoisingMLP(**model_cfg)
model.load_state_dict(ckpt["model"])
model.to(_device).eval()
typer.echo(f"loaded checkpoint: {checkpoint}")
gconfig.warn_if_checkpoint_config_mismatch(checkpoint)
# --- Output path ---
if out is None:
stem = data.stem if data.is_file() else data.name
suffix = (
"_predicted_local.parquet" if coord == Coord.local else "_predicted.parquet"
)
out = data.parent / f"{stem}{suffix}"
typer.echo(f"output: {out}")
# --- Stream input, generate predictions, write output ---
files = find_parquet_files(data)
typer.echo(f"found {len(files)} parquet file(s)")
writer: pq.ParquetWriter | None = None
total = 0
chunk_iter = iter_file_chunks if coord == Coord.local else iter_cond_chunks
for path in files:
for chunk in chunk_iter(path):
N = len(chunk["event_id"])
if coord == Coord.local:
cond_cont, cond_cat, target_raw, _, _ = build_features(
chunk, pdg_map, mat_map
)
cond_cont = cond_norm.transform(cond_cont)
else:
cond_cont, cond_cat = build_cond_features(
chunk, pdg_map, mat_map, cond_norm
)
# Inference in batch_size slices
pred_parts = []
for start in range(0, N, bs):
end = min(start + bs, N)
cc = torch.from_numpy(cond_cont[start:end]).float().to(_device)
ck = torch.from_numpy(cond_cat[start:end]).long().to(_device)
pred_parts.append(sample_flow(model, cc, ck, steps=steps).cpu().numpy())
pred = np.concatenate(pred_parts, axis=0) # (N, 9) normalised
# Inverse-normalise → local frame, log-scaled scalars
raw = tgt_norm.inverse_transform(pred)
if coord == Coord.local:
table = pa.table(
{
"event_id": chunk["event_id"],
"pdg": chunk["pdg"],
"pre_x": chunk["pre_pos"][:, 0],
"pre_y": chunk["pre_pos"][:, 1],
"pre_z": chunk["pre_pos"][:, 2],
"pre_E": chunk["pre_E"],
"pre_dx": chunk["pre_dir"][:, 0],
"pre_dy": chunk["pre_dir"][:, 1],
"pre_dz": chunk["pre_dir"][:, 2],
"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)
},
}
)
else:
step_length = inv_log_transform(raw[:, 0])
delta_e = inv_log_transform(raw[:, 1])
edep = inv_log_transform(raw[:, 2])
# Normalise predicted direction then rotate back to world frame
post_dir_local = raw[:, 3:6].copy()
norms = np.linalg.norm(post_dir_local, axis=1, keepdims=True)
post_dir_local /= np.where(norms < 1e-8, 1.0, norms)
post_dir_world = inv_local_frame_rotation(
chunk["pre_dir"], post_dir_local
)
# Same for the travel direction, then reconstruct post_pos from
# the single shared step_length so the two stay consistent.
travel_dir_local = raw[:, 6:9].copy()
norms = np.linalg.norm(travel_dir_local, axis=1, keepdims=True)
travel_dir_local /= np.where(norms < 1e-8, 1.0, norms)
post_pos_world = reconstruct_post_pos(
chunk["pre_pos"], chunk["pre_dir"], step_length, travel_dir_local
)
table = pa.table(
{
"event_id": chunk["event_id"],
"pdg": chunk["pdg"],
"pre_x": chunk["pre_pos"][:, 0],
"pre_y": chunk["pre_pos"][:, 1],
"pre_z": chunk["pre_pos"][:, 2],
"pre_E": chunk["pre_E"],
"pre_dx": chunk["pre_dir"][:, 0],
"pre_dy": chunk["pre_dir"][:, 1],
"pre_dz": chunk["pre_dir"][:, 2],
"material": chunk["material"],
"layer_id": chunk["layer_id"],
"n_sec": chunk["n_sec"],
"step_length": step_length,
"delta_e": delta_e,
"edep": edep,
"post_dx": post_dir_world[:, 0],
"post_dy": post_dir_world[:, 1],
"post_dz": post_dir_world[:, 2],
"post_x": post_pos_world[:, 0],
"post_y": post_pos_world[:, 1],
"post_z": post_pos_world[:, 2],
}
)
table = table.replace_schema_metadata(
{
PREDICT_COORD_METADATA_KEY: coord.value,
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
}
)
if writer is None:
writer = pq.ParquetWriter(out, table.schema)
writer.write_table(table)
total += N
if writer is not None:
writer.close()
typer.echo(f"wrote {total:,} rows → {out}")
if __name__ == "__main__":
app()