Add multi-rollout support to giant analyze (gitea #77)
CI / Lint (ruff check) (push) Successful in 32s
CI / Format (ruff format) (push) Successful in 30s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 35s
CI / Lint (ruff check) (pull_request) Successful in 33s
CI / Format (ruff format) (pull_request) Successful in 30s
CI / Type check (ty) (pull_request) Successful in 34s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Tests (push) Successful in 5m59s
CI / Bump version, tag, and update changelog on merge to master (push) Has been skipped
CI / Tests (pull_request) Successful in 4m22s
CI / Bump version, tag, and update changelog on merge to master (pull_request) Has been skipped

giant analyze compares N rollout YAMLs against one shared reference file
(all must name the same dataset, checked up front) instead of exactly one
rollout vs one reference, rendering each rollout as its own colored series
against a single reference line/panel. Series names come from a repeated
--label flag, else the YAML stem, else "rollout" for a single YAML — a
single-rollout run keeps rendering identically to before this change.

Bundle now holds a name-keyed dict of rollout sides instead of one fixed
pair, every catalog compute_partial/finalize builds a Reduced.payload
keyed the same way ("series": {name: ...}, "reference": ... as the one
distinguished non-rollout entry), and every renderer draws N series (or
N panels, for the two heatmap-shaped specs and the router/type-embedding
diagnostics, which are inherently one-matrix/one-checkpoint per rollout)
against the reference's fixed dashed-ink style.
This commit is contained in:
2026-08-24 13:23:50 +02:00
parent b8f8965338
commit ebd3e0dc71
18 changed files with 1346 additions and 623 deletions
+94 -31
View File
@@ -14,12 +14,21 @@ from giant.analysis.catalog import (
_ks_statistic,
)
from giant.analysis.context import Context, build_context
from giant.analysis.sources import RolloutSpec
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
def _build_ctx() -> Context:
r, t = _rollout_frame(), _reference_frame()
return build_context(r, t, n_energy_bins=2, n_marginal_bins=10, top_k_pdg=3, sample_rows=1000)
return build_context(
[RolloutSpec("rollout", r)], t, n_energy_bins=2, n_marginal_bins=10, top_k_pdg=3, sample_rows=1000
)
def _two_rollout_specs() -> list[RolloutSpec]:
# Two distinct rollout sources so multi-series merging/finalize code is
# exercised even though the underlying frame is the same fixture.
return [RolloutSpec("flow", _rollout_frame()), RolloutSpec("wgan", _rollout_frame())]
@pytest.fixture(scope="module")
@@ -27,9 +36,20 @@ def ctx() -> Context:
return _build_ctx()
@pytest.fixture(scope="module")
def two_ctx() -> Context:
t = _reference_frame()
return build_context(_two_rollout_specs(), t, n_energy_bins=2, n_marginal_bins=10, top_k_pdg=3, sample_rows=1000)
@pytest.fixture(scope="module")
def bundle(ctx: Context) -> Bundle:
return Bundle.open(_rollout_frame(), _reference_frame(), ctx)
return Bundle.open([RolloutSpec("rollout", _rollout_frame())], _reference_frame(), ctx)
@pytest.fixture(scope="module")
def two_bundle(two_ctx: Context) -> Bundle:
return Bundle.open(_two_rollout_specs(), _reference_frame(), two_ctx)
def test_catalog_ids_unique_and_nonempty():
@@ -64,46 +84,71 @@ def test_every_spec_computes_valid_reduced(bundle: Bundle):
"unavailable",
}
assert r.title and r.xlabel
_validate_payload(r)
_validate_payload(r, ["rollout"])
def _validate_payload(r) -> None:
def test_every_spec_computes_valid_reduced_with_two_rollouts(two_bundle: Bundle):
for spec in build_catalog():
r = spec.finalize([spec.compute_partial(two_bundle)], two_bundle.ctx)
assert r.id == spec.id
_validate_payload(r, ["flow", "wgan"])
def _validate_payload(r, names: list[str]) -> None:
p = r.payload
if r.kind == "overlay_hist":
n = len(p["edges"]) - 1
assert len(p["rollout"]) == n and len(p["reference"]) == n
assert list(p["series"]) == names
for v in p["series"].values():
assert len(v) == n
assert len(p["reference"]) == n
elif r.kind == "single_hist":
assert len(p["rollout"]) == len(p["edges"]) - 1
assert list(p["series"]) == names
for v in p["series"].values():
assert len(v) == len(p["edges"]) - 1
elif r.kind == "grouped_hist":
n = len(p["edges"]) - 1
assert p["groups"], "grouped hist must have at least one group"
for g in p["groups"].values():
assert len(g["rollout"]) == n and len(g["reference"]) == n
assert list(g["series"]) == names
for v in g["series"].values():
assert len(v) == n
assert len(g["reference"]) == n
elif r.kind == "profile":
n = len(p["edges"]) - 1
for k in ("rollout_mean", "rollout_std", "reference_mean", "reference_std"):
assert len(p[k]) == n
assert list(p["series"]) == names
for side in p["series"].values():
assert len(side["mean"]) == n and len(side["std"]) == n
assert len(p["reference"]["mean"]) == n and len(p["reference"]["std"]) == n
elif r.kind == "bar":
assert len(p["labels"]) == len(p["rollout"]) == len(p["reference"])
assert list(p["series"]) == names
for v in p["series"].values():
assert len(p["labels"]) == len(v)
assert len(p["labels"]) == len(p["reference"])
elif r.kind == "unavailable":
assert p["note"]
elif r.kind == "router_gating":
for side in ("rollout", "reference"):
if side in p:
assert len(p[side]["centers"]) == len(p[side]["means"])
elif r.kind == "router_share":
for cat in p["categories"]:
for entry in p["series"].values():
for side in ("rollout", "reference"):
if side in p:
assert cat in p[side]
if side in entry:
assert len(entry[side]["centers"]) == len(entry[side]["means"])
elif r.kind == "router_share":
for entry in p["series"].values():
for cat in entry["categories"]:
for side in ("rollout", "reference"):
if side in entry:
assert cat in entry[side]
elif r.kind == "router_specialization":
for side in ("rollout", "reference"):
if side in p:
assert len(p[side]["centers"]) == len(p[side]["score"])
for entry in p["series"].values():
for side in ("rollout", "reference"):
if side in entry:
assert len(entry[side]["centers"]) == len(entry[side]["score"])
elif r.kind == "heatmap":
assert len(p["matrix"]) == len(p["row_labels"])
for row in p["matrix"]:
assert len(row) == len(p["col_labels"])
assert list(p["series"]) == names
for mat in p["series"].values():
assert len(mat) == len(p["row_labels"])
for row in mat:
assert len(row) == len(p["col_labels"])
# ---------------------------------------------------------------------------
@@ -150,20 +195,21 @@ def _assert_payload_close(a, b, path: str = "payload") -> None:
@pytest.mark.parametrize("spec_id", _CHUNK_EQUIVALENCE_IDS)
def test_chunked_matches_unchunked(ctx: Context, spec_id: str):
def test_chunked_matches_unchunked(two_ctx: Context, spec_id: str):
"""A plot computed over N event-disjoint chunks then merged must equal the
same plot computed in one unchunked pass — the core chunking correctness
guarantee (see the analysis-rollout-plots chunking plan)."""
guarantee (see the analysis-rollout-plots chunking plan). Exercised with
two rollout series so the per-rollout merge path is covered too."""
spec: PlotSpec = get_spec(spec_id)
r, t = _rollout_frame(), _reference_frame()
rollouts, t = _two_rollout_specs(), _reference_frame()
unchunked_bundle = Bundle.open(r, t, ctx)
unchunked = spec.finalize([spec.compute_partial(unchunked_bundle)], ctx)
unchunked_bundle = Bundle.open(rollouts, t, two_ctx)
unchunked = spec.finalize([spec.compute_partial(unchunked_bundle)], two_ctx)
# 4 chunks over only 2 distinct event_ids also exercises empty chunks.
n_chunks = 4 if spec.chunkable else 1
parts = [spec.compute_partial(Bundle.open(r, t, ctx, chunk=(k, n_chunks))) for k in range(n_chunks)]
chunked = spec.finalize(parts, ctx)
parts = [spec.compute_partial(Bundle.open(rollouts, t, two_ctx, chunk=(k, n_chunks))) for k in range(n_chunks)]
chunked = spec.finalize(parts, two_ctx)
assert chunked.id == unchunked.id
assert chunked.kind == unchunked.kind
@@ -196,6 +242,14 @@ def test_integer_confusion_caps_pathological_outliers():
assert mat.sum() == 2
def test_integer_confusion_explicit_cap_overrides_local_range():
# Even though this pair's own max is 1, an explicit shared cap forces a
# wider (and so cross-rollout-consistent) label set.
labels, mat = _integer_confusion(np.array([1, 1]), np.array([0, 1]), cap=3)
assert labels == ["0", "1", "2", "3+"]
assert mat.shape == (4, 4)
def test_containment_depths_simple_ramp():
# one event, edep concentrated in the first bin -> 90%/95% containment
# depth is the first bin's right edge; a zero-energy event is dropped.
@@ -209,4 +263,13 @@ def test_n_sec_confusion_spec(bundle):
spec = get_spec("n_sec_confusion")
r = spec.finalize([spec.compute_partial(bundle)], bundle.ctx)
assert r.payload["row_labels"] == r.payload["col_labels"] == ["0", "1+"]
assert r.payload["matrix"] == [[0, 0], [1, 1]]
assert r.payload["series"]["rollout"] == [[0, 0], [1, 1]]
def test_n_sec_confusion_shares_one_cap_across_rollouts(two_bundle):
spec = get_spec("n_sec_confusion")
r = spec.finalize([spec.compute_partial(two_bundle)], two_bundle.ctx)
assert list(r.payload["series"]) == ["flow", "wgan"]
# both rollouts share the same fixture data here, so their matrices (and
# the shared label set) must be identical.
assert r.payload["series"]["flow"] == r.payload["series"]["wgan"]
+143 -30
View File
@@ -1,4 +1,4 @@
"""Tests for the rollout-YAML → run-directory flow, compute, and submit."""
"""Tests for the rollout-YAML(s) → run-directory flow, compute, and submit."""
from __future__ import annotations
@@ -17,6 +17,7 @@ from giant.analysis import (
compute_reduced,
derive_run_dir,
load_rollout_yaml,
load_rollout_yamls,
merge_one,
prep,
write_submit,
@@ -28,13 +29,17 @@ from giant.constants import PREDICT_COORD_METADATA_KEY, ROLLOUT_COORD_VALUE
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
def _write_rollout(path: Path) -> None:
tbl = _rollout_frame().collect().to_arrow()
tbl = tbl.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE})
pq.write_table(tbl, path)
def _write_inputs(tmp_path: Path) -> Path:
"""Materialize rollout+reference parquet and a rollout YAML; return the YAML path."""
rollout = tmp_path / "rollout.parquet"
reference = tmp_path / "reference.parquet"
tbl = _rollout_frame().collect().to_arrow()
tbl = tbl.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE})
pq.write_table(tbl, rollout)
_write_rollout(rollout)
_reference_frame().collect().write_parquet(reference)
yaml_path = tmp_path / "run.yaml"
@@ -54,6 +59,33 @@ def _write_inputs(tmp_path: Path) -> Path:
return yaml_path
def _write_two_inputs(tmp_path: Path) -> tuple[Path, Path]:
"""Two rollout YAMLs (distinct output files) sharing one reference file."""
reference = tmp_path / "reference.parquet"
_reference_frame().collect().write_parquet(reference)
paths = []
for tag, pred_id in (("a", "aaaa1111ef"), ("b", "bbbb2222ef")):
rollout = tmp_path / f"rollout_{tag}.parquet"
_write_rollout(rollout)
yaml_path = tmp_path / f"run_{tag}.yaml"
yaml_path.write_text(
yaml.safe_dump(
{
"prediction_id": pred_id,
"output": str(rollout),
"dataset": str(reference),
"checkpoint": f"/ckpt/{tag}.pt",
"kind": "rollout",
"energy_cutoff": 0.1,
"steps": 10,
}
)
)
paths.append(yaml_path)
return paths[0], paths[1]
def _fake_venv(repo_dir: Path) -> None:
"""Stand in for a `uv sync`'d venv: write_submit checks `.venv/bin/giant` exists."""
giant = repo_dir / ".venv" / "bin" / "giant"
@@ -62,12 +94,13 @@ def _fake_venv(repo_dir: Path) -> None:
giant.chmod(0o755)
def _prep(rollout_yaml: Path, run_dir: str | Path | None = None, chunks: int = 1) -> Path:
def _prep(rollout_yamls, run_dir: str | Path | None = None, chunks: int = 1, labels=None) -> Path:
"""``prep`` with small test-sized context bins/sampling."""
return prep(
rollout_yaml,
rollout_yamls,
run_dir,
n_chunks=chunks,
labels=labels,
n_energy_bins=2,
n_marginal_bins=8,
top_k_pdg=3,
@@ -82,39 +115,108 @@ def test_load_rollout_yaml_requires_paths(tmp_path: Path):
load_rollout_yaml(bad)
def test_load_rollout_yamls_single_defaults_to_rollout_name(tmp_path: Path):
yaml_path = _write_inputs(tmp_path)
loaded, reference = load_rollout_yamls([yaml_path])
assert [lr.name for lr in loaded] == ["rollout"]
assert reference.endswith("reference.parquet")
def test_load_rollout_yamls_multi_defaults_to_stem(tmp_path: Path):
a, b = _write_two_inputs(tmp_path)
loaded, _ = load_rollout_yamls([a, b])
assert [lr.name for lr in loaded] == ["run_a", "run_b"]
def test_load_rollout_yamls_explicit_labels(tmp_path: Path):
a, b = _write_two_inputs(tmp_path)
loaded, _ = load_rollout_yamls([a, b], labels=["flow", "wgan"])
assert [lr.name for lr in loaded] == ["flow", "wgan"]
def test_load_rollout_yamls_label_count_mismatch(tmp_path: Path):
a, b = _write_two_inputs(tmp_path)
with pytest.raises(ValueError, match="--label"):
load_rollout_yamls([a, b], labels=["only-one"])
def test_load_rollout_yamls_rejects_duplicate_names(tmp_path: Path):
a, b = _write_two_inputs(tmp_path)
with pytest.raises(ValueError, match="collide"):
load_rollout_yamls([a, b], labels=["same", "same"])
def test_load_rollout_yamls_rejects_mismatched_reference(tmp_path: Path):
a, _ = _write_two_inputs(tmp_path)
other_ref = tmp_path / "other_reference.parquet"
_reference_frame().collect().write_parquet(other_ref)
c = tmp_path / "run_c.yaml"
c.write_text(
yaml.safe_dump(
{"prediction_id": "cccc3333ef", "output": str(tmp_path / "rollout_c.parquet"), "dataset": str(other_ref)}
)
)
_write_rollout(tmp_path / "rollout_c.parquet")
with pytest.raises(ValueError, match="same reference"):
load_rollout_yamls([a, c])
def test_derive_run_dir_next_to_rollout():
y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"}
assert derive_run_dir(y) == Path("/data/analysis_abcd1234")
assert derive_run_dir(y, "/somewhere") == Path("/somewhere")
assert derive_run_dir([y]) == Path("/data/analysis_abcd1234")
assert derive_run_dir([y], "/somewhere") == Path("/somewhere")
def test_derive_run_dir_default_base():
y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"}
assert derive_run_dir(y, default_base="/work/lbogner/giant2/analysis_runs") == Path(
assert derive_run_dir([y], default_base="/work/lbogner/giant2/analysis_runs") == Path(
"/work/lbogner/giant2/analysis_runs/analysis_abcd1234"
)
# an explicit run_dir still wins over default_base
assert derive_run_dir(y, "/somewhere", default_base="/other") == Path("/somewhere")
assert derive_run_dir([y], "/somewhere", default_base="/other") == Path("/somewhere")
def test_derive_run_dir_multi_rollout_joins_tags():
ys = [{"output": f"/data/roll_{i}.parquet", "prediction_id": f"tag{i}xxxx", "dataset": "d"} for i in range(2)]
assert derive_run_dir(ys, default_base="/base") == Path("/base/analysis_tag0xxxx-tag1xxxx")
def test_derive_run_dir_many_rollouts_truncates_with_plus_count():
ys = [{"output": f"/data/roll_{i}.parquet", "prediction_id": f"tag{i}xxxx", "dataset": "d"} for i in range(5)]
run_dir = derive_run_dir(ys, default_base="/base")
assert run_dir == Path("/base/analysis_tag0xxxx-tag1xxxx-tag2xxxx-plus2")
def test_prep_lays_out_run_dir(tmp_path: Path):
yaml_path = _write_inputs(tmp_path)
run_dir = _prep(yaml_path)
run_dir = _prep([yaml_path])
assert run_dir == tmp_path / "analysis_abcd1234"
assert (run_dir / "shared.json").exists()
ctx = Context.load(run_dir / "shared.json")
assert set(ctx.var_ranges) == {"step_length", "edep", "delta_e", "post_E"}
meta = RunMeta.load(run_dir / "run_meta.json")
assert meta.reference.endswith("reference.parquet")
assert meta.plot_meta["checkpoint"] == "/ckpt/best.pt"
assert [ro["name"] for ro in meta.rollouts] == ["rollout"]
assert meta.rollouts[0]["plot_meta"]["checkpoint"] == "/ckpt/best.pt"
assert "best.pt" in meta.title
assert meta.n_chunks == 1
assert meta.rows_per_chunk == [meta.total_rows] # single chunk holds everything
assert meta.total_rows == 8 # 5 rollout rows + 3 reference rows
def test_prep_multi_rollout_lays_out_run_dir(tmp_path: Path):
a, b = _write_two_inputs(tmp_path)
run_dir = _prep([a, b], labels=["flow", "wgan"])
meta = RunMeta.load(run_dir / "run_meta.json")
assert [ro["name"] for ro in meta.rollouts] == ["flow", "wgan"]
assert meta.rollouts[0]["plot_meta"]["checkpoint"] == "/ckpt/a.pt"
assert meta.rollouts[1]["plot_meta"]["checkpoint"] == "/ckpt/b.pt"
# 5 rows from each rollout + 3 from the shared reference
assert meta.total_rows == 13
def test_prep_splits_rows_per_chunk(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path), chunks=2)
run_dir = _prep([_write_inputs(tmp_path)], chunks=2)
meta = RunMeta.load(run_dir / "run_meta.json")
assert len(meta.rows_per_chunk) == 2
assert sum(meta.rows_per_chunk) == meta.total_rows == 8
@@ -125,7 +227,7 @@ def test_reprep_clears_stale_partials_from_a_different_chunk_count(tmp_path: Pat
partials on disk for merge_one to silently merge against the new
context (they'd be keyed/sized for the old n_chunks)."""
yaml_path = _write_inputs(tmp_path)
run_dir = _prep(yaml_path, chunks=2)
run_dir = _prep([yaml_path], chunks=2)
compute_one("marginal_edep", run_dir, chunk_index=0)
compute_one("marginal_edep", run_dir, chunk_index=1)
stale = run_dir / "reduced_partial" / "marginal_edep__0.json"
@@ -133,7 +235,7 @@ def test_reprep_clears_stale_partials_from_a_different_chunk_count(tmp_path: Pat
(run_dir / "reduced").mkdir(exist_ok=True)
(run_dir / "reduced" / "marginal_edep.json").write_text("{}")
_prep(yaml_path, run_dir, chunks=1)
_prep([yaml_path], run_dir, chunks=1)
assert not stale.exists()
assert not (run_dir / "reduced" / "marginal_edep.json").exists()
@@ -141,20 +243,22 @@ def test_reprep_clears_stale_partials_from_a_different_chunk_count(tmp_path: Pat
def test_compute_one_from_run_dir(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path))
run_dir = _prep([_write_inputs(tmp_path)])
out = compute_one("marginal_edep", run_dir)
assert out == run_dir / "reduced_partial" / "marginal_edep__0.json"
partial = Partial.load(out)
assert partial.id == "marginal_edep" and partial.chunk == 0
assert "r" in partial.data and "t" in partial.data
assert list(partial.data["r"]) == ["rollout"]
def test_compute_reduced_explicit_paths(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path))
run_dir = _prep([_write_inputs(tmp_path)])
meta = RunMeta.load(run_dir / "run_meta.json")
rollouts = [{"name": ro["name"], "path": ro["path"]} for ro in meta.rollouts]
out = compute_reduced(
"marginal_step_length",
meta.rollout,
rollouts,
meta.reference,
run_dir / "shared.json",
tmp_path / "r.json",
@@ -163,17 +267,17 @@ def test_compute_reduced_explicit_paths(tmp_path: Path):
def test_merge_one_produces_reduced(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path))
run_dir = _prep([_write_inputs(tmp_path)])
compute_one("marginal_edep", run_dir)
out = merge_one("marginal_edep", run_dir)
assert out == run_dir / "reduced" / "marginal_edep.json"
reduced = Reduced.load(out)
assert reduced.id == "marginal_edep"
assert len(reduced.payload["rollout"]) == len(reduced.payload["edges"]) - 1
assert len(reduced.payload["series"]["rollout"]) == len(reduced.payload["edges"]) - 1
def test_merge_one_fails_loudly_on_missing_chunk(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path), chunks=2)
run_dir = _prep([_write_inputs(tmp_path)], chunks=2)
compute_one("marginal_edep", run_dir, chunk_index=0) # chunk 1 never computed
with pytest.raises(FileNotFoundError, match="missing chunk"):
merge_one("marginal_edep", run_dir)
@@ -182,11 +286,11 @@ def test_merge_one_fails_loudly_on_missing_chunk(tmp_path: Path):
def test_chunked_compute_and_merge_matches_unchunked(tmp_path: Path):
(tmp_path / "a").mkdir()
(tmp_path / "b").mkdir()
unchunked_dir = _prep(_write_inputs(tmp_path / "a"))
unchunked_dir = _prep([_write_inputs(tmp_path / "a")])
compute_one("marginal_step_length", unchunked_dir)
unchunked = Reduced.load(merge_one("marginal_step_length", unchunked_dir))
chunked_dir = _prep(_write_inputs(tmp_path / "b"), chunks=2)
chunked_dir = _prep([_write_inputs(tmp_path / "b")], chunks=2)
for k in range(2):
compute_one("marginal_step_length", chunked_dir, chunk_index=k)
chunked = Reduced.load(merge_one("marginal_step_length", chunked_dir))
@@ -194,14 +298,23 @@ def test_chunked_compute_and_merge_matches_unchunked(tmp_path: Path):
assert chunked.payload == unchunked.payload
def test_two_rollout_compute_and_merge_produces_both_series(tmp_path: Path):
a, b = _write_two_inputs(tmp_path)
run_dir = _prep([a, b], labels=["flow", "wgan"])
compute_one("marginal_edep", run_dir)
reduced = Reduced.load(merge_one("marginal_edep", run_dir))
assert list(reduced.payload["series"]) == ["flow", "wgan"]
assert "reference" in reduced.payload
def test_compute_reduced_rejects_out_of_range_chunk(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path)) # n_chunks=1 (default)
run_dir = _prep([_write_inputs(tmp_path)]) # n_chunks=1 (default)
with pytest.raises(ValueError, match="out of range"):
compute_one("marginal_edep", run_dir, chunk_index=1)
def test_write_submit_description(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path))
run_dir = _prep([_write_inputs(tmp_path)])
_fake_venv(tmp_path)
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path)
txt = write_submit(cfg).read_text()
@@ -223,7 +336,7 @@ def test_write_submit_description(tmp_path: Path):
def test_write_submit_requires_synced_venv(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
run_dir = _prep(_write_inputs(tmp_path))
run_dir = _prep([_write_inputs(tmp_path)])
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path)
# No `giant` next to the (fake) active interpreter, so this falls through
# to repo_dir/.venv/bin/giant, which _write_inputs/_prep also didn't create.
@@ -233,7 +346,7 @@ def test_write_submit_requires_synced_venv(tmp_path: Path, monkeypatch: pytest.M
def test_write_submit_remote_flag(tmp_path: Path):
run_dir = _prep(_write_inputs(tmp_path))
run_dir = _prep([_write_inputs(tmp_path)])
_fake_venv(tmp_path)
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, remote=True)
txt = write_submit(cfg).read_text()
@@ -243,7 +356,7 @@ def test_write_submit_remote_flag(tmp_path: Path):
def test_write_submit_chunks_respect_chunkable(tmp_path: Path):
assert get_spec("router_gating").chunkable is False
run_dir = _prep(_write_inputs(tmp_path), chunks=4)
run_dir = _prep([_write_inputs(tmp_path)], chunks=4)
_fake_venv(tmp_path)
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, n_chunks=4)
write_submit(cfg)
@@ -260,7 +373,7 @@ def test_write_submit_rejects_n_chunks_mismatch_with_run_meta(tmp_path: Path):
with — RunMeta.rows_per_chunk is sized to the prepped value, so a
mismatch would otherwise surface as a confusing IndexError deep inside
_job_walltimes instead of a clear error here."""
run_dir = _prep(_write_inputs(tmp_path), chunks=2)
run_dir = _prep([_write_inputs(tmp_path)], chunks=2)
_fake_venv(tmp_path)
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, n_chunks=4)
with pytest.raises(ValueError, match="n_chunks"):
@@ -282,7 +395,7 @@ def test_write_submit_walltime_grows_with_chunk_rows(tmp_path: Path):
"""A chunked run's later job walltimes track that chunk's row count."""
from giant.analysis.runtime_estimate import estimate_runtime_s
run_dir = _prep(_write_inputs(tmp_path), chunks=2)
run_dir = _prep([_write_inputs(tmp_path)], chunks=2)
meta = RunMeta.load(run_dir / "run_meta.json")
_fake_venv(tmp_path)
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, n_chunks=2)
+136 -43
View File
@@ -25,41 +25,96 @@ def test_render_router_diagnostics_and_edge_cases(tmp_path: Path):
reduced = [
Reduced(
"rg",
"router",
"model",
"router_gating",
"Router gating",
"pre-step energy [MeV]",
{
"n_experts": 2,
"log_x": True,
"router_type": "energy",
"rollout": {
"centers": [1.0, 10.0, 100.0],
"means": [[0.6, 0.4], [0.5, 0.5], [0.4, 0.6]],
},
"reference": {
"centers": [1.0, 10.0, 100.0],
"means": [[0.55, 0.45], [0.5, 0.5], [0.45, 0.55]],
"series": {
"flow": {
"n_experts": 2,
"router_type": "energy",
"rollout": {
"centers": [1.0, 10.0, 100.0],
"means": [[0.6, 0.4], [0.5, 0.5], [0.4, 0.6]],
},
"reference": {
"centers": [1.0, 10.0, 100.0],
"means": [[0.55, 0.45], [0.5, 0.5], [0.45, 0.55]],
},
},
"wgan": {
"n_experts": 2,
"router_type": "energy",
"rollout": {"centers": [1.0], "means": [[0.5, 0.5]]},
"reference": {"centers": [1.0], "means": [[0.5, 0.5]]},
},
},
},
),
Reduced(
"rs",
"router",
"model",
"router_share",
"Router share",
"species",
{
"categories": ["e-", "gamma"],
"n_experts": 2,
"router_type": "energy",
"rollout": {"e-": [0.7, 0.3], "gamma": [0.2, 0.8]},
"reference": {"e-": [0.6, 0.4], "gamma": [0.3, 0.7]},
"series": {
"flow": {
"categories": ["e-", "gamma"],
"n_experts": 2,
"router_type": "energy",
"rollout": {"e-": [0.7, 0.3], "gamma": [0.2, 0.8]},
"reference": {"e-": [0.6, 0.4], "gamma": [0.3, 0.7]},
},
},
},
),
Reduced(
"rp",
"model",
"router_share",
"Router share by process (reference-only)",
"process",
{
"series": {
"flow": {
"categories": ["compt", "phot"],
"n_experts": 2,
"router_type": "energy",
"reference": {"compt": [0.4, 0.6], "phot": [0.9, 0.1]},
},
},
},
),
Reduced(
"rz",
"model",
"router_specialization",
"Router specialization",
"pre-step energy [MeV]",
{
"log_x": True,
"series": {
"flow": {
"n_experts": 2,
"chance_level": 0.5,
"rollout": {"centers": [1.0, 10.0], "score": [0.6, 0.7]},
"reference": {"centers": [1.0, 10.0], "score": [0.55, 0.65]},
},
"wgan": {
"n_experts": 4,
"chance_level": 0.25,
"rollout": {"centers": [1.0, 10.0], "score": [0.3, 0.4]},
"reference": {"centers": [], "score": []},
},
},
},
),
Reduced(
"ru",
"router",
"model",
"unavailable",
"Router unavailable",
"x",
@@ -73,7 +128,10 @@ def test_render_router_diagnostics_and_edge_cases(tmp_path: Path):
"x",
{
"edges": [0, 1, 2],
"groups": {lbl: {"rollout": [1, 2], "reference": [2, 1]} for lbl in ("a", "b", "c", "d")},
"groups": {
lbl: {"series": {"flow": [1, 2], "wgan": [2, 1]}, "reference": [2, 1]}
for lbl in ("a", "b", "c", "d")
},
"log_y": True,
},
),
@@ -83,7 +141,22 @@ def test_render_router_diagnostics_and_edge_cases(tmp_path: Path):
"single_hist",
"Single (log-x)",
"x",
{"edges": [1, 10, 100], "rollout": [5, 1], "log_x": True, "log_y": True},
{"edges": [1, 10, 100], "series": {"flow": [5, 1], "wgan": [3, 2]}, "log_x": True, "log_y": True},
),
Reduced(
"hm",
"quality",
"heatmap",
"Distance summary (2 rollouts)",
"grouping axis",
{
"series": {"flow": [[0.1, 0.2], [0.3, 0.4]], "wgan": [[0.5, 0.6], [0.7, 0.8]]},
"row_labels": ["step_length", "edep"],
"col_labels": ["overall", "energy"],
"cbar_label": "KS statistic",
"vmin": 0.0,
"vmax": 1.0,
},
),
]
try:
@@ -117,7 +190,7 @@ def test_render_all_run_gallery_invokes_subprocess(tmp_path: Path, monkeypatch):
"single_hist",
"Single",
"x",
{"edges": [0, 1, 2], "rollout": [5, 1]},
{"edges": [0, 1, 2], "series": {"rollout": [5, 1]}},
)
]
for r in reduced:
@@ -142,15 +215,14 @@ def test_render_run_glues_condor_run_meta_into_render_all(tmp_path: Path, monkey
merge_calls = []
monkeypatch.setattr(condor_mod, "merge_all", lambda rd: merge_calls.append(Path(rd)))
meta = condor_mod.RunMeta(
rollout="rollout.parquet",
rollouts=[{"name": "rollout", "path": "rollout.parquet", "plot_meta": {"checkpoint": "ckpt/best.pt"}}],
reference="reference.parquet",
run_dir=str(run_dir),
title="my-run",
plot_meta={"checkpoint": "ckpt/best.pt"},
)
monkeypatch.setattr(condor_mod.RunMeta, "load", classmethod(lambda cls, p: meta))
Reduced("s", "species", "single_hist", "Single", "x", {"edges": [0, 1], "rollout": [1]}).save(
Reduced("s", "species", "single_hist", "Single", "x", {"edges": [0, 1], "series": {"rollout": [1]}}).save(
run_dir / "reduced" / "s.json"
)
@@ -177,7 +249,7 @@ def test_render_one_of_each_kind(tmp_path: Path):
"x",
{
"edges": [0, 1, 2, 3],
"rollout": [1, 2, 3],
"series": {"flow": [1, 2, 3], "wgan": [2, 2, 2]},
"reference": [3, 2, 1],
"log_y": False,
},
@@ -190,7 +262,7 @@ def test_render_one_of_each_kind(tmp_path: Path):
"x",
{
"edges": [0, 1, 2],
"groups": {"a": {"rollout": [1, 2], "reference": [2, 1]}},
"groups": {"a": {"series": {"flow": [1, 2]}, "reference": [2, 1]}},
"log_y": False,
},
),
@@ -202,10 +274,8 @@ def test_render_one_of_each_kind(tmp_path: Path):
"depth",
{
"edges": [0, 1, 2],
"rollout_mean": [1, 2],
"rollout_std": [0.1, 0.2],
"reference_mean": [1.1, 1.9],
"reference_std": [0.1, 0.1],
"series": {"flow": {"mean": [1, 2], "std": [0.1, 0.2]}},
"reference": {"mean": [1.1, 1.9], "std": [0.1, 0.1]},
"ylabel": "e",
},
),
@@ -217,7 +287,7 @@ def test_render_one_of_each_kind(tmp_path: Path):
"species",
{
"labels": ["e-", "gamma"],
"rollout": [0.6, 0.4],
"series": {"flow": [0.6, 0.4], "wgan": [0.55, 0.45]},
"reference": [0.5, 0.5],
"ylabel": "frac",
},
@@ -228,7 +298,20 @@ def test_render_one_of_each_kind(tmp_path: Path):
"single_hist",
"Single",
"x",
{"edges": [0, 1, 2], "rollout": [5, 1], "log_y": True},
{"edges": [0, 1, 2], "series": {"flow": [5, 1]}, "log_y": True},
),
Reduced(
"hm1",
"secondaries",
"heatmap",
"Confusion (single rollout)",
"predicted",
{
"series": {"flow": [[1, 0], [0, 1]]},
"row_labels": ["0", "1+"],
"col_labels": ["0", "1+"],
"cbar_label": "count",
},
),
]
try:
@@ -276,8 +359,8 @@ def test_figure_params_v2_basics_and_router_and_epoch():
},
"conditioning": {"particle": {"type": "physical"}},
}
run_meta = {"training_epoch": 12, "best_val_loss": 0.123456, "steps": 10}
params = render_mod._figure_params(run_meta | {"model_config": mc})
meta = {"training_epoch": 12, "best_val_loss": 0.123456, "steps": 10, "model_config": mc}
params = render_mod._figure_params({"rollouts": {"rollout": meta}})
assert params == {
"hidden_dim": 256,
"n_res_blocks": 4,
@@ -297,8 +380,8 @@ def test_figure_params_v2_wgan_reports_noise_dim_not_steps():
"wgan": {"noise_dim": 32},
},
}
run_meta = {"model_config": mc, "steps": 10}
params = render_mod._figure_params(run_meta)
meta = {"model_config": mc, "steps": 10}
params = render_mod._figure_params({"rollouts": {"rollout": meta}})
assert params["mode"] == "wgan"
assert params["noise_dim"] == 32
assert "steps" not in params
@@ -309,18 +392,18 @@ def test_figure_params_v2_reports_mode_s2_only_when_it_differs():
"stage1_model": {"generator": "flow"},
"stage2_model": {"generator": "flow"},
}
assert "mode_s2" not in render_mod._figure_params({"model_config": same})
assert "mode_s2" not in render_mod._figure_params({"rollouts": {"rollout": {"model_config": same}}})
mixed = {
"stage1_model": {"generator": "flow"},
"stage2_model": {"generator": "wgan"},
}
params = render_mod._figure_params({"model_config": mixed})
params = render_mod._figure_params({"rollouts": {"rollout": {"model_config": mixed}}})
assert params["mode_s2"] == "wgan"
def test_figure_params_old_shape_basics():
run_meta = {
meta = {
"model_config": {
"hidden_dim": 128,
"n_blocks": 3,
@@ -332,7 +415,7 @@ def test_figure_params_old_shape_basics():
"best_val_loss": 0.5,
"steps": 20,
}
params = render_mod._figure_params(run_meta)
params = render_mod._figure_params({"rollouts": {"rollout": meta}})
assert params == {
"hidden_dim": 128,
"n_blocks": 3,
@@ -346,20 +429,30 @@ def test_figure_params_old_shape_basics():
def test_figure_params_old_shape_wgan_reports_noise_dim_not_steps():
run_meta = {
meta = {
"model_config": {"mode": "wgan", "noise_dim": 16},
"steps": 20,
}
params = render_mod._figure_params(run_meta)
params = render_mod._figure_params({"rollouts": {"rollout": meta}})
assert params["noise_dim"] == 16
assert "steps" not in params
def test_figure_params_multi_rollout_names_the_series():
run_meta = {"rollouts": {"flow": {"model_config": {"mode": "flow"}}, "wgan": {"model_config": {"mode": "wgan"}}}}
assert render_mod._figure_params(run_meta) == {"rollouts": "flow, wgan"}
def test_figure_params_empty_rollouts_is_empty():
assert render_mod._figure_params({}) == {}
assert render_mod._figure_params({"rollouts": {}}) == {}
def test_plot_metadata_includes_note_and_run_meta_parameters():
r = Reduced("u", "router", "unavailable", "Unavailable", "x", {"note": "no router data"})
meta = render_mod._plot_metadata(r, {"title": "run-1", "checkpoint": "ckpt.pt"})
meta = render_mod._plot_metadata(r, {"title": "run-1", "reference": "ref.parquet", "rollouts": {"rollout": {}}})
assert meta["note"] == "no router data"
assert meta["parameters"] == {"checkpoint": "ckpt.pt"}
assert meta["parameters"] == {"reference": "ref.parquet", "rollouts": {"rollout": {}}}
assert "title" not in meta["parameters"]
+52 -12
View File
@@ -10,7 +10,9 @@ from giant.analysis.router_gating import (
compute_router_gating,
compute_router_share_by_pdg,
compute_router_share_by_process,
compute_router_specialization,
)
from giant.analysis.sources import RolloutSide
from giant.data.transforms import Normalizer
from giant.model.network import build_models
@@ -34,7 +36,7 @@ def _model_cfg() -> dict:
}
def _write_checkpoint(tmp_path) -> str:
def _write_checkpoint(tmp_path, name: str = "ckpt.pt") -> str:
cfg = _model_cfg()
stage1 = build_models(cfg)["stage1"]
assert stage1 is not None
@@ -48,7 +50,7 @@ def _write_checkpoint(tmp_path) -> str:
"mat_map": _MAT_MAP,
"normalizer": {"cond": norm.to_dict()},
}
path = tmp_path / "ckpt.pt"
path = tmp_path / name
torch.save(ckpt, path)
return str(path)
@@ -86,42 +88,80 @@ def _steps_frame(process: bool = False) -> pl.LazyFrame:
return pl.DataFrame(data).lazy()
def _side(checkpoint: str | None, lf: pl.LazyFrame) -> RolloutSide:
return RolloutSide(all=lf, phys=lf, checkpoint=checkpoint)
def test_compute_router_gating_shapes(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
lf = _steps_frame()
r = compute_router_gating(checkpoint, lf, lf)
r = compute_router_gating({"rollout": _side(checkpoint, lf)}, lf)
assert r.kind == "router_gating"
assert r.payload["n_experts"] == 2
assert list(r.payload["series"]) == ["rollout"]
entry = r.payload["series"]["rollout"]
assert entry["n_experts"] == 2
for side in ("rollout", "reference"):
means = r.payload[side]["means"]
means = entry[side]["means"]
assert means, f"{side} produced no bins"
assert all(abs(sum(row) - 1.0) < 1e-5 for row in means)
def test_compute_router_gating_missing_checkpoint_is_unavailable():
lf = _steps_frame()
r = compute_router_gating(None, lf, lf)
r = compute_router_gating({"rollout": _side(None, lf)}, lf)
assert r.kind == "unavailable"
assert "note" in r.payload
assert r.title
def test_compute_router_gating_two_rollouts_only_moe_ones_included(tmp_path):
lf = _steps_frame()
ckpt = _write_checkpoint(tmp_path)
rollouts = {"flow": _side(None, lf), "moe": _side(ckpt, lf)}
r = compute_router_gating(rollouts, lf)
assert list(r.payload["series"]) == ["moe"]
def test_compute_router_specialization_two_rollouts(tmp_path):
lf = _steps_frame()
ckpt_a = _write_checkpoint(tmp_path, "a.pt")
ckpt_b = _write_checkpoint(tmp_path, "b.pt")
rollouts = {"a": _side(ckpt_a, lf), "b": _side(ckpt_b, lf)}
r = compute_router_specialization(rollouts, lf)
assert r.kind == "router_specialization"
assert list(r.payload["series"]) == ["a", "b"]
for entry in r.payload["series"].values():
assert entry["chance_level"] == 0.5
assert len(entry["rollout"]["centers"]) == len(entry["rollout"]["score"])
def test_compute_router_share_by_pdg(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
lf = _steps_frame()
r = compute_router_share_by_pdg(checkpoint, lf, lf, top_pdgs=[11, 22])
r = compute_router_share_by_pdg({"rollout": _side(checkpoint, lf)}, lf, top_pdgs=[11, 22])
assert r.kind == "router_share"
entry = r.payload["series"]["rollout"]
for side in ("rollout", "reference"):
assert set(r.payload[side]) == {"e-", "gamma"}
for shares in r.payload[side].values():
assert set(entry[side]) == {"e-", "gamma"}
for shares in entry[side].values():
assert abs(sum(shares) - 1.0) < 1e-5
def test_compute_router_share_by_process(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
lf = _steps_frame(process=True)
r = compute_router_share_by_process(checkpoint, lf)
r = compute_router_share_by_process({"rollout": _side(checkpoint, lf)}, lf)
assert r.kind == "router_share"
assert set(r.payload["categories"]) <= {"eIoni", "compt"}
for shares in r.payload["reference"].values():
entry = r.payload["series"]["rollout"]
assert set(entry["categories"]) <= {"eIoni", "compt"}
for shares in entry["reference"].values():
assert abs(sum(shares) - 1.0) < 1e-5
def test_no_moe_rollouts_are_unavailable(tmp_path):
lf = _steps_frame()
rollouts = {"flow": _side(None, lf), "wgan": _side(None, lf)}
assert compute_router_gating(rollouts, lf).kind == "unavailable"
assert compute_router_share_by_pdg(rollouts, lf, top_pdgs=[11, 22]).kind == "unavailable"
assert compute_router_share_by_process(rollouts, lf).kind == "unavailable"
assert compute_router_specialization(rollouts, lf).kind == "unavailable"
+30 -6
View File
@@ -3,6 +3,9 @@
from __future__ import annotations
import polars as pl
from giant.analysis.sources import RolloutSide
from giant.analysis.type_embedding_distance import compute_type_embedding_l1_distance
@@ -18,26 +21,47 @@ def _summary(n=100):
}
def _side(l1_dist: dict | None) -> RolloutSide:
empty = pl.LazyFrame()
return RolloutSide(all=empty, phys=empty, type_embedding_l1_dist=l1_dist)
def test_none_is_unavailable():
r = compute_type_embedding_l1_distance(None)
r = compute_type_embedding_l1_distance({"rollout": _side(None)})
assert r.kind == "unavailable"
assert r.id == "type_embedding_l1_distance"
assert r.payload["note"]
def test_summary_produces_single_hist():
r = compute_type_embedding_l1_distance(_summary())
r = compute_type_embedding_l1_distance({"rollout": _side(_summary())})
assert r.kind == "single_hist"
assert r.id == "type_embedding_l1_distance"
assert r.payload["edges"] == [0.0, 1.0, 2.0, 3.0]
assert r.payload["rollout"] == [30, 40, 30]
assert r.payload["series"]["rollout"] == [30, 40, 30]
assert r.payload["log_x"] is True
assert r.payload["log_y"] is True
assert "n=100" in r.payload["note"]
def test_single_hist_payload_shape_matches_render_contract():
"""_render_single (giant.analysis.render) requires len(rollout) ==
"""_render_single (giant.analysis.render) requires each series' length ==
len(edges) - 1."""
r = compute_type_embedding_l1_distance(_summary())
assert len(r.payload["rollout"]) == len(r.payload["edges"]) - 1
r = compute_type_embedding_l1_distance({"rollout": _side(_summary())})
assert len(r.payload["series"]["rollout"]) == len(r.payload["edges"]) - 1
def test_two_rollouts_both_populated():
r = compute_type_embedding_l1_distance({"flow": _side(_summary(50)), "wgan": _side(_summary(80))})
assert list(r.payload["series"]) == ["flow", "wgan"]
assert "n=50" in r.payload["note"] and "n=80" in r.payload["note"]
def test_one_of_two_rollouts_populated_only_that_one_appears():
r = compute_type_embedding_l1_distance({"flow": _side(None), "wgan": _side(_summary())})
assert list(r.payload["series"]) == ["wgan"]
def test_none_populated_across_rollouts_is_unavailable():
r = compute_type_embedding_l1_distance({"flow": _side(None), "wgan": _side(None)})
assert r.kind == "unavailable"