diff --git a/giant/cli.py b/giant/cli.py index 9523e8f..ef0a3cc 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -1,5 +1,3 @@ -import subprocess -import tomllib from enum import Enum from pathlib import Path from typing import Optional @@ -7,18 +5,17 @@ from typing import Optional import numpy as np import torch import typer -from torch.utils.data import DataLoader from typing_extensions import Annotated import pyarrow as pa import pyarrow.parquet as pq +from giant import config as gconfig +from giant.constants import LOCAL_TARGET_NAMES from giant.data.loader import ( find_parquet_files, - load_event_ids, iter_file_chunks, iter_cond_chunks, - build_index_maps_from_files, ) from giant.data.transforms import ( build_features, @@ -26,28 +23,14 @@ from giant.data.transforms import ( inv_local_frame_rotation, inv_log_transform, reconstruct_post_pos, - _WelfordAccumulator, Normalizer, ) -from giant.data.dataset import make_event_split, StreamingStepsDataset from giant.model.network import DenoisingMLP +from giant.pipeline import run_train_job from giant.sample import sample_flow -from giant.train import train as run_training app = typer.Typer(no_args_is_help=True) -_LOCAL_TARGET_NAMES = [ - "log_step_length", - "log_delta_e", - "log_edep", - "post_dx", - "post_dy", - "post_dz", - "travel_dx", - "travel_dy", - "travel_dz", -] - @app.callback() def _main() -> None: @@ -64,38 +47,6 @@ class Coord(str, Enum): local = "local" -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")], @@ -108,42 +59,26 @@ def train( n_blocks: Annotated[Optional[int], typer.Option()] = None, emb_dim: Annotated[Optional[int], typer.Option()] = None, val_fraction: Annotated[Optional[float], typer.Option()] = None, + seed: Annotated[Optional[int], typer.Option(help="Random seed for reproducibility")] = 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, + resume: Annotated[Optional[Path], typer.Option(help="Checkpoint .pt to resume training from")] = 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, + "val_fraction": val_fraction, "num_workers": num_workers, "seed": seed, }.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) - + cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, cli_train, cli_model) t, m = cfg["train"], cfg["model"] - _device = torch.device(device) if device else _auto_device() + _device = torch.device(device) if device else gconfig.auto_device() out_dir = out or Path( f"checkpoints/{t['mode']}" f"_h{m['hidden_dim']}" @@ -156,89 +91,10 @@ def train( 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(9) - 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=t["num_workers"], pin_memory=pin, - ) - val_loader = DataLoader( - val_ds, batch_size=None, - 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) - - 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"], - } - - 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, + run_train_job( + data=data, cfg=cfg, out_dir=out_dir, device=_device, + shuffle_buffer=shuffle_buffer, num_workers=t["num_workers"], + resume=resume, echo=typer.echo, ) @@ -258,7 +114,7 @@ def predict( device: Annotated[Optional[str], typer.Option(help="cpu | cuda | mps (default: auto)")] = None, ) -> None: """Run trained model on a parquet file and save predictions.""" - _device = torch.device(device) if device else _auto_device() + _device = torch.device(device) if device else gconfig.auto_device() typer.echo(f"device: {_device}") # --- Load checkpoint --- @@ -329,8 +185,8 @@ def predict( "material": chunk["material"], "layer_id": chunk["layer_id"], "n_sec": chunk["n_sec"], - **{f"pred_{name}": raw[:, j] for j, name in enumerate(_LOCAL_TARGET_NAMES)}, - **{f"true_{name}": target_raw[:, j] for j, name in enumerate(_LOCAL_TARGET_NAMES)}, + **{f"pred_{name}": raw[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)}, + **{f"true_{name}": target_raw[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)}, }) else: step_length = inv_log_transform(raw[:, 0]) diff --git a/giant/config.py b/giant/config.py new file mode 100644 index 0000000..87cd3e9 --- /dev/null +++ b/giant/config.py @@ -0,0 +1,106 @@ +import random +import subprocess +import sys +import tomllib +from datetime import datetime, timezone +from pathlib import Path + +import numpy as np +import torch + +DEFAULT_CONFIG: dict = { + "train": { + "mode": "flow", "epochs": 100, "batch_size": 4096, "lr": 3e-4, + "val_fraction": 0.1, "num_workers": 4, "seed": 0, + }, + "model": { + "hidden_dim": 256, "n_blocks": 6, "emb_dim": 16, + }, +} + + +def git_hash() -> str: + try: + return subprocess.check_output( + ["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL + ).decode().strip() + except Exception: + return "unknown" + + +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 load_toml(path: Path) -> dict: + with open(path, "rb") as f: + return tomllib.load(f) + + +def merge_cli_overrides( + defaults: dict, + config_path: Path | None, + train_overrides: dict, + model_overrides: dict, +) -> dict: + """Resolve config as defaults -> TOML file -> explicit CLI flags.""" + cfg = {"train": dict(defaults["train"]), "model": dict(defaults["model"])} + if config_path is not None: + file_cfg = load_toml(config_path) + for section in ("train", "model"): + cfg[section].update(file_cfg.get(section, {})) + cfg["train"].update(train_overrides) + cfg["model"].update(model_overrides) + return cfg + + +def seed_everything(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def save_config(cfg: dict, out_dir: Path, meta: dict) -> None: + lines = [] + for section, values in cfg.items(): + lines.append(f"[{section}]") + for k, v in values.items(): + lines.append(f"{k:<14} = {repr(v) if isinstance(v, str) else v}") + lines.append("") + + lines.append("[meta]") + for k, v in meta.items(): + lines.append(f"{k:<14} = {repr(v) if isinstance(v, str) else v}") + + (out_dir / "config.toml").write_text("\n".join(lines)) + + +def build_run_meta( + data: Path, + seed: int, + n_pdg_codes: int, + n_materials: int, + n_train_events: int, + n_val_events: int, + n_train_steps: int, +) -> dict: + return { + "git_hash": git_hash(), + "seed": seed, + "timestamp_utc": datetime.now(timezone.utc).isoformat(timespec="seconds"), + "python_version": sys.version.split()[0], + "torch_version": torch.__version__, + "command": " ".join(sys.argv), + "data_path": str(data), + "n_pdg_codes": n_pdg_codes, + "n_materials": n_materials, + "n_train_events": n_train_events, + "n_val_events": n_val_events, + "n_train_steps": n_train_steps, + } diff --git a/giant/constants.py b/giant/constants.py new file mode 100644 index 0000000..7b00422 --- /dev/null +++ b/giant/constants.py @@ -0,0 +1,13 @@ +X_DIM = 9 + +LOCAL_TARGET_NAMES = [ + "log_step_length", + "log_delta_e", + "log_edep", + "post_dx", + "post_dy", + "post_dz", + "travel_dx", + "travel_dy", + "travel_dz", +] diff --git a/giant/model/network.py b/giant/model/network.py index 5ae91a2..9ff099c 100644 --- a/giant/model/network.py +++ b/giant/model/network.py @@ -3,6 +3,8 @@ import math import torch import torch.nn as nn +from giant.constants import X_DIM + class SinusoidalEmbedding(nn.Module): def __init__(self, dim: int) -> None: @@ -73,7 +75,7 @@ class DenoisingMLP(nn.Module): emb_dim: int = 16, time_dim: int = 64, cond_out_dim: int = 128, - x_dim: int = 9, + x_dim: int = X_DIM, ) -> None: super().__init__() self.time_emb = SinusoidalEmbedding(time_dim) diff --git a/giant/pipeline.py b/giant/pipeline.py new file mode 100644 index 0000000..db8d5b9 --- /dev/null +++ b/giant/pipeline.py @@ -0,0 +1,129 @@ +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"], + ) + 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"], + } + + 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, + ) diff --git a/giant/sample.py b/giant/sample.py index 5ae4be5..e9b45c4 100644 --- a/giant/sample.py +++ b/giant/sample.py @@ -1,5 +1,7 @@ import torch +from giant.constants import X_DIM + @torch.no_grad() def sample_flow( @@ -12,7 +14,7 @@ def sample_flow( model.eval() B = cond_cont.size(0) device = cond_cont.device - x = torch.randn(B, 9, device=device) + x = torch.randn(B, X_DIM, device=device) dt = 1.0 / steps for i in range(steps): t = torch.full((B,), i * dt, device=device) @@ -32,7 +34,7 @@ def sample_ddpm( model.eval() B = cond_cont.size(0) device = cond_cont.device - x = torch.randn(B, 9, device=device) + x = torch.randn(B, X_DIM, device=device) T = schedule.T for i in reversed(range(T)): t_norm = torch.full((B,), i / T, device=device) @@ -63,7 +65,7 @@ def sample_ddim( device = cond_cont.device T = schedule.T timesteps = torch.linspace(T - 1, 0, steps, dtype=torch.long, device=device) - x = torch.randn(B, 9, device=device) + x = torch.randn(B, X_DIM, device=device) for step_idx, ts in enumerate(timesteps): t_idx = int(ts.item()) t_norm = torch.full((B,), t_idx / T, device=device) diff --git a/giant/train.py b/giant/train.py index e803b86..12d0b5b 100644 --- a/giant/train.py +++ b/giant/train.py @@ -1,3 +1,5 @@ +import csv +import time from pathlib import Path import torch @@ -6,6 +8,8 @@ from torch.utils.data import DataLoader from giant.model.schedule import CosineSchedule, flow_matching_loss +_METRICS_FIELDS = ["epoch", "train_loss", "val_loss", "lr", "epoch_time_s"] + def train( model: torch.nn.Module, @@ -20,6 +24,7 @@ def train( pdg_map: dict | None = None, mat_map: dict | None = None, model_config: dict | None = None, + resume_path: str | Path | None = None, ) -> None: out_dir = Path(out_dir) out_dir.mkdir(parents=True, exist_ok=True) @@ -29,9 +34,27 @@ def train( lr_sched = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) ddpm_schedule = CosineSchedule().to(device) if mode == "ddpm" else None - best_val_loss = float("inf") - for epoch in range(1, epochs + 1): + start_epoch = 1 + best_val_loss = float("inf") + if resume_path is not None: + ckpt = torch.load(resume_path, map_location=device, weights_only=False) + model.load_state_dict(ckpt["model"]) + optimizer.load_state_dict(ckpt["optimizer"]) + lr_sched.load_state_dict(ckpt["lr_sched"]) + start_epoch = ckpt.get("epoch", 0) + 1 + best_val_loss = ckpt.get("best_val_loss", float("inf")) + + metrics_path = out_dir / "metrics.csv" + write_header = not (resume_path is not None and metrics_path.exists()) + metrics_file = open(metrics_path, "a", newline="") + metrics_writer = csv.DictWriter(metrics_file, fieldnames=_METRICS_FIELDS) + if write_header: + metrics_writer.writeheader() + + for epoch in range(start_epoch, epochs + 1): + epoch_start = time.monotonic() + current_lr = optimizer.param_groups[0]["lr"] model.train() train_loss_sum = 0.0 train_n = 0 @@ -70,20 +93,39 @@ def train( val_loss_sum += loss.item() * x1.size(0) val_n += x1.size(0) val_loss = val_loss_sum / max(val_n, 1) + epoch_time = time.monotonic() - epoch_start - print(f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f}") + print( + f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f} " + f"lr {current_lr:.2e} {epoch_time:.1f}s" + ) + metrics_writer.writerow({ + "epoch": epoch, "train_loss": train_loss, "val_loss": val_loss, + "lr": current_lr, "epoch_time_s": epoch_time, + }) + metrics_file.flush() + + ckpt: dict = { + "model": model.state_dict(), + "optimizer": optimizer.state_dict(), + "lr_sched": lr_sched.state_dict(), + "epoch": epoch, + "best_val_loss": best_val_loss, + } + if normalizer_dict is not None: + ckpt["normalizer"] = normalizer_dict + if pdg_map is not None: + ckpt["pdg_map"] = pdg_map + if mat_map is not None: + ckpt["mat_map"] = mat_map + if model_config is not None: + ckpt["model_config"] = model_config if val_loss < best_val_loss: best_val_loss = val_loss - ckpt: dict = {"model": model.state_dict()} - if normalizer_dict is not None: - ckpt["normalizer"] = normalizer_dict - if pdg_map is not None: - ckpt["pdg_map"] = pdg_map - if mat_map is not None: - ckpt["mat_map"] = mat_map - if model_config is not None: - ckpt["model_config"] = model_config + ckpt["best_val_loss"] = best_val_loss torch.save(ckpt, out_dir / "best.pt") - torch.save({"model": model.state_dict()}, out_dir / "last.pt") + torch.save(ckpt, out_dir / "last.pt") + + metrics_file.close() diff --git a/giant/validate.py b/giant/validate.py index 83630ca..c9b64a3 100644 --- a/giant/validate.py +++ b/giant/validate.py @@ -2,20 +2,9 @@ import numpy as np import torch from torch.utils.data import DataLoader +from giant.constants import LOCAL_TARGET_NAMES from giant.sample import sample_flow, sample_ddpm, sample_ddim -_TARGET_NAMES = [ - "log_step_length", - "log_delta_e", - "log_edep", - "post_dx", - "post_dy", - "post_dz", - "travel_dx", - "travel_dy", - "travel_dz", -] - def validate_marginals( model: torch.nn.Module, @@ -56,7 +45,7 @@ def validate_marginals( header = f"{'Dim':<20} {'real_mean':>10} {'gen_mean':>10} {'real_std':>10} {'gen_std':>10}" print(f"\n{header}") print("-" * len(header)) - for j, name in enumerate(_TARGET_NAMES): + for j, name in enumerate(LOCAL_TARGET_NAMES): r, g = real[:, j], generated[:, j] print(f"{name:<20} {r.mean():>10.4f} {g.mean():>10.4f} {r.std():>10.4f} {g.std():>10.4f}") diff --git a/scripts/train.py b/scripts/train.py index 1c321da..87fc88a 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -1,54 +1,10 @@ import argparse -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 ( - 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 _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_config(path: str) -> 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)) +from giant import config as gconfig +from giant.pipeline import run_train_job def main() -> None: @@ -63,42 +19,28 @@ 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("--seed", type=int, help="Random seed for reproducibility") 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) + parser.add_argument("--resume", default=None, help="Checkpoint .pt to resume training from") args = parser.parse_args() - # 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 args.config: - file_cfg = _load_config(args.config) - for section in ("train", "model"): - cfg[section].update(file_cfg.get(section, {})) - cli_train = {k: v for k, v in { "mode": args.mode, "epochs": args.epochs, "batch_size": args.batch_size, "lr": args.lr, "val_fraction": args.val_fraction, "num_workers": args.num_workers, + "seed": args.seed, }.items() if v is not None} cli_model = {k: v for k, v in { "hidden_dim": args.hidden_dim, "n_blocks": args.n_blocks, "emb_dim": args.emb_dim, }.items() if v is not None} - cfg["train"].update(cli_train) - cfg["model"].update(cli_model) - + config_path = Path(args.config) if args.config else None + cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config_path, cli_train, cli_model) t, m = cfg["train"], cfg["model"] - device = torch.device(args.device) if args.device else _auto_device() + device = torch.device(args.device) if args.device else gconfig.auto_device() out_dir = Path(args.out or ( f"checkpoints/{t['mode']}" f"_h{m['hidden_dim']}" @@ -111,100 +53,10 @@ def main() -> None: print(f"device: {device}") print(f"out_dir: {out_dir}") - # --- Discover files --- - files = find_parquet_files(args.data) - print(f"found {len(files)} parquet file(s)") - - # --- 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") - - # --- 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(9) - 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, - batch_size=t["batch_size"], - 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, - 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=t["num_workers"], pin_memory=pin, - ) - val_loader = DataLoader( - val_ds, batch_size=None, - 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"], - ) - 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) - - 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()}, + run_train_job( + data=Path(args.data), cfg=cfg, out_dir=out_dir, device=device, + shuffle_buffer=args.shuffle_buffer, num_workers=t["num_workers"], + resume=Path(args.resume) if args.resume else None, echo=print, )