Add streaming data pipeline and giant CLI entry point

- Streaming pipeline: row-group-level parquet reading (PyArrow) so
  large files never fully land in RAM; Welford online algorithm for
  normalizer fitting; StreamingStepsDataset with shuffle buffer and
  multi-worker file striping; event-ID scan and vocab scan via cheap
  single-column reads
- giant/cli.py: typer-based CLI with `giant train` subcommand, mirroring
  scripts/train.py; --shuffle-buffer flag for RAM control
- pyproject.toml: add typer>=0.12 dependency and giant entry point
- train.py: replace len(loader.dataset) with local counters (compatible
  with IterableDataset which has no __len__)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-17 11:03:48 +02:00
parent 9277d79dff
commit 93c4d6b74d
8 changed files with 538 additions and 30 deletions
+68 -21
View File
@@ -3,12 +3,18 @@ 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 load_steps, build_index_maps
from giant.data.transforms import build_features
from giant.data.dataset import train_val_split
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
@@ -48,7 +54,7 @@ def _save_config(cfg: dict, out_dir: Path) -> None:
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("--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)
@@ -57,12 +63,14 @@ 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("--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)
args = parser.parse_args()
# Defaults, overridden by config file, then by explicit CLI flags.
# Defaults config file explicit CLI flags.
cfg: dict = {
"train": {
"mode": "flow", "epochs": 100, "batch_size": 4096, "lr": 3e-4,
@@ -91,7 +99,6 @@ def main() -> None:
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']}"
@@ -104,28 +111,69 @@ def main() -> None:
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")
# --- Discover files ---
files = find_parquet_files(args.data)
print(f"found {len(files)} parquet file(s)")
print("building features …")
cond_cont, cond_cat, target, cond_norm, tgt_norm = build_features(
data, pdg_map, mat_map, fit=True
)
# --- 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")
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=t["val_fraction"]
# --- 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(6)
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,
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,
shuffle=False,
)
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,
train_ds, batch_size=t["batch_size"],
num_workers=t["num_workers"], pin_memory=pin,
)
val_loader = DataLoader(
val_ds, batch_size=t["batch_size"], shuffle=False,
val_ds, batch_size=t["batch_size"],
num_workers=t["num_workers"], pin_memory=pin,
)
@@ -136,8 +184,7 @@ def main() -> None:
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")
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)