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:
+68
-21
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user