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:
+202
@@ -0,0 +1,202 @@
|
||||
import subprocess
|
||||
import tomllib
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import typer
|
||||
from torch.utils.data import DataLoader
|
||||
from typing_extensions import Annotated
|
||||
|
||||
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
|
||||
|
||||
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"
|
||||
|
||||
|
||||
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")],
|
||||
config: Annotated[Optional[Path], typer.Option(help="TOML config file (overridden by explicit flags)")] = None,
|
||||
mode: Annotated[Optional[Mode], typer.Option(help="Generative model: flow matching or DDPM")] = None,
|
||||
epochs: Annotated[Optional[int], typer.Option()] = None,
|
||||
batch_size: Annotated[Optional[int], typer.Option()] = None,
|
||||
lr: Annotated[Optional[float], typer.Option()] = None,
|
||||
hidden_dim: Annotated[Optional[int], typer.Option()] = None,
|
||||
n_blocks: Annotated[Optional[int], typer.Option()] = None,
|
||||
emb_dim: Annotated[Optional[int], typer.Option()] = None,
|
||||
val_fraction: Annotated[Optional[float], typer.Option()] = 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,
|
||||
) -> 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,
|
||||
}.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)
|
||||
|
||||
t, m = cfg["train"], cfg["model"]
|
||||
|
||||
_device = torch.device(device) if device else _auto_device()
|
||||
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}")
|
||||
|
||||
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(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()
|
||||
|
||||
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=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,
|
||||
)
|
||||
|
||||
pin = _device.type == "cuda"
|
||||
train_loader = DataLoader(
|
||||
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"],
|
||||
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)
|
||||
|
||||
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()},
|
||||
)
|
||||
+107
-1
@@ -1,8 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
from torch.utils.data import Dataset, IterableDataset
|
||||
|
||||
from giant.data.loader import iter_file_chunks
|
||||
from giant.data.transforms import Normalizer, build_features
|
||||
|
||||
|
||||
class StepsDataset(Dataset):
|
||||
@@ -44,3 +49,104 @@ def train_val_split(
|
||||
StepsDataset(cond_cont[train_mask], cond_cat[train_mask], target[train_mask]),
|
||||
StepsDataset(cond_cont[val_mask], cond_cat[val_mask], target[val_mask]),
|
||||
)
|
||||
|
||||
|
||||
def make_event_split(
|
||||
all_event_ids: np.ndarray,
|
||||
val_fraction: float = 0.1,
|
||||
seed: int = 42,
|
||||
) -> tuple[set, set]:
|
||||
"""Assign unique event_ids to train/val sets by event_id, not by row."""
|
||||
rng = np.random.default_rng(seed)
|
||||
unique = np.unique(all_event_ids)
|
||||
rng.shuffle(unique)
|
||||
n_val = max(1, int(len(unique) * val_fraction))
|
||||
val_set = set(unique[:n_val].tolist())
|
||||
train_set = set(unique[n_val:].tolist())
|
||||
return train_set, val_set
|
||||
|
||||
|
||||
class StreamingStepsDataset(IterableDataset):
|
||||
"""Streams parquet files one row-group at a time.
|
||||
|
||||
Never loads more than `shuffle_buffer` rows into RAM simultaneously.
|
||||
Files are split evenly across DataLoader workers via worker_info.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
files: list[Path],
|
||||
split_events: set,
|
||||
pdg_map: dict[int, int],
|
||||
mat_map: dict[int, int],
|
||||
cond_normalizer: Normalizer,
|
||||
target_normalizer: Normalizer,
|
||||
shuffle_buffer: int = 65536,
|
||||
shuffle: bool = True,
|
||||
) -> None:
|
||||
self.files = list(files)
|
||||
self.split_events = split_events
|
||||
self._events_arr = np.array(sorted(split_events)) # for np.isin
|
||||
self.pdg_map = pdg_map
|
||||
self.mat_map = mat_map
|
||||
self.cond_normalizer = cond_normalizer
|
||||
self.target_normalizer = target_normalizer
|
||||
self.shuffle_buffer = shuffle_buffer
|
||||
self.shuffle = shuffle
|
||||
|
||||
def __iter__(self):
|
||||
worker_info = torch.utils.data.get_worker_info()
|
||||
files = self.files
|
||||
if worker_info is not None:
|
||||
files = files[worker_info.id :: worker_info.num_workers]
|
||||
|
||||
if self.shuffle:
|
||||
files = list(files)
|
||||
np.random.default_rng().shuffle(files)
|
||||
|
||||
buf_cont: list[np.ndarray] = []
|
||||
buf_cat: list[np.ndarray] = []
|
||||
buf_tgt: list[np.ndarray] = []
|
||||
buf_n = 0
|
||||
|
||||
for path in files:
|
||||
for chunk in iter_file_chunks(path):
|
||||
mask = np.isin(chunk["event_id"], self._events_arr)
|
||||
if not mask.any():
|
||||
continue
|
||||
chunk = {k: v[mask] for k, v in chunk.items()}
|
||||
|
||||
cond_cont, cond_cat, target, _, _ = build_features(
|
||||
chunk, self.pdg_map, self.mat_map,
|
||||
cond_normalizer=self.cond_normalizer,
|
||||
target_normalizer=self.target_normalizer,
|
||||
)
|
||||
buf_cont.append(cond_cont)
|
||||
buf_cat.append(cond_cat)
|
||||
buf_tgt.append(target)
|
||||
buf_n += len(cond_cont)
|
||||
|
||||
if not self.shuffle or buf_n >= self.shuffle_buffer:
|
||||
yield from self._flush(buf_cont, buf_cat, buf_tgt)
|
||||
buf_cont, buf_cat, buf_tgt, buf_n = [], [], [], 0
|
||||
|
||||
if buf_n > 0:
|
||||
yield from self._flush(buf_cont, buf_cat, buf_tgt)
|
||||
|
||||
def _flush(
|
||||
self,
|
||||
buf_cont: list[np.ndarray],
|
||||
buf_cat: list[np.ndarray],
|
||||
buf_tgt: list[np.ndarray],
|
||||
):
|
||||
cont = np.concatenate(buf_cont)
|
||||
cat = np.concatenate(buf_cat)
|
||||
tgt = np.concatenate(buf_tgt)
|
||||
if self.shuffle:
|
||||
idx = np.random.permutation(len(cont))
|
||||
cont, cat, tgt = cont[idx], cat[idx], tgt[idx]
|
||||
t_cont = torch.from_numpy(cont).float()
|
||||
t_cat = torch.from_numpy(cat).long()
|
||||
t_tgt = torch.from_numpy(tgt).float()
|
||||
for i in range(len(cont)):
|
||||
yield t_cont[i], t_cat[i], t_tgt[i]
|
||||
|
||||
+45
-2
@@ -1,11 +1,22 @@
|
||||
from pathlib import Path
|
||||
from typing import Iterator
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
|
||||
def load_steps(path: str | Path) -> dict[str, np.ndarray]:
|
||||
df = pd.read_parquet(path)
|
||||
def find_parquet_files(path: str | Path) -> list[Path]:
|
||||
p = Path(path)
|
||||
if p.is_dir():
|
||||
files = sorted(p.glob("*.parquet"))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"no .parquet files found in {p}")
|
||||
return files
|
||||
return [p]
|
||||
|
||||
|
||||
def _df_to_dict(df: pd.DataFrame) -> dict[str, np.ndarray]:
|
||||
return {
|
||||
"event_id": df["event_id"].to_numpy(),
|
||||
"pdg": df["pdg"].to_numpy(dtype=np.int32),
|
||||
@@ -22,6 +33,22 @@ def load_steps(path: str | Path) -> dict[str, np.ndarray]:
|
||||
}
|
||||
|
||||
|
||||
def load_steps(path: str | Path) -> dict[str, np.ndarray]:
|
||||
return _df_to_dict(pd.read_parquet(path))
|
||||
|
||||
|
||||
def load_event_ids(path: str | Path) -> np.ndarray:
|
||||
"""Read only the event_id column — cheap scan for split assignment."""
|
||||
return pd.read_parquet(path, columns=["event_id"])["event_id"].to_numpy()
|
||||
|
||||
|
||||
def iter_file_chunks(path: str | Path) -> Iterator[dict[str, np.ndarray]]:
|
||||
"""Yield one parquet row-group at a time so a large file never fully loads."""
|
||||
pf = pq.ParquetFile(path)
|
||||
for i in range(pf.num_row_groups):
|
||||
yield _df_to_dict(pf.read_row_group(i).to_pandas())
|
||||
|
||||
|
||||
def build_index_maps(
|
||||
data: dict[str, np.ndarray],
|
||||
) -> tuple[dict[int, int], dict[int, int]]:
|
||||
@@ -31,3 +58,19 @@ def build_index_maps(
|
||||
{v: i for i, v in enumerate(pdg_vals)},
|
||||
{v: i for i, v in enumerate(mat_vals)},
|
||||
)
|
||||
|
||||
|
||||
def build_index_maps_from_files(
|
||||
files: list[Path],
|
||||
) -> tuple[dict[int, int], dict[int, int]]:
|
||||
"""Scan only pdg and material_id columns across all files (2-column read)."""
|
||||
pdg_vals: set[int] = set()
|
||||
mat_vals: set[int] = set()
|
||||
for path in files:
|
||||
df = pd.read_parquet(path, columns=["pdg", "material_id"])
|
||||
pdg_vals.update(int(v) for v in df["pdg"].unique())
|
||||
mat_vals.update(int(v) for v in df["material_id"].unique())
|
||||
return (
|
||||
{v: i for i, v in enumerate(sorted(pdg_vals))},
|
||||
{v: i for i, v in enumerate(sorted(mat_vals))},
|
||||
)
|
||||
|
||||
@@ -65,6 +65,39 @@ class Normalizer:
|
||||
return obj
|
||||
|
||||
|
||||
class _WelfordAccumulator:
|
||||
"""Streaming mean/variance (Welford's online algorithm, batch update).
|
||||
|
||||
Use to fit a Normalizer over data that doesn't fit in memory:
|
||||
acc = _WelfordAccumulator(n_features)
|
||||
for chunk in data:
|
||||
acc.update(chunk)
|
||||
normalizer = acc.to_normalizer()
|
||||
"""
|
||||
|
||||
def __init__(self, n_features: int) -> None:
|
||||
self.n = 0
|
||||
self._mean = np.zeros(n_features, dtype=np.float64)
|
||||
self._M2 = np.zeros(n_features, dtype=np.float64)
|
||||
|
||||
def update(self, X: np.ndarray) -> None:
|
||||
X = np.asarray(X, dtype=np.float64)
|
||||
B = X.shape[0]
|
||||
new_n = self.n + B
|
||||
delta = X - self._mean
|
||||
self._mean += delta.sum(0) / new_n
|
||||
delta2 = X - self._mean
|
||||
self._M2 += (delta * delta2).sum(0)
|
||||
self.n = new_n
|
||||
|
||||
def to_normalizer(self) -> "Normalizer":
|
||||
norm = Normalizer()
|
||||
norm.mean = self._mean.astype(np.float32)
|
||||
std = np.sqrt(self._M2 / max(self.n, 1)).astype(np.float32)
|
||||
norm.std = np.where(std < _EPS, 1.0, std).astype(np.float32)
|
||||
return norm
|
||||
|
||||
|
||||
def build_features(
|
||||
data: dict[str, np.ndarray],
|
||||
pdg_map: dict[int, int],
|
||||
|
||||
+10
-6
@@ -32,7 +32,8 @@ def train(
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
model.train()
|
||||
train_loss = 0.0
|
||||
train_loss_sum = 0.0
|
||||
train_n = 0
|
||||
for cond_cont, cond_cat, x1 in train_loader:
|
||||
cond_cont = cond_cont.to(device)
|
||||
cond_cat = cond_cat.to(device)
|
||||
@@ -47,13 +48,15 @@ def train(
|
||||
loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
||||
optimizer.step()
|
||||
train_loss += loss.item() * x1.size(0)
|
||||
train_loss_sum += loss.item() * x1.size(0)
|
||||
train_n += x1.size(0)
|
||||
|
||||
train_loss /= len(train_loader.dataset)
|
||||
train_loss = train_loss_sum / max(train_n, 1)
|
||||
lr_sched.step()
|
||||
|
||||
model.eval()
|
||||
val_loss = 0.0
|
||||
val_loss_sum = 0.0
|
||||
val_n = 0
|
||||
with torch.no_grad():
|
||||
for cond_cont, cond_cat, x1 in val_loader:
|
||||
cond_cont = cond_cont.to(device)
|
||||
@@ -63,8 +66,9 @@ def train(
|
||||
loss = flow_matching_loss(model, x1, cond_cont, cond_cat)
|
||||
else:
|
||||
loss = ddpm_schedule.loss(model, x1, cond_cont, cond_cat)
|
||||
val_loss += loss.item() * x1.size(0)
|
||||
val_loss /= len(val_loader.dataset)
|
||||
val_loss_sum += loss.item() * x1.size(0)
|
||||
val_n += x1.size(0)
|
||||
val_loss = val_loss_sum / max(val_n, 1)
|
||||
|
||||
print(f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f}")
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ dependencies = [
|
||||
"numpy>=1.26",
|
||||
"pandas>=2.2",
|
||||
"pyarrow>=16",
|
||||
"typer>=0.12",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
@@ -16,6 +17,9 @@ dev = [
|
||||
"pytest>=8",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
giant = "giant.cli:app"
|
||||
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
+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)
|
||||
|
||||
@@ -10,6 +10,15 @@ resolution-markers = [
|
||||
"python_full_version < '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "annotated-doc"
|
||||
version = "0.0.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/57/ba/046ceea27344560984e26a590f90bc7f4a75b06701f653222458922b558c/annotated_doc-0.0.4.tar.gz", hash = "sha256:fbcda96e87e9c92ad167c2e53839e57503ecfda18804ea28102353485033faa4", size = 7288, upload-time = "2025-11-10T22:07:42.062Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/1e/d3/26bf1008eb3d2daa8ef4cacc7f3bfdc11818d111f7e2d0201bc6e3b49d45/annotated_doc-0.0.4-py3-none-any.whl", hash = "sha256:571ac1dc6991c450b25a9c2d84a3705e2ae7a53467b5d111c24fa8baabbed320", size = 5303, upload-time = "2025-11-10T22:07:40.673Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "colorama"
|
||||
version = "0.4.6"
|
||||
@@ -112,6 +121,7 @@ dependencies = [
|
||||
{ name = "pandas" },
|
||||
{ name = "pyarrow" },
|
||||
{ name = "torch" },
|
||||
{ name = "typer" },
|
||||
]
|
||||
|
||||
[package.optional-dependencies]
|
||||
@@ -126,6 +136,7 @@ requires-dist = [
|
||||
{ name = "pyarrow", specifier = ">=16" },
|
||||
{ name = "pytest", marker = "extra == 'dev'", specifier = ">=8" },
|
||||
{ name = "torch", specifier = ">=2.3" },
|
||||
{ name = "typer", specifier = ">=0.12" },
|
||||
]
|
||||
provides-extras = ["dev"]
|
||||
|
||||
@@ -150,6 +161,18 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/62/a1/3d680cbfd5f4b8f15abc1d571870c5fc3e594bb582bc3b64ea099db13e56/jinja2-3.1.6-py3-none-any.whl", hash = "sha256:85ece4451f492d0c13c5dd7c13a64681a86afae63a5f347908daf103ce6d2f67", size = 134899, upload-time = "2025-03-05T20:05:00.369Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "markdown-it-py"
|
||||
version = "4.2.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "mdurl" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/06/ff/7841249c247aa650a76b9ee4bbaeae59370dc8bfd2f6c01f3630c35eb134/markdown_it_py-4.2.0.tar.gz", hash = "sha256:04a21681d6fbb623de53f6f364d352309d4094dd4194040a10fd51833e418d49", size = 82454, upload-time = "2026-05-07T12:08:28.36Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/b3/81/4da04ced5a082363ecfa159c010d200ecbd959ae410c10c0264a38cac0f5/markdown_it_py-4.2.0-py3-none-any.whl", hash = "sha256:9f7ebbcd14fe59494226453aed97c1070d83f8d24b6fc3a3bcf9a38092641c4a", size = 91687, upload-time = "2026-05-07T12:08:27.182Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "markupsafe"
|
||||
version = "3.0.3"
|
||||
@@ -213,6 +236,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/70/bc/6f1c2f612465f5fa89b95bead1f44dcb607670fd42891d8fdcd5d039f4f4/markupsafe-3.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:32001d6a8fc98c8cb5c947787c5d08b0a50663d139f1305bac5885d98d9b40fa", size = 14146, upload-time = "2025-09-27T18:37:28.327Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mdurl"
|
||||
version = "0.1.2"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/d6/54/cfe61301667036ec958cb99bd3efefba235e65cdeb9c84d24a8293ba1d90/mdurl-0.1.2.tar.gz", hash = "sha256:bb413d29f5eea38f31dd4754dd7377d4465116fb207585f97bf925588687c1ba", size = 8729, upload-time = "2022-08-14T12:40:10.846Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979, upload-time = "2022-08-14T12:40:09.779Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mpmath"
|
||||
version = "1.3.0"
|
||||
@@ -594,6 +626,19 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/ec/57/56b9bcc3c9c6a792fcbaf139543cee77261f3651ca9da0c93f5c1221264b/python_dateutil-2.9.0.post0-py2.py3-none-any.whl", hash = "sha256:a8b2bc7bffae282281c8140a97d3aa9c14da0b136dfe83f850eea9a5f7470427", size = 229892, upload-time = "2024-03-01T18:36:18.57Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rich"
|
||||
version = "15.0.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "markdown-it-py" },
|
||||
{ name = "pygments" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/c0/8f/0722ca900cc807c13a6a0c696dacf35430f72e0ec571c4275d2371fca3e9/rich-15.0.0.tar.gz", hash = "sha256:edd07a4824c6b40189fb7ac9bc4c52536e9780fbbfbddf6f1e2502c31b068c36", size = 230680, upload-time = "2026-04-12T08:24:00.75Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/82/3b/64d4899d73f91ba49a8c18a8ff3f0ea8f1c1d75481760df8c68ef5235bf5/rich-15.0.0-py3-none-any.whl", hash = "sha256:33bd4ef74232fb73fe9279a257718407f169c09b78a87ad3d296f548e27de0bb", size = 310654, upload-time = "2026-04-12T08:24:02.83Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "setuptools"
|
||||
version = "81.0.0"
|
||||
@@ -603,6 +648,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e1/e3/c164c88b2e5ce7b24d667b9bd83589cf4f3520d97cad01534cd3c4f55fdb/setuptools-81.0.0-py3-none-any.whl", hash = "sha256:fdd925d5c5d9f62e4b74b30d6dd7828ce236fd6ed998a08d81de62ce5a6310d6", size = 1062021, upload-time = "2026-02-06T21:10:37.175Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "shellingham"
|
||||
version = "1.5.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/58/15/8b3609fd3830ef7b27b655beb4b4e9c62313a4e8da8c676e142cc210d58e/shellingham-1.5.4.tar.gz", hash = "sha256:8dbca0739d487e5bd35ab3ca4b36e11c4078f3a234bfce294b0a0291363404de", size = 10310, upload-time = "2023-10-24T04:13:40.426Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e0/f9/0595336914c5619e5f28a1fb793285925a8cd4b432c9da0a987836c7f822/shellingham-1.5.4-py2.py3-none-any.whl", hash = "sha256:7ecfff8f2fd72616f7481040475a65b2bf8af90a56c89140852d1120324e8686", size = 9755, upload-time = "2023-10-24T04:13:38.866Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "six"
|
||||
version = "1.17.0"
|
||||
@@ -685,6 +739,21 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/c1/68/fa86e5a39608000f645535b2c124920126327ab731f8c4fafd5b07ff8d4b/triton-3.7.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:ce061073102714b725f3660ec6939d94a1da7984b3aa99c921417cae273672f5", size = 201546766, upload-time = "2026-05-07T18:46:42.088Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "typer"
|
||||
version = "0.26.7"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "annotated-doc" },
|
||||
{ name = "colorama", marker = "sys_platform == 'win32'" },
|
||||
{ name = "rich" },
|
||||
{ name = "shellingham" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/5e/ed/ef06584ccdd5c410df0837951ecd7e15d9a6144ea1bd4c73cecab1a89891/typer-0.26.7.tar.gz", hash = "sha256:e314a34c617e419c091b2830dda3ea1f257134ff593061a8f5b9717ab8dddb3a", size = 201709, upload-time = "2026-06-03T07:18:06.843Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/24/25/2201973529af2c954de0bb725323c3aaed6d7f0ceee8f550dec9185df013/typer-0.26.7-py3-none-any.whl", hash = "sha256:5c87cfbc5d34491c5346ebf49c23e18d56ccb863268d3a8d592b26087c2f5e58", size = 122456, upload-time = "2026-06-03T07:18:05.732Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "typing-extensions"
|
||||
version = "4.15.0"
|
||||
|
||||
Reference in New Issue
Block a user