The merged proc_idx/proc_map plumbing wasn't run through ruff format before merging; reflow only, no logic changes.
194 lines
6.2 KiB
Python
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,
|
|
)
|