Files
giant/giant/pipeline.py
T
lars 05d5dee606 Apply ruff format after merging phase2-secondary-prediction
The merged proc_idx/proc_map plumbing wasn't run through ruff format
before merging; reflow only, no logic changes.
2026-07-15 10:00:36 +02:00

194 lines
6.2 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 COND_DIM, EMB_DIM, K_MAX, SEC_SLOT_DIM, X_DIM
from giant.data.loader import (
find_parquet_files,
load_event_ids,
iter_file_chunks,
build_index_maps_from_files,
build_process_map_from_files,
)
from giant.data.transforms import build_features, _WelfordAccumulator
from giant.data.dataset import make_event_split, StreamingStepsDataset
from giant.model.network import build_models
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())
total_train_batches = n_train_steps // t["batch_size"]
echo(
f" {len(all_event_ids):,} steps | "
f"{len(train_events)} train events (~{n_train_steps:,} steps, ~{total_train_batches:,} batches) | "
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")
router_cfg = m["router"]
proc_map: dict[str, int] | None = None
if router_cfg.get("enabled") and router_cfg.get("type") == "process":
echo("building process vocabulary …")
proc_map = build_process_map_from_files(
files, n_experts=router_cfg["n_experts"]
)
echo(
f" {len(proc_map)} process labels mapped to {router_cfg['n_experts']} experts"
)
echo("fitting normalizer (streaming) …")
cond_acc = _WelfordAccumulator(COND_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_s1, _n_sec, _sec_cont, _sec_pdg, _proc, _, _ = (
build_features(
chunk_tr,
pdg_map,
mat_map,
proc_map=proc_map,
require_secondaries=True,
)
)
cond_acc.update(cond_cont)
tgt_acc.update(target_s1)
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,
proc_map=proc_map,
)
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,
proc_map=proc_map,
)
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,
)
emb_dim = m.get("emb_dim", EMB_DIM)
# SEC_SLOT_DIM must match constants (1 stick + 3 dir + emb_dim)
assert SEC_SLOT_DIM == 1 + 3 + emb_dim, (
f"SEC_SLOT_DIM={SEC_SLOT_DIM} must equal 1+3+emb_dim={1 + 3 + emb_dim}; "
"update giant/constants.py if emb_dim changed"
)
model_config = {
"pdg_vocab": len(pdg_map),
"mat_vocab": len(mat_map),
"hidden_dim": m["hidden_dim"],
"n_blocks": m["n_blocks"],
"emb_dim": emb_dim,
"dropout": m["dropout"],
"k_max": K_MAX,
"sec_slot_dim": SEC_SLOT_DIM,
"router": dict(router_cfg),
"expert_hidden_dim": router_cfg["expert_hidden_dim"],
"expert_n_blocks": router_cfg["expert_n_blocks"],
}
stage1_model, sec_decoder = build_models(model_config)
echo(
f"stage1: {sum(p.numel() for p in stage1_model.parameters()):,} parameters | "
f"sec_decoder: {sum(p.numel() for p in sec_decoder.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)
run_training(
stage1_model=stage1_model,
sec_decoder=sec_decoder,
train_loader=train_loader,
val_loader=val_loader,
mode=t["mode"],
epochs=t["epochs"],
lr=t["lr"],
warmup_epochs=t["warmup_epochs"],
device=device,
out_dir=out_dir,
lambda_nsec=t.get("lambda_nsec", 0.1),
lambda_s2=t.get("lambda_s2", 1.0),
lambda_balance=router_cfg.get("lambda_balance", 0.0),
lambda_proc=router_cfg.get("lambda_proc", 0.0),
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()},
proc_map=proc_map,
model_config=model_config,
resume_path=resume,
validate_every=t["validate_every"],
validate_steps=t["validate_steps"],
total_train_batches=total_train_batches,
)