da84801625
Wire a dropout hyperparameter (default 0.1) through the config, model, training pipeline, and CLI. Persisted in saved model_config so checkpoints reconstruct the architecture correctly. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
133 lines
4.4 KiB
Python
133 lines
4.4 KiB
Python
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import torch
|
|
from torch.utils.data import DataLoader
|
|
|
|
from giant import config
|
|
from giant.constants import X_DIM
|
|
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
|
|
|
|
|
|
def run_train_job(
|
|
data: Path,
|
|
cfg: dict,
|
|
out_dir: Path,
|
|
device: torch.device,
|
|
shuffle_buffer: int,
|
|
num_workers: int,
|
|
resume: Path | None = None,
|
|
echo=print,
|
|
) -> None:
|
|
t, m = cfg["train"], cfg["model"]
|
|
config.seed_everything(t["seed"])
|
|
|
|
out_dir = Path(out_dir)
|
|
|
|
files = find_parquet_files(data)
|
|
echo(f"found {len(files)} parquet file(s)")
|
|
|
|
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())
|
|
echo(
|
|
f" {len(all_event_ids):,} steps | "
|
|
f"{len(train_events)} train events (~{n_train_steps:,} steps) | "
|
|
f"{len(val_events)} val events"
|
|
)
|
|
|
|
echo("building vocabulary maps …")
|
|
pdg_map, mat_map = build_index_maps_from_files(files)
|
|
echo(f" {len(pdg_map)} PDG codes | {len(mat_map)} materials")
|
|
|
|
echo("fitting normalizer (streaming) …")
|
|
cond_acc = _WelfordAccumulator(X_DIM)
|
|
tgt_acc = _WelfordAccumulator(X_DIM)
|
|
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,
|
|
batch_size=t["batch_size"],
|
|
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,
|
|
batch_size=t["batch_size"],
|
|
shuffle=False,
|
|
)
|
|
|
|
# Dataset yields whole batches already, so batch_size=None tells DataLoader
|
|
# to pass them through instead of re-collating row-by-row in Python.
|
|
pin = device.type == "cuda"
|
|
train_loader = DataLoader(
|
|
train_ds, batch_size=None,
|
|
num_workers=num_workers, pin_memory=pin,
|
|
)
|
|
val_loader = DataLoader(
|
|
val_ds, batch_size=None,
|
|
num_workers=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"],
|
|
dropout=m["dropout"],
|
|
)
|
|
echo(f"model: {sum(p.numel() for p in model.parameters()):,} parameters")
|
|
|
|
out_dir.mkdir(parents=True, exist_ok=True)
|
|
meta = config.build_run_meta(
|
|
data=data,
|
|
seed=t["seed"],
|
|
n_pdg_codes=len(pdg_map),
|
|
n_materials=len(mat_map),
|
|
n_train_events=len(train_events),
|
|
n_val_events=len(val_events),
|
|
n_train_steps=n_train_steps,
|
|
)
|
|
config.save_config(cfg, out_dir, meta)
|
|
|
|
model_config = {
|
|
"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"],
|
|
"dropout": m["dropout"],
|
|
}
|
|
|
|
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()},
|
|
model_config=model_config,
|
|
resume_path=resume,
|
|
validate_every=t["validate_every"],
|
|
)
|