1115301eb9
ty correctly flagged this as unsound: _CTX's inferred dict[str, int] type doesn't rule out a "run_dir" key, which would silently bind to prep's own run_dir: str | Path | None parameter instead of falling through to **ctx_kwargs. Passing the context kwargs by name in a small test helper removes the ambiguity. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
138 lines
4.6 KiB
Python
138 lines
4.6 KiB
Python
"""Tests for the rollout-YAML → run-directory flow, compute, and submit."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from pathlib import Path
|
|
|
|
import pyarrow.parquet as pq
|
|
import pytest
|
|
import yaml
|
|
|
|
from giant.analysis import (
|
|
RunMeta,
|
|
SubmitConfig,
|
|
catalog_ids,
|
|
compute_one,
|
|
compute_reduced,
|
|
derive_run_dir,
|
|
load_rollout_yaml,
|
|
prep,
|
|
write_submit,
|
|
)
|
|
from giant.analysis.condor import Context
|
|
from giant.analysis.reduced import Reduced
|
|
from giant.constants import PREDICT_COORD_METADATA_KEY, ROLLOUT_COORD_VALUE
|
|
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
|
|
|
|
|
|
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)
|
|
_reference_frame().collect().write_parquet(reference)
|
|
|
|
yaml_path = tmp_path / "run.yaml"
|
|
yaml_path.write_text(
|
|
yaml.safe_dump(
|
|
{
|
|
"prediction_id": "abcd1234ef",
|
|
"output": str(rollout),
|
|
"dataset": str(reference),
|
|
"checkpoint": "/ckpt/best.pt",
|
|
"kind": "rollout",
|
|
"energy_cutoff": 0.1,
|
|
"steps": 10,
|
|
}
|
|
)
|
|
)
|
|
return yaml_path
|
|
|
|
|
|
def _prep(rollout_yaml: Path, run_dir: str | Path | None = None) -> Path:
|
|
"""``prep`` with small test-sized context bins/sampling."""
|
|
return prep(
|
|
rollout_yaml,
|
|
run_dir,
|
|
n_energy_bins=2,
|
|
n_marginal_bins=8,
|
|
top_k_pdg=3,
|
|
sample_rows=1000,
|
|
)
|
|
|
|
|
|
def test_load_rollout_yaml_requires_paths(tmp_path: Path):
|
|
bad = tmp_path / "bad.yaml"
|
|
bad.write_text(yaml.safe_dump({"output": "x.parquet"})) # no dataset
|
|
with pytest.raises(ValueError):
|
|
load_rollout_yaml(bad)
|
|
|
|
|
|
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")
|
|
|
|
|
|
def test_prep_lays_out_run_dir(tmp_path: Path):
|
|
yaml_path = _write_inputs(tmp_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 "best.pt" in meta.title
|
|
|
|
|
|
def test_compute_one_from_run_dir(tmp_path: Path):
|
|
run_dir = _prep(_write_inputs(tmp_path))
|
|
out = compute_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
|
|
|
|
|
|
def test_compute_reduced_explicit_paths(tmp_path: Path):
|
|
run_dir = _prep(_write_inputs(tmp_path))
|
|
meta = RunMeta.load(run_dir / "run_meta.json")
|
|
out = compute_reduced(
|
|
"marginal_step_length",
|
|
meta.rollout,
|
|
meta.reference,
|
|
run_dir / "shared.json",
|
|
tmp_path / "r.json",
|
|
)
|
|
assert Reduced.load(out).id == "marginal_step_length"
|
|
|
|
|
|
def test_write_submit_description(tmp_path: Path):
|
|
run_dir = _prep(_write_inputs(tmp_path))
|
|
cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path)
|
|
txt = write_submit(cfg).read_text()
|
|
assert "universe = docker" in txt
|
|
assert "docker_image = mschnepf/slc7-condocker" in txt
|
|
assert "requirements = TARGET.ProvidesETPResources" in txt
|
|
assert "accounting_group = cms" in txt
|
|
assert "queue plotid from" in txt
|
|
assert (run_dir / "plotids.txt").read_text().split() == catalog_ids()
|
|
wrapper = run_dir / "run_compute.sh"
|
|
assert wrapper.exists() and (wrapper.stat().st_mode & 0o111)
|
|
body = wrapper.read_text()
|
|
assert "giant analyze compute-one --id" in body and "--run-dir" in body
|
|
|
|
|
|
def test_write_submit_remote_flag(tmp_path: Path):
|
|
run_dir = _prep(_write_inputs(tmp_path))
|
|
cfg = SubmitConfig(
|
|
run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, remote=True
|
|
)
|
|
txt = write_submit(cfg).read_text()
|
|
assert "+RemoteJob = True" in txt
|
|
assert "ProvidesETPResources" not in txt
|