diff --git a/CLAUDE.md b/CLAUDE.md index 97228bc..f7a1a8c 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -15,6 +15,16 @@ python scripts/train.py --data path/to/steps.parquet --mode ddpm # train (DDPM `cpu` and `cuda` are mutually exclusive — pick one to select the torch build (pinned to 2.3.x; newer torch requires newer NVIDIA drivers). Plain `uv sync` with no extra will not install torch at all; uv has no concept of a "default extra", so `--extra cpu` should always be included unless you need GPU support. +### Lint and type checking + +```bash +uv run ruff check . # lint +uv run ruff format . # format +uv run ty check . # type check +``` + +Part of the `dev` extra. Run these periodically (not just at commit time) to catch drift early. + ## Architecture GIANT is a conditional generative surrogate for the Geant4 step function. It replaces the stochastic physics engine: given a pre-step particle state (conditioning), it samples a post-step outcome. diff --git a/giant/analysis.py b/giant/analysis.py index 90a3a42..c2d2c10 100644 --- a/giant/analysis.py +++ b/giant/analysis.py @@ -74,7 +74,9 @@ from giant.validate import _histogram_kl _N_LOG_DIMS = 3 # log_step_length, log_delta_e, log_edep are the first 3 target dims RAW_TARGET_NAMES = [n.removeprefix("log_") for n in LOCAL_TARGET_NAMES] -_LOG_EPS = 1e-8 # mirrors giant.data.transforms._EPS, duplicated for use in polars exprs +_LOG_EPS = ( + 1e-8 # mirrors giant.data.transforms._EPS, duplicated for use in polars exprs +) def _hist_edges(*arrays: np.ndarray, bins: int) -> np.ndarray: @@ -151,9 +153,15 @@ def load_model_bundle( schedule = CosineSchedule().to(device) if mode != "flow" else None return ModelBundle( - model=model, cond_normalizer=cond_normalizer, target_normalizer=target_normalizer, - pdg_map=pdg_map, mat_map=mat_map, model_config=ckpt["model_config"], - mode=mode, schedule=schedule, device=device, + model=model, + cond_normalizer=cond_normalizer, + target_normalizer=target_normalizer, + pdg_map=pdg_map, + mat_map=mat_map, + model_config=ckpt["model_config"], + mode=mode, + schedule=schedule, + device=device, ) @@ -174,10 +182,15 @@ def make_val_loader( full = {k: np.concatenate([c[k] for c in chunks]) for k in chunks[0]} cond_cont, cond_cat, target, _, _ = build_features( - full, bundle.pdg_map, bundle.mat_map, - cond_normalizer=bundle.cond_normalizer, target_normalizer=bundle.target_normalizer, + full, + bundle.pdg_map, + bundle.mat_map, + cond_normalizer=bundle.cond_normalizer, + target_normalizer=bundle.target_normalizer, + ) + _, val_ds = train_val_split( + full, cond_cont, cond_cat, target, val_fraction=val_fraction, seed=seed ) - _, val_ds = train_val_split(full, cond_cont, cond_cat, target, val_fraction=val_fraction, seed=seed) return DataLoader(val_ds, batch_size=batch_size, shuffle=False) @@ -230,10 +243,13 @@ def collect_samples( material = np.array([idx_to_mat[i] for i in cond_cat_arr[:, 1].astype(np.int64)]) return SampleCollection( - cond_cont_raw=cond_cont_raw, pdg=pdg, material=material, + cond_cont_raw=cond_cont_raw, + pdg=pdg, + material=material, real_raw=_to_raw_targets(real_norm, bundle.target_normalizer), gen_raw=_to_raw_targets(gen_norm, bundle.target_normalizer), - real_norm=real_norm, gen_norm=gen_norm, + real_norm=real_norm, + gen_norm=gen_norm, ) @@ -268,7 +284,15 @@ def _check_predict_metadata(path: Path) -> None: _COND_CONT_COLS = [ - "pre_x", "pre_y", "pre_z", "pre_E", "pre_dx", "pre_dy", "pre_dz", "layer_id", "n_sec", + "pre_x", + "pre_y", + "pre_z", + "pre_E", + "pre_dx", + "pre_dy", + "pre_dz", + "layer_id", + "n_sec", ] @@ -318,6 +342,7 @@ def load_predicted_local(path: str | Path) -> SampleCollection: # Tier 1: stratified marginals # --------------------------------------------------------------------------- + def _group_labels( collection: SampleCollection, group_by: str | None, @@ -329,14 +354,17 @@ def _group_labels( if group_by == "pdg": return [(f"pdg={v}", collection.pdg == v) for v in np.unique(collection.pdg)] if group_by == "material": - return [(f"material={v}", collection.material == v) for v in np.unique(collection.material)] + return [ + (f"material={v}", collection.material == v) + for v in np.unique(collection.material) + ] if group_by == "energy": pre_E = collection.cond_cont_raw[:, 3] edges = np.quantile(pre_E, np.linspace(0, 1, n_energy_bins + 1)) edges[-1] += 1e-6 bin_idx = np.digitize(pre_E, edges[1:-1]) return [ - (f"E∈[{edges[i]:.3g},{edges[i+1]:.3g})", bin_idx == i) + (f"E∈[{edges[i]:.3g},{edges[i + 1]:.3g})", bin_idx == i) for i in range(n_energy_bins) ] raise ValueError(f"unknown group_by={group_by!r}") @@ -360,13 +388,23 @@ def marginal_table( continue real, gen = collection.real_raw[mask], collection.gen_raw[mask] for j, name in enumerate(RAW_TARGET_NAMES): - rows.append({ - "group": label, "dim": name, "n": int(mask.sum()), - "real_mean": real[:, j].mean(), "gen_mean": gen[:, j].mean(), - "real_std": real[:, j].std(), "gen_std": gen[:, j].std(), - "kl_real_gen": _histogram_kl(real[:, j], gen[:, j], bins=bins), - }) - return pd.DataFrame(rows).sort_values("kl_real_gen", ascending=False).reset_index(drop=True) + rows.append( + { + "group": label, + "dim": name, + "n": int(mask.sum()), + "real_mean": real[:, j].mean(), + "gen_mean": gen[:, j].mean(), + "real_std": real[:, j].std(), + "gen_std": gen[:, j].std(), + "kl_real_gen": _histogram_kl(real[:, j], gen[:, j], bins=bins), + } + ) + return ( + pd.DataFrame(rows) + .sort_values("kl_real_gen", ascending=False) + .reset_index(drop=True) + ) # --------------------------------------------------------------------------- @@ -381,7 +419,10 @@ def marginal_table( # would be too large to hold in memory at once. # --------------------------------------------------------------------------- -def _histogram_kl_pl(p: pl.Series, q: pl.Series, bins: int = 50, eps: float = 1e-8) -> float: + +def _histogram_kl_pl( + p: pl.Series, q: pl.Series, bins: int = 50, eps: float = 1e-8 +) -> float: """Polars duplicate of `giant.validate._histogram_kl`, binning via `Series.hist`.""" lo, hi = min(p.min(), q.min()), max(p.max(), q.max()) if hi <= lo: @@ -409,7 +450,9 @@ def _scan_predicted_local(source: str | Path | pl.LazyFrame) -> pl.LazyFrame: def _group_filters_pl( - lf: pl.LazyFrame, group_by: str | None, n_energy_bins: int, + lf: pl.LazyFrame, + group_by: str | None, + n_energy_bins: int, ) -> list[tuple[str, pl.Expr]]: if group_by is None: return [("all", pl.lit(True))] @@ -425,7 +468,7 @@ def _group_filters_pl( edges[-1] += 1e-6 return [ ( - f"E∈[{edges[i]:.3g},{edges[i+1]:.3g})", + f"E∈[{edges[i]:.3g},{edges[i + 1]:.3g})", (pl.col("pre_E") >= edges[i]) & (pl.col("pre_E") < edges[i + 1]), ) for i in range(n_energy_bins) @@ -457,18 +500,26 @@ def marginal_table_pl( if n < 2: continue for j, name in enumerate(RAW_TARGET_NAMES): - pair = glf.select([ - _raw_expr(true_cols[j], j).alias("real"), - _raw_expr(pred_cols[j], j).alias("gen"), - ]).collect() + pair = glf.select( + [ + _raw_expr(true_cols[j], j).alias("real"), + _raw_expr(pred_cols[j], j).alias("gen"), + ] + ).collect() real_s, gen_s = pair["real"], pair["gen"] - rows.append({ - "group": label, "dim": name, "n": n, - "real_mean": real_s.mean(), "gen_mean": gen_s.mean(), - # ddof=0 to match numpy's (population-std) default used by marginal_table - "real_std": real_s.std(ddof=0), "gen_std": gen_s.std(ddof=0), - "kl_real_gen": _histogram_kl_pl(real_s, gen_s, bins=bins), - }) + rows.append( + { + "group": label, + "dim": name, + "n": n, + "real_mean": real_s.mean(), + "gen_mean": gen_s.mean(), + # ddof=0 to match numpy's (population-std) default used by marginal_table + "real_std": real_s.std(ddof=0), + "gen_std": gen_s.std(ddof=0), + "kl_real_gen": _histogram_kl_pl(real_s, gen_s, bins=bins), + } + ) return pl.DataFrame(rows).sort("kl_real_gen", descending=True) @@ -492,14 +543,20 @@ def plot_marginals( groups = _group_labels(collection, group_by, n_energy_bins) if group_by is not None: - table = marginal_table(collection, group_by=group_by, n_energy_bins=n_energy_bins, bins=bins) - worst_first = table.groupby("group")["kl_real_gen"].max().sort_values(ascending=False) + table = marginal_table( + collection, group_by=group_by, n_energy_bins=n_energy_bins, bins=bins + ) + worst_first = ( + table.groupby("group")["kl_real_gen"].max().sort_values(ascending=False) + ) order = {label: rank for rank, label in enumerate(worst_first.index)} groups = sorted(groups, key=lambda g: order[g[0]])[:max_groups] n_rows, n_cols = len(groups), len(dims) fig, axes = plt.subplots( - n_rows, n_cols, squeeze=False, + n_rows, + n_cols, + squeeze=False, figsize=(figsize_per_axis[0] * n_cols, figsize_per_axis[1] * n_rows), ) for row, (label, mask) in enumerate(groups): @@ -523,6 +580,7 @@ def plot_marginals( # Tier 2: joint structure # --------------------------------------------------------------------------- + def correlation_matrices(collection: SampleCollection) -> tuple[np.ndarray, np.ndarray]: """Pearson correlation matrices of the raw targets, real vs generated.""" return ( @@ -581,9 +639,9 @@ def plot_pairwise( fig, axes = plt.subplots(2, len(pairs), squeeze=False, figsize=(4 * len(pairs), 7)) for col, (a, b) in enumerate(pairs): ia, ib = RAW_TARGET_NAMES.index(a), RAW_TARGET_NAMES.index(b) - for row, (data, title) in enumerate([ - (collection.real_raw, "real"), (collection.gen_raw, "generated") - ]): + for row, (data, title) in enumerate( + [(collection.real_raw, "real"), (collection.gen_raw, "generated")] + ): ax = axes[row][col] ax.scatter(data[idx, ia], data[idx, ib], s=3, alpha=0.3) ax.set_xlabel(a) @@ -600,11 +658,13 @@ def direction_alignment(collection: SampleCollection) -> tuple[np.ndarray, np.nd These two unit vectors are coupled through the scattering physics, so their joint alignment is a check the per-dimension marginals can't see. """ + def cos_angle(raw: np.ndarray) -> np.ndarray: post, travel = raw[:, 3:6], raw[:, 6:9] return np.sum(post * travel, axis=1) / ( np.linalg.norm(post, axis=1) * np.linalg.norm(travel, axis=1) + 1e-8 ) + return cos_angle(collection.real_raw), cos_angle(collection.gen_raw) @@ -624,7 +684,10 @@ def plot_direction_alignment(collection: SampleCollection, bins: int = 50): # Tier 3: physical constraints # --------------------------------------------------------------------------- -def constraint_report(collection: SampleCollection, norm_tol: float = 0.05) -> pd.DataFrame: + +def constraint_report( + collection: SampleCollection, norm_tol: float = 0.05 +) -> pd.DataFrame: """Rate of physical-constraint violations in the generated raw-space samples. The model is an unconstrained MLP, so nothing forces post_dir_local / @@ -649,15 +712,19 @@ def constraint_report(collection: SampleCollection, norm_tol: float = 0.05) -> p }, ] for j, name in enumerate(RAW_TARGET_NAMES[:_N_LOG_DIMS]): - rows.append({ - "check": f"{name} >= 0", - "violation_rate": float(np.mean(gen[:, j] < 0)), - "mean_abs_error": float(np.mean(np.clip(-gen[:, j], 0, None))), - }) + rows.append( + { + "check": f"{name} >= 0", + "violation_rate": float(np.mean(gen[:, j] < 0)), + "mean_abs_error": float(np.mean(np.clip(-gen[:, j], 0, None))), + } + ) return pd.DataFrame(rows) -def constraint_report_pl(source: str | Path | pl.LazyFrame, norm_tol: float = 0.05) -> pl.DataFrame: +def constraint_report_pl( + source: str | Path | pl.LazyFrame, norm_tol: float = 0.05 +) -> pl.DataFrame: """Polars duplicate of `constraint_report`, reading straight from a predict parquet. Computed as a single set of lazy aggregation expressions over the @@ -671,16 +738,30 @@ def constraint_report_pl(source: str | Path | pl.LazyFrame, norm_tol: float = 0. travel_norm = sum(pl.col(pred_cols[k]) ** 2 for k in range(6, 9)).sqrt() raw_log_dims = [_raw_expr(pred_cols[j], j) for j in range(_N_LOG_DIMS)] - agg = lf.select([ - ((post_norm - 1).abs() > norm_tol).mean().alias("post_dir_violation_rate"), - (post_norm - 1).abs().mean().alias("post_dir_mean_abs_error"), - ((travel_norm - 1).abs() > norm_tol).mean().alias("travel_dir_violation_rate"), - (travel_norm - 1).abs().mean().alias("travel_dir_mean_abs_error"), - *[(raw < 0).mean().alias(f"{name}_violation_rate") - for raw, name in zip(raw_log_dims, RAW_TARGET_NAMES[:_N_LOG_DIMS])], - *[raw.clip(upper_bound=0).abs().mean().alias(f"{name}_mean_abs_error") - for raw, name in zip(raw_log_dims, RAW_TARGET_NAMES[:_N_LOG_DIMS])], - ]).collect().row(0, named=True) + agg = ( + lf.select( + [ + ((post_norm - 1).abs() > norm_tol) + .mean() + .alias("post_dir_violation_rate"), + (post_norm - 1).abs().mean().alias("post_dir_mean_abs_error"), + ((travel_norm - 1).abs() > norm_tol) + .mean() + .alias("travel_dir_violation_rate"), + (travel_norm - 1).abs().mean().alias("travel_dir_mean_abs_error"), + *[ + (raw < 0).mean().alias(f"{name}_violation_rate") + for raw, name in zip(raw_log_dims, RAW_TARGET_NAMES[:_N_LOG_DIMS]) + ], + *[ + raw.clip(upper_bound=0).abs().mean().alias(f"{name}_mean_abs_error") + for raw, name in zip(raw_log_dims, RAW_TARGET_NAMES[:_N_LOG_DIMS]) + ], + ] + ) + .collect() + .row(0, named=True) + ) rows = [ { @@ -695,11 +776,13 @@ def constraint_report_pl(source: str | Path | pl.LazyFrame, norm_tol: float = 0. }, ] for name in RAW_TARGET_NAMES[:_N_LOG_DIMS]: - rows.append({ - "check": f"{name} >= 0", - "violation_rate": agg[f"{name}_violation_rate"], - "mean_abs_error": agg[f"{name}_mean_abs_error"], - }) + rows.append( + { + "check": f"{name} >= 0", + "violation_rate": agg[f"{name}_violation_rate"], + "mean_abs_error": agg[f"{name}_mean_abs_error"], + } + ) return pl.DataFrame(rows) @@ -711,7 +794,10 @@ def plot_constraint_violations(collection: SampleCollection): n_panels = 2 + _N_LOG_DIMS fig, axes = plt.subplots(1, n_panels, figsize=(4 * n_panels, 3.5)) - for ax, norm, title in [(axes[0], post_norm, "||post_dir||"), (axes[1], travel_norm, "||travel_dir||")]: + for ax, norm, title in [ + (axes[0], post_norm, "||post_dir||"), + (axes[1], travel_norm, "||travel_dir||"), + ]: ax.hist(norm, bins=_hist_edges(norm, bins=50)) ax.axvline(1.0, color="k", linestyle="--", linewidth=1) ax.set_title(title) diff --git a/giant/cli.py b/giant/cli.py index 88cbf14..a248e2e 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -54,9 +54,16 @@ class Coord(str, Enum): @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, + 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, @@ -64,25 +71,55 @@ 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, - validate_every: Annotated[Optional[int], typer.Option(help="Run marginal+KL validation every N epochs (0 disables)")] = 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, + seed: Annotated[ + Optional[int], typer.Option(help="Random seed for reproducibility") + ] = None, + validate_every: Annotated[ + Optional[int], + typer.Option(help="Run marginal+KL validation every N epochs (0 disables)"), + ] = 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, + resume: Annotated[ + Optional[Path], typer.Option(help="Checkpoint .pt to resume training from") + ] = None, ) -> None: """Train the GIANT surrogate model.""" - 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, "seed": seed, - "validate_every": validate_every, - }.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 = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, cli_train, cli_model) + 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, + "seed": seed, + "validate_every": validate_every, + }.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 = 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 gconfig.auto_device() @@ -99,26 +136,45 @@ def train( typer.echo(f"out_dir: {out_dir}") 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, + data=data, + cfg=cfg, + out_dir=out_dir, + device=_device, + shuffle_buffer=shuffle_buffer, + num_workers=t["num_workers"], + resume=resume, + echo=typer.echo, ) @app.command() def predict( - data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")], - checkpoint: Annotated[Path, typer.Option(help="Path to checkpoint .pt file (best.pt or last.pt)")], - coord: Annotated[Coord, typer.Option( - help="global: full physical units, world frame (default). " - "local: raw 9D model output (denormalised only, local frame, " - "log-scaled scalars) alongside the matching ground-truth target " - "for the same input file — requires post-step columns." - )] = Coord.global_, - out: Annotated[Optional[Path], typer.Option(help="Output parquet path (default: _predicted[_local].parquet)")] = None, + data: Annotated[ + Path, typer.Argument(help="Parquet file or directory of parquet files") + ], + checkpoint: Annotated[ + Path, typer.Option(help="Path to checkpoint .pt file (best.pt or last.pt)") + ], + coord: Annotated[ + Coord, + typer.Option( + help="global: full physical units, world frame (default). " + "local: raw 9D model output (denormalised only, local frame, " + "log-scaled scalars) alongside the matching ground-truth target " + "for the same input file — requires post-step columns." + ), + ] = Coord.global_, + out: Annotated[ + Optional[Path], + typer.Option( + help="Output parquet path (default: _predicted[_local].parquet)" + ), + ] = None, batch_size: Annotated[int, typer.Option(help="Inference batch size")] = 4096, steps: Annotated[int, typer.Option(help="Flow matching ODE steps")] = 10, - device: Annotated[Optional[str], typer.Option(help="cpu | cuda | mps (default: auto)")] = None, + 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 gconfig.auto_device() @@ -127,7 +183,10 @@ def predict( # --- Load checkpoint --- ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False) if "model_config" not in ckpt: - typer.echo("error: checkpoint has no model_config — retrain with the current code", err=True) + typer.echo( + "error: checkpoint has no model_config — retrain with the current code", + err=True, + ) raise typer.Exit(1) model_cfg = ckpt["model_config"] @@ -145,7 +204,9 @@ def predict( # --- Output path --- if out is None: stem = data.stem if data.is_file() else data.name - suffix = "_predicted_local.parquet" if coord == Coord.local else "_predicted.parquet" + suffix = ( + "_predicted_local.parquet" if coord == Coord.local else "_predicted.parquet" + ) out = data.parent / f"{stem}{suffix}" typer.echo(f"output: {out}") @@ -162,10 +223,14 @@ def predict( N = len(chunk["event_id"]) if coord == Coord.local: - cond_cont, cond_cat, target_raw, _, _ = build_features(chunk, pdg_map, mat_map) + cond_cont, cond_cat, target_raw, _, _ = build_features( + chunk, pdg_map, mat_map + ) cond_cont = cond_norm.transform(cond_cont) else: - cond_cont, cond_cat = build_cond_features(chunk, pdg_map, mat_map, cond_norm) + cond_cont, cond_cat = build_cond_features( + chunk, pdg_map, mat_map, cond_norm + ) # Inference in batch_size slices pred_parts = [] @@ -180,22 +245,30 @@ def predict( raw = tgt_norm.inverse_transform(pred) if coord == Coord.local: - table = pa.table({ - "event_id": chunk["event_id"], - "pdg": chunk["pdg"], - "pre_x": chunk["pre_pos"][:, 0], - "pre_y": chunk["pre_pos"][:, 1], - "pre_z": chunk["pre_pos"][:, 2], - "pre_E": chunk["pre_E"], - "pre_dx": chunk["pre_dir"][:, 0], - "pre_dy": chunk["pre_dir"][:, 1], - "pre_dz": chunk["pre_dir"][:, 2], - "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)}, - }) + table = pa.table( + { + "event_id": chunk["event_id"], + "pdg": chunk["pdg"], + "pre_x": chunk["pre_pos"][:, 0], + "pre_y": chunk["pre_pos"][:, 1], + "pre_z": chunk["pre_pos"][:, 2], + "pre_E": chunk["pre_E"], + "pre_dx": chunk["pre_dir"][:, 0], + "pre_dy": chunk["pre_dir"][:, 1], + "pre_dz": chunk["pre_dir"][:, 2], + "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) + }, + } + ) else: step_length = inv_log_transform(raw[:, 0]) delta_e = inv_log_transform(raw[:, 1]) @@ -205,7 +278,9 @@ def predict( post_dir_local = raw[:, 3:6].copy() norms = np.linalg.norm(post_dir_local, axis=1, keepdims=True) post_dir_local /= np.where(norms < 1e-8, 1.0, norms) - post_dir_world = inv_local_frame_rotation(chunk["pre_dir"], post_dir_local) + post_dir_world = inv_local_frame_rotation( + chunk["pre_dir"], post_dir_local + ) # Same for the travel direction, then reconstruct post_pos from # the single shared step_length so the two stay consistent. @@ -216,34 +291,38 @@ def predict( chunk["pre_pos"], chunk["pre_dir"], step_length, travel_dir_local ) - table = pa.table({ - "event_id": chunk["event_id"], - "pdg": chunk["pdg"], - "pre_x": chunk["pre_pos"][:, 0], - "pre_y": chunk["pre_pos"][:, 1], - "pre_z": chunk["pre_pos"][:, 2], - "pre_E": chunk["pre_E"], - "pre_dx": chunk["pre_dir"][:, 0], - "pre_dy": chunk["pre_dir"][:, 1], - "pre_dz": chunk["pre_dir"][:, 2], - "material": chunk["material"], - "layer_id": chunk["layer_id"], - "n_sec": chunk["n_sec"], - "step_length": step_length, - "delta_e": delta_e, - "edep": edep, - "post_dx": post_dir_world[:, 0], - "post_dy": post_dir_world[:, 1], - "post_dz": post_dir_world[:, 2], - "post_x": post_pos_world[:, 0], - "post_y": post_pos_world[:, 1], - "post_z": post_pos_world[:, 2], - }) + table = pa.table( + { + "event_id": chunk["event_id"], + "pdg": chunk["pdg"], + "pre_x": chunk["pre_pos"][:, 0], + "pre_y": chunk["pre_pos"][:, 1], + "pre_z": chunk["pre_pos"][:, 2], + "pre_E": chunk["pre_E"], + "pre_dx": chunk["pre_dir"][:, 0], + "pre_dy": chunk["pre_dir"][:, 1], + "pre_dz": chunk["pre_dir"][:, 2], + "material": chunk["material"], + "layer_id": chunk["layer_id"], + "n_sec": chunk["n_sec"], + "step_length": step_length, + "delta_e": delta_e, + "edep": edep, + "post_dx": post_dir_world[:, 0], + "post_dy": post_dir_world[:, 1], + "post_dz": post_dir_world[:, 2], + "post_x": post_pos_world[:, 0], + "post_y": post_pos_world[:, 1], + "post_z": post_pos_world[:, 2], + } + ) - table = table.replace_schema_metadata({ - PREDICT_COORD_METADATA_KEY: coord.value, - PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION, - }) + table = table.replace_schema_metadata( + { + PREDICT_COORD_METADATA_KEY: coord.value, + PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION, + } + ) if writer is None: writer = pq.ParquetWriter(out, table.schema) @@ -255,5 +334,6 @@ def predict( typer.echo(f"wrote {total:,} rows → {out}") + if __name__ == "__main__": app() diff --git a/giant/config.py b/giant/config.py index c11106a..0872044 100644 --- a/giant/config.py +++ b/giant/config.py @@ -10,20 +10,33 @@ 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, "validate_every": 10, + "mode": "flow", + "epochs": 100, + "batch_size": 4096, + "lr": 3e-4, + "val_fraction": 0.1, + "num_workers": 4, + "seed": 0, + "validate_every": 10, }, "model": { - "hidden_dim": 256, "n_blocks": 6, "emb_dim": 16, "dropout": 0.1, + "hidden_dim": 256, + "n_blocks": 6, + "emb_dim": 16, + "dropout": 0.1, }, } def git_hash() -> str: try: - return subprocess.check_output( - ["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL - ).decode().strip() + return ( + subprocess.check_output( + ["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL + ) + .decode() + .strip() + ) except Exception: return "unknown" diff --git a/giant/data/dataset.py b/giant/data/dataset.py index c97b8dc..05feed6 100644 --- a/giant/data/dataset.py +++ b/giant/data/dataset.py @@ -124,7 +124,9 @@ class StreamingStepsDataset(IterableDataset): chunk = {k: v[mask] for k, v in chunk.items()} cond_cont, cond_cat, target, _, _ = build_features( - chunk, self.pdg_map, self.mat_map, + chunk, + self.pdg_map, + self.mat_map, cond_normalizer=self.cond_normalizer, target_normalizer=self.target_normalizer, ) diff --git a/giant/data/loader.py b/giant/data/loader.py index 2c908e9..9ee8506 100644 --- a/giant/data/loader.py +++ b/giant/data/loader.py @@ -51,10 +51,18 @@ def iter_file_chunks(path: str | Path) -> Iterator[dict[str, np.ndarray]]: _COND_COLS = [ - "event_id", "pdg", - "pre_x", "pre_y", "pre_z", "pre_E", - "pre_dx", "pre_dy", "pre_dz", - "material", "layer_id", "child_track_ids", + "event_id", + "pdg", + "pre_x", + "pre_y", + "pre_z", + "pre_E", + "pre_dx", + "pre_dy", + "pre_dz", + "material", + "layer_id", + "child_track_ids", ] diff --git a/giant/data/transforms.py b/giant/data/transforms.py index f92b14b..c78903a 100644 --- a/giant/data/transforms.py +++ b/giant/data/transforms.py @@ -21,7 +21,7 @@ def local_frame_rotation(pre_dir: np.ndarray, post_dir: np.ndarray) -> np.ndarra z = np.array([[0.0, 0.0, 1.0]], dtype=np.float32) cos_t = np.clip((pre_dir * z).sum(axis=1, keepdims=True), -1.0, 1.0) # (N,1) - sin_t = np.sqrt(np.maximum(0.0, 1.0 - cos_t ** 2)) # (N,1) + sin_t = np.sqrt(np.maximum(0.0, 1.0 - cos_t**2)) # (N,1) axis = np.cross(pre_dir, z) # (N,3); zero when pre_dir ∥ ẑ axis_norm = np.linalg.norm(axis, axis=1, keepdims=True) # (N,1) @@ -34,7 +34,9 @@ def local_frame_rotation(pre_dir: np.ndarray, post_dir: np.ndarray) -> np.ndarra kxv = np.cross(axis, post_dir) # (N,3) kdv = (axis * post_dir).sum(axis=1, keepdims=True) # (N,1) - return (post_dir * cos_t + kxv * sin_t + axis * kdv * (1.0 - cos_t)).astype(np.float32) + return (post_dir * cos_t + kxv * sin_t + axis * kdv * (1.0 - cos_t)).astype( + np.float32 + ) class Normalizer: @@ -109,7 +111,9 @@ def travel_direction(pre_pos: np.ndarray, post_pos: np.ndarray) -> np.ndarray: disp = post_pos - pre_pos norm = np.linalg.norm(disp, axis=1, keepdims=True) safe_norm = np.where(norm < 1e-7, 1.0, norm) - return np.where(norm < 1e-7, np.array([[0.0, 0.0, 1.0]]), disp / safe_norm).astype(np.float32) + return np.where(norm < 1e-7, np.array([[0.0, 0.0, 1.0]]), disp / safe_norm).astype( + np.float32 + ) def reconstruct_post_pos( @@ -128,7 +132,9 @@ def reconstruct_post_pos( return (pre_pos + step_length.reshape(-1, 1) * travel_dir_world).astype(np.float32) -def inv_local_frame_rotation(pre_dir: np.ndarray, post_dir_local: np.ndarray) -> np.ndarray: +def inv_local_frame_rotation( + pre_dir: np.ndarray, post_dir_local: np.ndarray +) -> np.ndarray: """Inverse of local_frame_rotation: rotate from local frame back to world frame. Applies R^T (same axis, negative angle) to post_dir_local. @@ -136,7 +142,7 @@ def inv_local_frame_rotation(pre_dir: np.ndarray, post_dir_local: np.ndarray) -> z = np.array([[0.0, 0.0, 1.0]], dtype=np.float32) cos_t = np.clip((pre_dir * z).sum(axis=1, keepdims=True), -1.0, 1.0) - sin_t = np.sqrt(np.maximum(0.0, 1.0 - cos_t ** 2)) + sin_t = np.sqrt(np.maximum(0.0, 1.0 - cos_t**2)) axis = np.cross(pre_dir, z) axis_norm = np.linalg.norm(axis, axis=1, keepdims=True) @@ -147,7 +153,9 @@ def inv_local_frame_rotation(pre_dir: np.ndarray, post_dir_local: np.ndarray) -> kdv = (axis * post_dir_local).sum(axis=1, keepdims=True) # Negative angle: sin_t → -sin_t - return (post_dir_local * cos_t - kxv * sin_t + axis * kdv * (1.0 - cos_t)).astype(np.float32) + return (post_dir_local * cos_t - kxv * sin_t + axis * kdv * (1.0 - cos_t)).astype( + np.float32 + ) def build_cond_features( @@ -157,13 +165,15 @@ def build_cond_features( cond_normalizer: "Normalizer | None" = None, ) -> tuple[np.ndarray, np.ndarray]: """Build conditioning arrays only — no target, no post-step variables.""" - cond_cont = np.column_stack([ - data["pre_pos"], - log_transform(data["pre_E"]), - data["pre_dir"], - data["layer_id"].astype(np.float32), - data["n_sec"].astype(np.float32), - ]).astype(np.float32) + cond_cont = np.column_stack( + [ + data["pre_pos"], + log_transform(data["pre_E"]), + data["pre_dir"], + data["layer_id"].astype(np.float32), + data["n_sec"].astype(np.float32), + ] + ).astype(np.float32) pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64) mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64) @@ -192,21 +202,25 @@ def build_features( data["pre_dir"], travel_direction(data["pre_pos"], data["post_pos"]) ) - target = np.column_stack([ - log_transform(data["step_length"]), - log_transform(data["delta_e"]), - log_transform(data["edep"]), - post_dir_local, - travel_dir_local, - ]).astype(np.float32) # (N, 9) + target = np.column_stack( + [ + log_transform(data["step_length"]), + log_transform(data["delta_e"]), + log_transform(data["edep"]), + post_dir_local, + travel_dir_local, + ] + ).astype(np.float32) # (N, 9) - cond_cont = np.column_stack([ - data["pre_pos"], - log_transform(data["pre_E"]), - data["pre_dir"], - data["layer_id"].astype(np.float32), - data["n_sec"].astype(np.float32), - ]).astype(np.float32) # (N, 9) + cond_cont = np.column_stack( + [ + data["pre_pos"], + log_transform(data["pre_E"]), + data["pre_dir"], + data["layer_id"].astype(np.float32), + data["n_sec"].astype(np.float32), + ] + ).astype(np.float32) # (N, 9) pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64) mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64) diff --git a/giant/model/network.py b/giant/model/network.py index e619fdb..1fa4ec5 100644 --- a/giant/model/network.py +++ b/giant/model/network.py @@ -12,7 +12,9 @@ class SinusoidalEmbedding(nn.Module): assert dim % 2 == 0, "dim must be even" half = dim // 2 freqs = torch.exp( - -math.log(10000) * torch.arange(half, dtype=torch.float32) / max(half - 1, 1) + -math.log(10000) + * torch.arange(half, dtype=torch.float32) + / max(half - 1, 1) ) self.register_buffer("freqs", freqs) @@ -90,9 +92,12 @@ class DenoisingMLP(nn.Module): ) merged_cond_dim = time_dim + cond_out_dim self.input_proj = nn.Linear(x_dim, hidden_dim) - self.blocks = nn.ModuleList([ - ResBlock(hidden_dim, merged_cond_dim, dropout=dropout) for _ in range(n_blocks) - ]) + self.blocks = nn.ModuleList( + [ + ResBlock(hidden_dim, merged_cond_dim, dropout=dropout) + for _ in range(n_blocks) + ] + ) self.out_proj = nn.Linear(hidden_dim, x_dim) def forward( @@ -102,9 +107,9 @@ class DenoisingMLP(nn.Module): cond_cont: torch.Tensor, cond_cat: torch.Tensor, ) -> torch.Tensor: - t_emb = self.time_emb(t) # (B, time_dim) - c_emb = self.cond_enc(cond_cont, cond_cat) # (B, cond_out_dim) - cond = torch.cat([t_emb, c_emb], dim=-1) # (B, time_dim+cond_out_dim) + t_emb = self.time_emb(t) # (B, time_dim) + c_emb = self.cond_enc(cond_cont, cond_cat) # (B, cond_out_dim) + cond = torch.cat([t_emb, c_emb], dim=-1) # (B, time_dim+cond_out_dim) x = self.input_proj(x_t) for block in self.blocks: x = block(x, cond) diff --git a/giant/model/schedule.py b/giant/model/schedule.py index 8352847..a4e942f 100644 --- a/giant/model/schedule.py +++ b/giant/model/schedule.py @@ -11,7 +11,9 @@ class CosineSchedule: steps = np.arange(T + 1, dtype=np.float64) f = np.cos(((steps / T + s) / (1.0 + s)) * np.pi / 2.0) ** 2 alpha_bars = (f / f[0]).astype(np.float32) - betas = np.clip(1.0 - alpha_bars[1:] / alpha_bars[:-1], 0.0, 0.999).astype(np.float32) + betas = np.clip(1.0 - alpha_bars[1:] / alpha_bars[:-1], 0.0, 0.999).astype( + np.float32 + ) self.betas = torch.from_numpy(betas) self.alphas = torch.from_numpy((1.0 - betas)) diff --git a/giant/pipeline.py b/giant/pipeline.py index 5909c05..565ce00 100644 --- a/giant/pipeline.py +++ b/giant/pipeline.py @@ -38,7 +38,9 @@ def run_train_job( 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"]) + 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( @@ -67,16 +69,23 @@ def run_train_job( 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, + 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, + 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, + 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, ) @@ -85,17 +94,24 @@ def run_train_job( # 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, + 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, + 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"], + 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") @@ -113,16 +129,23 @@ def run_train_job( 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"], + "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, + 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()}, diff --git a/giant/sample.py b/giant/sample.py index e9b45c4..e922057 100644 --- a/giant/sample.py +++ b/giant/sample.py @@ -43,11 +43,9 @@ def sample_ddpm( alpha = schedule.alphas[i] alpha_bar = schedule.alpha_bars[i] z = torch.randn_like(x) if i > 0 else torch.zeros_like(x) - x = ( - (1.0 / alpha.sqrt()) - * (x - (1.0 - alpha) / (1.0 - alpha_bar).sqrt() * eps_pred) - + beta.sqrt() * z - ) + x = (1.0 / alpha.sqrt()) * ( + x - (1.0 - alpha) / (1.0 - alpha_bar).sqrt() * eps_pred + ) + beta.sqrt() * z return x diff --git a/giant/train.py b/giant/train.py index 3cace01..2db15fa 100644 --- a/giant/train.py +++ b/giant/train.py @@ -29,7 +29,8 @@ class _GracefulShutdown: def __init__(self) -> None: self.requested = False self._previous: dict[ - int, Callable[[int, FrameType | None], object] | signal.Handlers | int | None + int, + Callable[[int, FrameType | None], object] | signal.Handlers | int | None, ] = {} def __enter__(self) -> "_GracefulShutdown": @@ -155,15 +156,22 @@ def train( 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_writer.writerow( + { + "epoch": epoch, + "train_loss": train_loss, + "val_loss": val_loss, + "lr": current_lr, + "epoch_time_s": epoch_time, + } + ) metrics_file.flush() if validate_every > 0 and epoch % validate_every == 0: print(f"[epoch {epoch}] marginal validation:") - validate_marginals(model, val_loader, mode=mode, schedule=ddpm_schedule, device=device) + validate_marginals( + model, val_loader, mode=mode, schedule=ddpm_schedule, device=device + ) ckpt: dict = { "model": model.state_dict(), diff --git a/giant/validate.py b/giant/validate.py index d2ed619..1814b91 100644 --- a/giant/validate.py +++ b/giant/validate.py @@ -6,7 +6,9 @@ from giant.constants import LOCAL_TARGET_NAMES from giant.sample import sample_flow, sample_ddpm, sample_ddim -def _histogram_kl(p_samples: np.ndarray, q_samples: np.ndarray, bins: int = 50, eps: float = 1e-8) -> float: +def _histogram_kl( + p_samples: np.ndarray, q_samples: np.ndarray, bins: int = 50, eps: float = 1e-8 +) -> float: """KL(P || Q) between two 1D samples, estimated via a shared histogram.""" lo = min(p_samples.min(), q_samples.min()) hi = max(p_samples.max(), q_samples.max()) @@ -61,10 +63,12 @@ def validate_marginals( real = np.concatenate(all_real, axis=0) generated = np.concatenate(all_gen, axis=0) - kl_divergence = np.array([ - _histogram_kl(real[:, j], generated[:, j], bins=kl_bins) - for j in range(real.shape[1]) - ]) + kl_divergence = np.array( + [ + _histogram_kl(real[:, j], generated[:, j], bins=kl_bins) + for j in range(real.shape[1]) + ] + ) header = ( f"{'Dim':<20} {'real_mean':>10} {'gen_mean':>10} " diff --git a/scripts/steps_to_parquet.py b/scripts/steps_to_parquet.py index 6c70ab7..5d848f0 100644 --- a/scripts/steps_to_parquet.py +++ b/scripts/steps_to_parquet.py @@ -96,7 +96,9 @@ def main() -> None: help="Uproot read batch size (default: '100 MB'). E.g. '50 MB', '500000' (rows).", ) parser.add_argument( - "--tree", default="Steps", help="Tree name inside the ROOT file (default: Steps)" + "--tree", + default="Steps", + help="Tree name inside the ROOT file (default: Steps)", ) parser.add_argument( "--compression", diff --git a/scripts/train.py b/scripts/train.py index 84b3540..067ceec 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -10,7 +10,9 @@ from giant.pipeline import run_train_job 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 parquet file or directory") + 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) @@ -18,49 +20,91 @@ def main() -> None: parser.add_argument("--hidden-dim", type=int) parser.add_argument("--n-blocks", type=int) parser.add_argument("--emb-dim", type=int) - parser.add_argument("--dropout", type=float, help="Dropout probability in ResBlocks (default: 0.1)") + parser.add_argument( + "--dropout", type=float, help="Dropout probability in ResBlocks (default: 0.1)" + ) parser.add_argument("--val-fraction", type=float) parser.add_argument("--seed", type=int, help="Random seed for reproducibility") - parser.add_argument("--validate-every", type=int, - help="Run marginal+KL validation every N epochs (0 disables)") - 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( + "--validate-every", + type=int, + help="Run marginal+KL validation every N epochs (0 disables)", + ) + 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") + parser.add_argument( + "--resume", default=None, help="Checkpoint .pt to resume training from" + ) args = parser.parse_args() - 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, "validate_every": args.validate_every, - }.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, - "dropout": args.dropout, - }.items() if v is not None} + 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, + "validate_every": args.validate_every, + }.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, + "dropout": args.dropout, + }.items() + if v is not None + } config_path = Path(args.config) if args.config else None - cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config_path, cli_train, cli_model) + 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 gconfig.auto_device() - out_dir = Path(args.out or ( - 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']}" - )) + out_dir = Path( + args.out + or ( + 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']}" + ) + ) print(f"device: {device}") print(f"out_dir: {out_dir}") 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, + 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, ) diff --git a/tests/test_analysis.py b/tests/test_analysis.py index fd2de86..04da8a7 100644 --- a/tests/test_analysis.py +++ b/tests/test_analysis.py @@ -39,27 +39,36 @@ def _unit_vectors(rng, n): def _make_collection(n=200, seed=0, gen_offset=0.0) -> SampleCollection: rng = np.random.default_rng(seed) - real = np.column_stack([ - rng.uniform(0.1, 5.0, n), # step_length - rng.uniform(0.1, 5.0, n), # delta_e - rng.uniform(0.1, 5.0, n), # edep - _unit_vectors(rng, n), # post_dir - _unit_vectors(rng, n), # travel_dir - ]).astype(np.float32) + real = np.column_stack( + [ + rng.uniform(0.1, 5.0, n), # step_length + rng.uniform(0.1, 5.0, n), # delta_e + rng.uniform(0.1, 5.0, n), # edep + _unit_vectors(rng, n), # post_dir + _unit_vectors(rng, n), # travel_dir + ] + ).astype(np.float32) gen = real + gen_offset pre_E = rng.uniform(1.0, 100.0, n).astype(np.float32) - cond_cont_raw = np.column_stack([ - rng.standard_normal((n, 3)), pre_E, rng.standard_normal((n, 3)), - rng.integers(0, 5, n), rng.integers(0, 3, n), - ]).astype(np.float32) + cond_cont_raw = np.column_stack( + [ + rng.standard_normal((n, 3)), + pre_E, + rng.standard_normal((n, 3)), + rng.integers(0, 5, n), + rng.integers(0, 3, n), + ] + ).astype(np.float32) return SampleCollection( cond_cont_raw=cond_cont_raw, pdg=rng.choice([11, -11, 22], size=n), material=rng.choice(["W", "Pb"], size=n), - real_raw=real, gen_raw=gen, - real_norm=real, gen_norm=gen, + real_raw=real, + gen_raw=gen, + real_norm=real, + gen_norm=gen, ) @@ -150,25 +159,35 @@ def _write_predicted_local_parquet(path, n=50, metadata=None, rng=None): """Mimic `giant predict --coord local`'s output schema for the loader tests.""" rng = rng or np.random.default_rng(0) true_log_local = rng.standard_normal((n, 9)).astype(np.float32) - true_log_local[:, :3] = log_transform(rng.uniform(0.1, 5.0, (n, 3)).astype(np.float32)) + true_log_local[:, :3] = log_transform( + rng.uniform(0.1, 5.0, (n, 3)).astype(np.float32) + ) pred_log_local = true_log_local + rng.normal(0, 0.01, (n, 9)).astype(np.float32) - table = pa.table({ - "event_id": rng.integers(0, 10, n), - "pdg": rng.choice([11, -11, 22], n), - "pre_x": rng.standard_normal(n).astype(np.float32), - "pre_y": rng.standard_normal(n).astype(np.float32), - "pre_z": rng.standard_normal(n).astype(np.float32), - "pre_E": rng.uniform(1.0, 100.0, n).astype(np.float32), - "pre_dx": rng.standard_normal(n).astype(np.float32), - "pre_dy": rng.standard_normal(n).astype(np.float32), - "pre_dz": rng.standard_normal(n).astype(np.float32), - "material": rng.choice(["W", "Pb"], n), - "layer_id": rng.integers(0, 10, n).astype(np.int32), - "n_sec": rng.integers(0, 3, n).astype(np.int32), - **{f"pred_{name}": pred_log_local[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)}, - **{f"true_{name}": true_log_local[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)}, - }) + table = pa.table( + { + "event_id": rng.integers(0, 10, n), + "pdg": rng.choice([11, -11, 22], n), + "pre_x": rng.standard_normal(n).astype(np.float32), + "pre_y": rng.standard_normal(n).astype(np.float32), + "pre_z": rng.standard_normal(n).astype(np.float32), + "pre_E": rng.uniform(1.0, 100.0, n).astype(np.float32), + "pre_dx": rng.standard_normal(n).astype(np.float32), + "pre_dy": rng.standard_normal(n).astype(np.float32), + "pre_dz": rng.standard_normal(n).astype(np.float32), + "material": rng.choice(["W", "Pb"], n), + "layer_id": rng.integers(0, 10, n).astype(np.int32), + "n_sec": rng.integers(0, 3, n).astype(np.int32), + **{ + f"pred_{name}": pred_log_local[:, j] + for j, name in enumerate(LOCAL_TARGET_NAMES) + }, + **{ + f"true_{name}": true_log_local[:, j] + for j, name in enumerate(LOCAL_TARGET_NAMES) + }, + } + ) if metadata is not None: table = table.replace_schema_metadata(metadata) pq.write_table(table, path) @@ -263,14 +282,21 @@ def test_marginal_table_pl_matches_numpy_version(tmp_path, group_by): path = _predicted_local_path(tmp_path) collection = load_predicted_local(path) - expected = marginal_table(collection, group_by=group_by).sort_values(["group", "dim"]) - actual = marginal_table_pl(path, group_by=group_by).sort(["group", "dim"]).to_pandas() + expected = marginal_table(collection, group_by=group_by).sort_values( + ["group", "dim"] + ) + actual = ( + marginal_table_pl(path, group_by=group_by).sort(["group", "dim"]).to_pandas() + ) assert list(expected["group"]) == list(actual["group"]) assert list(expected["n"]) == list(actual["n"]) for col in ["real_mean", "gen_mean", "real_std", "gen_std", "kl_real_gen"]: np.testing.assert_allclose( - expected[col].to_numpy(), actual[col].to_numpy(), atol=1e-4, rtol=1e-4, + expected[col].to_numpy(), + actual[col].to_numpy(), + atol=1e-4, + rtol=1e-4, ) @@ -290,8 +316,12 @@ def test_constraint_report_pl_matches_numpy_version(tmp_path): assert list(expected["check"]) == list(actual["check"]) np.testing.assert_allclose( - expected["violation_rate"].to_numpy(), actual["violation_rate"].to_numpy(), atol=1e-6, + expected["violation_rate"].to_numpy(), + actual["violation_rate"].to_numpy(), + atol=1e-6, ) np.testing.assert_allclose( - expected["mean_abs_error"].to_numpy(), actual["mean_abs_error"].to_numpy(), atol=1e-4, + expected["mean_abs_error"].to_numpy(), + actual["mean_abs_error"].to_numpy(), + atol=1e-4, ) diff --git a/tests/test_config.py b/tests/test_config.py index 768c42f..a035d3f 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -22,7 +22,10 @@ def test_merge_cli_overrides_applies_file_then_cli(tmp_path, monkeypatch): _write_config(path, "abc123") cfg = gconfig.merge_cli_overrides( - gconfig.DEFAULT_CONFIG, path, train_overrides={}, model_overrides={"hidden_dim": 128}, + gconfig.DEFAULT_CONFIG, + path, + train_overrides={}, + model_overrides={"hidden_dim": 128}, ) assert cfg["train"]["epochs"] == 5 # from file assert cfg["model"]["hidden_dim"] == 128 # CLI override wins over file @@ -42,7 +45,9 @@ def test_merge_cli_overrides_warns_on_git_hash_mismatch(tmp_path, monkeypatch, c assert "current999" in captured.err -def test_merge_cli_overrides_no_warning_on_matching_git_hash(tmp_path, monkeypatch, capsys): +def test_merge_cli_overrides_no_warning_on_matching_git_hash( + tmp_path, monkeypatch, capsys +): monkeypatch.setattr(gconfig, "git_hash", lambda: "same123") path = tmp_path / "config.toml" _write_config(path, "same123") @@ -51,7 +56,9 @@ def test_merge_cli_overrides_no_warning_on_matching_git_hash(tmp_path, monkeypat assert capsys.readouterr().err == "" -def test_merge_cli_overrides_no_warning_when_git_hash_unknown(tmp_path, monkeypatch, capsys): +def test_merge_cli_overrides_no_warning_when_git_hash_unknown( + tmp_path, monkeypatch, capsys +): monkeypatch.setattr(gconfig, "git_hash", lambda: "unknown") path = tmp_path / "config.toml" _write_config(path, "abc123") @@ -60,7 +67,9 @@ def test_merge_cli_overrides_no_warning_when_git_hash_unknown(tmp_path, monkeypa assert capsys.readouterr().err == "" -def test_merge_cli_overrides_no_warning_when_meta_section_absent(tmp_path, monkeypatch, capsys): +def test_merge_cli_overrides_no_warning_when_meta_section_absent( + tmp_path, monkeypatch, capsys +): monkeypatch.setattr(gconfig, "git_hash", lambda: "current999") path = tmp_path / "config.toml" path.write_text("[train]\nepochs = 5\n") @@ -69,7 +78,9 @@ def test_merge_cli_overrides_no_warning_when_meta_section_absent(tmp_path, monke assert capsys.readouterr().err == "" -def test_warn_if_checkpoint_config_mismatch_finds_sibling_toml(tmp_path, monkeypatch, capsys): +def test_warn_if_checkpoint_config_mismatch_finds_sibling_toml( + tmp_path, monkeypatch, capsys +): monkeypatch.setattr(gconfig, "git_hash", lambda: "current999") ckpt_path = tmp_path / "best.pt" ckpt_path.write_bytes(b"") # contents irrelevant, only its directory is used @@ -83,7 +94,9 @@ def test_warn_if_checkpoint_config_mismatch_finds_sibling_toml(tmp_path, monkeyp assert "current999" in captured.err -def test_warn_if_checkpoint_config_mismatch_no_warning_when_toml_absent(tmp_path, monkeypatch, capsys): +def test_warn_if_checkpoint_config_mismatch_no_warning_when_toml_absent( + tmp_path, monkeypatch, capsys +): monkeypatch.setattr(gconfig, "git_hash", lambda: "current999") ckpt_path = tmp_path / "best.pt" ckpt_path.write_bytes(b"") @@ -92,7 +105,9 @@ def test_warn_if_checkpoint_config_mismatch_no_warning_when_toml_absent(tmp_path assert capsys.readouterr().err == "" -def test_warn_if_checkpoint_config_mismatch_no_warning_when_hashes_match(tmp_path, monkeypatch, capsys): +def test_warn_if_checkpoint_config_mismatch_no_warning_when_hashes_match( + tmp_path, monkeypatch, capsys +): monkeypatch.setattr(gconfig, "git_hash", lambda: "same123") ckpt_path = tmp_path / "best.pt" ckpt_path.write_bytes(b"") diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 02a7eef..9e06fac 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -26,13 +26,17 @@ def test_dataset_item_shapes(): def test_split_sizes_sum_to_total(): data, cond_cont, cond_cat, target = _dummy(N=500) - train_ds, val_ds = train_val_split(data, cond_cont, cond_cat, target, val_fraction=0.2) + train_ds, val_ds = train_val_split( + data, cond_cont, cond_cat, target, val_fraction=0.2 + ) assert len(train_ds) + len(val_ds) == 500 def test_split_no_empty_sets(): data, cond_cont, cond_cat, target = _dummy(N=500, n_events=20) - train_ds, val_ds = train_val_split(data, cond_cont, cond_cat, target, val_fraction=0.2) + train_ds, val_ds = train_val_split( + data, cond_cont, cond_cat, target, val_fraction=0.2 + ) assert len(val_ds) > 0 assert len(train_ds) > 0 @@ -48,7 +52,9 @@ def test_split_event_leakage(): cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64) target = rng.standard_normal((N, 6)).astype(np.float32) - train_ds, val_ds = train_val_split(data, cond_cont, cond_cat, target, val_fraction=0.2) + train_ds, val_ds = train_val_split( + data, cond_cont, cond_cat, target, val_fraction=0.2 + ) # Recover which event_ids ended up in each split via the indices # (The dataset doesn't store event_ids, so we check via the original mask logic) diff --git a/tests/test_network.py b/tests/test_network.py index 52bc540..b2e1bd5 100644 --- a/tests/test_network.py +++ b/tests/test_network.py @@ -20,10 +20,13 @@ def test_denoising_mlp_output_shape(): x_t = torch.randn(B, 9) t = torch.rand(B) cond_cont = torch.randn(B, 9) - cond_cat = torch.stack([ - torch.randint(0, 5, (B,)), - torch.randint(0, 3, (B,)), - ], dim=1) + cond_cat = torch.stack( + [ + torch.randint(0, 5, (B,)), + torch.randint(0, 3, (B,)), + ], + dim=1, + ) out = model(x_t, t, cond_cont, cond_cat) assert out.shape == (B, 9) diff --git a/tests/test_transforms.py b/tests/test_transforms.py index 33a96c0..c8cf8e9 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -81,10 +81,14 @@ def test_reconstruct_post_pos_straight_line(): step_length = rng.uniform(0.1, 5.0, size=N).astype(np.float32) post_pos = pre_pos + step_length[:, None] * pre_dir - travel_dir_local = local_frame_rotation(pre_dir, travel_direction(pre_pos, post_pos)) + travel_dir_local = local_frame_rotation( + pre_dir, travel_direction(pre_pos, post_pos) + ) np.testing.assert_allclose(travel_dir_local, np.tile([0, 0, 1], (N, 1)), atol=1e-4) - reconstructed = reconstruct_post_pos(pre_pos, pre_dir, step_length, travel_dir_local) + reconstructed = reconstruct_post_pos( + pre_pos, pre_dir, step_length, travel_dir_local + ) np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4) @@ -98,8 +102,12 @@ def test_reconstruct_post_pos_general_roundtrip(): post_pos = pre_pos + rng.standard_normal((N, 3)).astype(np.float32) step_length = np.linalg.norm(post_pos - pre_pos, axis=1).astype(np.float32) - travel_dir_local = local_frame_rotation(pre_dir, travel_direction(pre_pos, post_pos)) - reconstructed = reconstruct_post_pos(pre_pos, pre_dir, step_length, travel_dir_local) + travel_dir_local = local_frame_rotation( + pre_dir, travel_direction(pre_pos, post_pos) + ) + reconstructed = reconstruct_post_pos( + pre_pos, pre_dir, step_length, travel_dir_local + ) np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)