feat: add eval-cost benchmark — Geant4 reference vs surrogate rollout timing #92

Merged
lars merged 1 commits from eval-cost-benchmark into master 2026-08-31 12:01:17 +02:00
11 changed files with 354 additions and 7 deletions
+83
View File
@@ -45,6 +45,7 @@ import numpy as np
import polars as pl
from giant.analysis.context import Context
from giant.analysis.geant4_reference import GEANT4_REFERENCE, geant4_per_step_us
from giant.analysis.grouping import (
energy_bin_labels,
event_energy_bins,
@@ -125,6 +126,7 @@ class Bundle:
phys=physical_steps(r_all, Side.rollout),
checkpoint=rs.checkpoint,
type_embedding_l1_dist=rs.type_embedding_l1_dist,
timing=rs.timing,
)
return cls(ctx=ctx, rollouts=sides, t_all=t_all, t_phys=physical_steps(t_all, Side.reference))
@@ -973,6 +975,80 @@ def _sec_cos_angle_finalize(parts: list[dict], ctx: Context) -> Reduced:
)
# ---------------------------------------------------------------------------
# eval cost (not chunked — metadata-only, no row scan)
# ---------------------------------------------------------------------------
_EVAL_COST_LABELS = ["sampling / simulation", "parquet write / convert", "total"]
_EVAL_COST_NOTE = (
"no rollout in this run carries a `timing` block — re-run `giant rollout` "
"(timing instrumentation added after this checkpoint's rollout run) to "
"populate this plot"
)
def _eval_cost_per_step(b: Bundle) -> Reduced:
"""Per-rollout µs/physical-step vs the measured Geant4 reference.
``timing`` (``giant.cli``'s ``rollout`` command) is metadata carried on
the rollout YAML, not derived from the row data, so this needs no chunked
scan — same shape as the router diagnostics above.
"""
series: dict[str, list[float]] = {}
speedup: dict[str, float] = {}
for name, rs in b.rollouts.items():
t = rs.timing
if not t or t.get("us_per_step") is None:
continue
sample_us = t["us_per_step"]
write_us = t.get("write_us_per_step") or 0.0
series[name] = [sample_us, write_us, sample_us + write_us]
if not series:
return Reduced(
id="eval_cost_per_step",
family="cost",
kind="unavailable",
title="Eval cost per step: surrogate vs Geant4",
xlabel="n/a",
payload={"note": _EVAL_COST_NOTE},
)
g4 = geant4_per_step_us()
reference = [g4["sim_us_per_step"], g4["convert_us_per_step"], g4["total_us_per_step"]]
for name, vals in series.items():
speedup[name] = reference[-1] / vals[-1] if vals[-1] else float("inf")
return Reduced(
id="eval_cost_per_step",
family="cost",
kind="bar",
title="Eval cost per step: surrogate vs Geant4",
xlabel="phase",
payload={
"labels": _EVAL_COST_LABELS,
"series": series,
"reference": reference,
"ylabel": "µs per physical step",
"log_y": True,
},
meta={
"speedup_vs_geant4_total": speedup,
"geant4_provenance": GEANT4_REFERENCE["provenance"],
"caveat": (
"The Geant4 reference is measured single-threaded on one CPU core "
"(see giant.analysis.geant4_reference); a rollout's timing is "
"whatever device it actually ran on (see each series' device in "
"run_meta.json's plot_meta). This is a deployment-speedup ratio, "
"not a same-hardware or per-FLOP comparison."
),
},
)
_eval_cost_per_step_partial, _eval_cost_per_step_finalize = _unchunkable(_eval_cost_per_step)
# ---------------------------------------------------------------------------
# router diagnostics (not chunked — already bounded/subsampled)
# ---------------------------------------------------------------------------
@@ -1160,6 +1236,13 @@ def build_catalog() -> list[PlotSpec]:
compute_partial=_sec_cos_angle_partial,
finalize=_sec_cos_angle_finalize,
),
PlotSpec(
"eval_cost_per_step",
"cost",
compute_partial=_eval_cost_per_step_partial,
finalize=_eval_cost_per_step_finalize,
chunkable=False,
),
PlotSpec(
"router_gating",
"model",
+6 -3
View File
@@ -82,6 +82,7 @@ _PLOT_META_KEYS = (
"rollout_seed",
"n_rows",
"termination_reason_counts",
"timing",
"model_config",
"training_epoch",
"best_val_loss",
@@ -329,9 +330,9 @@ def compute_reduced(
) -> Path:
"""Core: run one (plot, chunk)'s partial reduction against explicit paths.
``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?},
...]``, one per rollout series (insertion order preserved through to every
plot's ``Reduced.payload["series"]``).
``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?,
"timing"?}, ...]``, one per rollout series (insertion order preserved
through to every plot's ``Reduced.payload["series"]``).
Writes a ``Partial`` JSON — the raw, not-yet-merged output of
``PlotSpec.compute_partial`` — never a finished ``Reduced``; ``merge_one``
@@ -352,6 +353,7 @@ def compute_reduced(
source=r["path"],
checkpoint=r.get("checkpoint"),
type_embedding_l1_dist=r.get("type_embedding_l1_dist"),
timing=r.get("timing"),
)
for r in rollouts
]
@@ -377,6 +379,7 @@ def compute_one(spec_id: str, run_dir: str | Path, chunk_index: int = 0) -> Path
"path": ro["path"],
"checkpoint": ro["plot_meta"].get("checkpoint"),
"type_embedding_l1_dist": ro["plot_meta"].get("type_embedding_l1_dist"),
"timing": ro["plot_meta"].get("timing"),
}
for ro in meta.rollouts
]
+76
View File
@@ -0,0 +1,76 @@
"""Measured Geant4 (miniCaloSim) per-step eval cost — the reference line for
``eval_cost_per_step`` in ``catalog.py``.
Mirrors the precedent set by ``runtime_estimate.py``'s ``_COST_MODEL``: a
constant table measured once on a specific machine and pasted in, with the
methodology and provenance recorded in this docstring rather than derived at
runtime (there is no live Geant4 install on the machines that run
``giant analyze``, and re-measuring per invocation would be both slow and
noisy see the module docstring precedent).
**Methodology** (``scratchpad/bench_geant4.py``, a one-off, not a `dwarf`
subcommand): ``run_pbwo4`` (the default homogeneous-PbWO4 miniCaloSim
executable, see ``~/Programming/minicalosim``) was timed at 3 beam energies
(1/10/50 GeV) and **4 event counts each**, converting each run's ROOT output
to Parquet with ``giant.tools.steps_to_parquet.convert_steps_to_parquet``
immediately after. Event counts were scaled down as energy rose (100/400/
1000/2000 at 1 GeV, 30/100/200/300 at 10 GeV, 10/25/45/60 at 50 GeV) to keep
every run's row count under ~8.1M — a naive 50/200 pair at 50 GeV produces
~27M steps and OOM'd the conversion step on a 14GB laptop. Per-energy linear
fits (``t = intercept + slope * n``) separate Geant4's one-time init (physics
tables, geometry construction) from its true marginal per-event cost the
slope, not a naive ``t / n_events`` from a single run, is what feeds
``sim_us_per_step`` below. The per-step denominator is the produced
``Steps``-tree/Parquet row count, matching the "physical step" unit
``giant rollout``'s ``timing.n_physical_rows`` uses on the surrogate side.
Both stages ran single-threaded (default Geant4 threading), pinned to one
CPU core.
``sim_us_per_step``/``convert_us_per_step``/``sim_ms_per_event`` below are
the mean across the 3 energies. With 4 event-count points per energy (up
from an initial 2-point pass, which had ~80% spread and nonsensical negative
fitted intercepts at 10/50 GeV an artifact of extrapolating a 2-point
line), both quantities are now energy-flat as physically expected:
``sim_us_per_step`` spread ~5%, ``convert_us_per_step`` spread ~13.5%. Treat
these as reliable to about that precision.
**Caveat hardware asymmetry**: this reference is single-core CPU. A
surrogate rollout's ``timing`` block will typically be measured on a batched
GPU. The resulting ratio in ``eval_cost_per_step`` is a *deployment* speedup
(what you'd actually see swapping Geant4 for the surrogate in a production
pipeline), not a same-hardware or per-FLOP comparison state this whenever
quoting the number.
**Staleness**: re-run ``scratchpad/bench_geant4.py`` (and update this file)
if measured on different hardware, after a miniCaloSim/Geant4 version bump,
or if this reference is more than a year or two stale.
"""
from __future__ import annotations
GEANT4_REFERENCE: dict = {
"sim_us_per_step": 11.2903,
"convert_us_per_step": 11.1014,
"sim_ms_per_event": 609.6848,
"provenance": {
"cpu": "AMD Ryzen 7 PRO 4750U with Radeon Graphics",
"geant4_version": "11.4.1",
"minicalosim_sha": "ea917da",
"measured": "2026-08-31",
"energies_gev": [1.0, 10.0, 50.0],
"spread_pct_sim": 4.96,
"spread_pct_convert": 13.52,
"threads": 1,
},
}
def geant4_per_step_us() -> dict[str, float]:
"""Sim / convert / total microseconds per physical step, from ``GEANT4_REFERENCE``."""
sim = GEANT4_REFERENCE["sim_us_per_step"]
convert = GEANT4_REFERENCE["convert_us_per_step"]
return {
"sim_us_per_step": sim,
"convert_us_per_step": convert,
"total_us_per_step": sim + convert,
}
+2
View File
@@ -274,6 +274,8 @@ def _render_bar(r: Reduced, params: dict):
ax.set_xticks(x)
ax.set_xticklabels(labels, rotation=45, ha="right")
ax.set_ylabel(r.payload.get("ylabel", "value"))
if r.payload.get("log_y"):
ax.set_yscale("log")
ps.style_legend(ax, title="source")
return fig
+4
View File
@@ -96,6 +96,10 @@ _COST_MODEL: dict[str, tuple[float, float]] = {
"sec_count_per_species": (0.0, 4.963e-07),
"sec_energy": (0.0, 4.727e-07),
"sec_cos_angle": (0.0, 2.749e-06),
# Metadata-only (YAML-carried `timing`, no row scan) — same shape as the
# router diagnostics' fixed cost, just cheaper since there's no live
# torch checkpoint to load.
"eval_cost_per_step": (0.0, 0.0),
}
+5
View File
@@ -110,6 +110,7 @@ class RolloutSpec:
source: str | Path | pl.LazyFrame
checkpoint: str | None = None
type_embedding_l1_dist: dict | None = None
timing: dict | None = None
@dataclass
@@ -124,6 +125,10 @@ class RolloutSide:
# only. Unlike checkpoint, this needs no live model: it's already a
# finished histogram, just passed through.
type_embedding_l1_dist: dict | None = None
# Wall-clock cost of this rollout run (giant.cli's rollout command),
# from the rollout YAML — eval_cost_per_step only. None on rollout runs
# that predate timing instrumentation.
timing: dict | None = None
def _check_rollout_metadata(path: Path) -> None:
+77 -3
View File
@@ -6,7 +6,7 @@ from enum import Enum
import math
from pathlib import Path
import re
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, Optional, cast
import uuid as uuid_mod
import typer
@@ -205,6 +205,45 @@ def _write_prediction_ref(
return ref_path
def _build_rollout_timing(
*,
setup_s: float,
rollout_s: float,
write_s: float,
n_rows: int,
termination_reason_counts: dict[str, int],
n_seed_events: int,
device: str,
torch_threads: int,
) -> dict:
"""Assemble ``giant rollout``'s ``timing`` sidecar block.
``n_physical_rows`` excludes the synthetic termination rows (escape/
unknown-pdg/energy-cutoff/max-steps markers `giant.rollout` emits but
Geant4 never does) so ``us_per_step`` is comparable to
``giant.analysis.geant4_reference``'s per-step Geant4 measurement — see
``giant/analysis/catalog.py``'s ``eval_cost_per_step`` spec.
"""
from giant.analysis.sources import SYNTHETIC_TERMINATION_REASONS
sample_s = rollout_s - write_s
n_synthetic_rows = sum(termination_reason_counts.get(reason, 0) for reason in SYNTHETIC_TERMINATION_REASONS)
n_physical_rows = n_rows - n_synthetic_rows
return {
"setup_s": setup_s,
"rollout_s": rollout_s,
"write_s": write_s,
"sample_s": sample_s,
"n_rows": n_rows,
"n_physical_rows": n_physical_rows,
"us_per_step": (sample_s / n_physical_rows * 1e6) if n_physical_rows else None,
"write_us_per_step": (write_s / n_physical_rows * 1e6) if n_physical_rows else None,
"ms_per_event": (rollout_s / n_seed_events * 1e3) if n_seed_events else None,
"device": device,
"torch_threads": torch_threads,
}
@app.callback()
def _main() -> None:
"""GIANT — Geant4 step-function surrogate."""
@@ -1455,6 +1494,8 @@ def rollout(
] = None,
) -> None:
"""Roll the surrogate forward into full showers (autoregressive)."""
import time
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
@@ -1464,7 +1505,9 @@ def rollout(
from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference
from giant.data.loader import find_parquet_files
from giant.geometry import GeometryOracle
from giant.rollout import L1DistCollector, rollout as run_rollout
from giant.rollout import L1DistCollector, RolloutSummary, rollout as run_rollout
_t_setup_start = time.perf_counter()
if seed is not None:
torch.manual_seed(seed)
@@ -1511,9 +1554,11 @@ def rollout(
# avg_tracks_per_event) — mirrors the row-group streaming `giant predict`
# already does on its input side.
writer: pq.ParquetWriter | None = None
_write_s = 0.0
def _write_chunk(row: dict[str, np.ndarray]) -> None:
nonlocal writer
nonlocal writer, _write_s
_t0 = time.perf_counter()
table = pa.table(row)
if writer is None:
table = table.replace_schema_metadata(
@@ -1524,11 +1569,14 @@ def rollout(
)
writer = pq.ParquetWriter(out, table.schema)
writer.write_table(table)
_write_s += time.perf_counter() - _t0
# Only meaningful under particle_type.target="embedding" — a
# no-op collector otherwise, cheaper than branching the call itself.
l1_dist_collector = L1DistCollector()
_setup_s = time.perf_counter() - _t_setup_start
_t_rollout_start = time.perf_counter()
summary = run_rollout(
model,
sec_decoder,
@@ -1560,6 +1608,23 @@ def rollout(
)
if writer is not None:
writer.close()
# on_chunk=_write_chunk is always passed above, so rollout() always
# returns the streaming-summary shape (RolloutSummary), never the
# materialized dict[str, np.ndarray] alternative its return type allows.
summary = cast(RolloutSummary, summary)
_rollout_s = time.perf_counter() - _t_rollout_start
timing = _build_rollout_timing(
setup_s=_setup_s,
rollout_s=_rollout_s,
write_s=_write_s,
n_rows=summary["n_rows"],
termination_reason_counts=summary["termination_reason_counts"],
n_seed_events=len(seeds["event_id"]),
device=str(_device),
torch_threads=torch.get_num_threads(),
)
_sample_s = timing["sample_s"]
n_physical_rows = timing["n_physical_rows"]
l1_summary = l1_dist_collector.summary()
@@ -1582,6 +1647,10 @@ def rollout(
"rollout_seed": seed,
"n_rows": summary["n_rows"],
"termination_reason_counts": summary["termination_reason_counts"],
# Wall-clock cost of this run, normalized per physical step (the
# comparable unit against giant.analysis.geant4_reference) — see
# eval_cost_per_step in giant/analysis/catalog.py.
"timing": timing,
# Diagnostic — only present under
# stage2_model.particle_type.target="embedding"; omitted (not
# written as null) otherwise, so giant.analysis can tell "not
@@ -1605,6 +1674,11 @@ def rollout(
typer.echo(f"wrote {summary['n_rows']:,} step rows → {out}")
typer.echo(f"terminations: {summary['termination_reason_counts']}")
if timing["us_per_step"] is not None:
typer.echo(
f"timing: {_rollout_s:.1f}s total ({_sample_s:.1f}s sample + {_write_s:.1f}s write), "
f"{timing['us_per_step']:.1f} us/step over {n_physical_rows:,} physical steps"
)
typer.echo(f"reference: {ref_path}")
+29
View File
@@ -261,3 +261,32 @@ def test_sec_count_per_step_by_species_zero_row_is_per_species(bundle):
for j, _ in enumerate(cols):
if j != g:
assert ref[0][j] == 3 and sum(row[j] for row in ref[1:]) == 0
# ---------------------------------------------------------------------------
# eval_cost_per_step
# ---------------------------------------------------------------------------
def test_eval_cost_per_step_unavailable_without_timing(bundle: Bundle):
# `bundle`'s RolloutSpec carries no `timing` -> no rollout to compare.
spec = get_spec("eval_cost_per_step")
r = spec.finalize([spec.compute_partial(bundle)], bundle.ctx)
assert r.kind == "unavailable"
assert r.payload["note"]
def test_eval_cost_per_step_bar_with_timing(ctx: Context):
spec = get_spec("eval_cost_per_step")
rs = RolloutSpec(
"rollout",
_rollout_frame(),
timing={"us_per_step": 12.5, "write_us_per_step": 2.5},
)
b = Bundle.open([rs], _reference_frame(), ctx)
r = spec.finalize([spec.compute_partial(b)], ctx)
assert r.kind == "bar"
assert r.payload["series"]["rollout"] == [12.5, 2.5, 15.0]
assert len(r.payload["reference"]) == 3
assert r.payload["log_y"] is True
assert "rollout" in r.meta["speedup_vs_geant4_total"]
+42 -1
View File
@@ -8,11 +8,52 @@ from __future__ import annotations
import torch
from typer.testing import CliRunner
from giant.cli import app
from giant.cli import _build_rollout_timing, app
runner = CliRunner()
def test_build_rollout_timing_excludes_synthetic_rows_from_per_step_cost():
# 100 rows total, 30 of them synthetic termination markers (escape) ->
# us_per_step should be normalized over the 70 physical rows only, the
# same unit giant.analysis.geant4_reference measures Geant4 in.
timing = _build_rollout_timing(
setup_s=1.0,
rollout_s=10.0,
write_s=2.0,
n_rows=100,
termination_reason_counts={"escaped": 30, "natural_end": 70},
n_seed_events=5,
device="cpu",
torch_threads=4,
)
assert timing["n_rows"] == 100
assert timing["n_physical_rows"] == 70
assert timing["n_physical_rows"] < timing["n_rows"]
assert timing["sample_s"] == 8.0 # rollout_s - write_s
assert timing["us_per_step"] == 8.0 / 70 * 1e6
assert timing["write_us_per_step"] == 2.0 / 70 * 1e6
assert timing["ms_per_event"] == 10.0 / 5 * 1e3
assert timing["device"] == "cpu" and timing["torch_threads"] == 4
def test_build_rollout_timing_handles_zero_physical_rows_and_events():
timing = _build_rollout_timing(
setup_s=1.0,
rollout_s=1.0,
write_s=0.0,
n_rows=5,
termination_reason_counts={"escaped": 5},
n_seed_events=0,
device="cpu",
torch_threads=1,
)
assert timing["n_physical_rows"] == 0
assert timing["us_per_step"] is None
assert timing["write_us_per_step"] is None
assert timing["ms_per_event"] is None
def test_rollout_exits_1_on_checkpoint_missing_model_config(tmp_path):
checkpoint = tmp_path / "bad.pt"
torch.save({"sec_decoder": {}, "normalizer": {"sec_phys": {}}}, checkpoint)
+16
View File
@@ -252,6 +252,22 @@ def test_compute_one_from_run_dir(tmp_path: Path):
assert list(partial.data["r"]) == ["rollout"]
def test_timing_survives_plot_meta_to_compute_one(tmp_path: Path):
yaml_path = _write_inputs(tmp_path)
d = yaml.safe_load(yaml_path.read_text())
d["timing"] = {"us_per_step": 7.0, "write_us_per_step": 1.0}
yaml_path.write_text(yaml.safe_dump(d))
run_dir = _prep([yaml_path])
meta = RunMeta.load(run_dir / "run_meta.json")
assert meta.rollouts[0]["plot_meta"]["timing"] == {"us_per_step": 7.0, "write_us_per_step": 1.0}
out = compute_one("eval_cost_per_step", run_dir)
reduced = Reduced(**Partial.load(out).data["reduced"])
assert reduced.kind == "bar"
assert reduced.payload["series"]["rollout"] == [7.0, 1.0, 8.0]
def test_compute_reduced_explicit_paths(tmp_path: Path):
run_dir = _prep([_write_inputs(tmp_path)])
meta = RunMeta.load(run_dir / "run_meta.json")
+14
View File
@@ -292,6 +292,20 @@ def test_render_one_of_each_kind(tmp_path: Path):
"ylabel": "frac",
},
),
Reduced(
"cost",
"cost",
"bar",
"Cost",
"phase",
{
"labels": ["sample", "write", "total"],
"series": {"flow": [10.0, 1.0, 11.0]},
"reference": [5.0, 0.5, 5.5],
"ylabel": "us/step",
"log_y": True,
},
),
Reduced(
"s",
"species",