Apply ruff format and document lint/type tooling in CLAUDE.md

First repo-wide ruff format pass, plus a note in CLAUDE.md to run
ruff and ty periodically.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-18 17:40:27 +02:00
parent 53fd2e4405
commit 8cebc4809d
20 changed files with 668 additions and 307 deletions
+10
View File
@@ -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.
+148 -62
View File
@@ -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)
+160 -80
View File
@@ -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: <data>_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: <data>_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()
+19 -6
View File
@@ -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"
+3 -1
View File
@@ -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,
)
+12 -4
View File
@@ -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",
]
+41 -27
View File
@@ -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)
+12 -7
View File
@@ -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)
+3 -1
View File
@@ -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))
+42 -19
View File
@@ -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()},
+3 -5
View File
@@ -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
+14 -6
View File
@@ -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(),
+9 -5
View File
@@ -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} "
+3 -1
View File
@@ -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",
+74 -30
View File
@@ -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,
)
+65 -35
View File
@@ -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,
)
+22 -7
View File
@@ -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"")
+9 -3
View File
@@ -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)
+7 -4
View File
@@ -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)
+12 -4
View File
@@ -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)