feat: add eval-cost benchmark — Geant4 reference vs surrogate rollout timing #92
@@ -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",
|
||||
|
||||
@@ -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
|
||||
]
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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}")
|
||||
|
||||
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user