92d38cbed4
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>
464 lines
15 KiB
Python
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()
|