Files
giant/tests/test_cli_predict.py
T
lars 4fc15ecdfc
CI / Lint (ruff check) (push) Successful in 26s
CI / Format (ruff format) (push) Successful in 28s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 28s
CI / Type check (ty) (push) Successful in 31s
CI / Format (ruff format) (pull_request) Successful in 35s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 36s
CI / Tests (pull_request) Successful in 1m41s
CI / Tests (push) Successful in 1m47s
v0.3.0 step 4: type map + particle_type.target = "onehot"/"embedding"
Builds the shared top-N-plus-other PDG/material maps (pooling both primary
and secondary occurrences for PDG, directly targeting the meeting's
species-collapse failure mode) and wires up conditioning.{particle,material}
= "onehot" plus stage2_model.particle_type.target in ("onehot", "embedding")
end-to-end: setup-cache persistence, Stage2OneShot's type_head (flow/ddpm)
vs. folded+ST-Gumbel-relaxed adversarial slice (wgan), and the corresponding
CE/MSE training losses. particle_type.target = "physical" stays byte-for-byte
unchanged, keeping the v0.2 migration shim's bit-identical guarantee intact.
giant predict/rollout fail loudly on a onehot/embedding checkpoint until
full decode support lands in step 6.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-08-06 15:43:48 +02:00

217 lines
6.7 KiB
Python

import uuid
import pytest
import typer
import yaml
from giant.cli import (
_CEPH_PREDICTIONS,
_check_v030_onehot_support,
_resolve_prediction_output,
_write_prediction_ref,
)
# ---------------------------------------------------------------------------
# _check_v030_onehot_support
# ---------------------------------------------------------------------------
def _nested_model_cfg(
particle_type="physical", material_type="physical", target="physical"
):
return {
"conditioning": {
"particle": {"type": particle_type, "emb_dim": 8},
"material": {"type": material_type, "emb_dim": 8},
},
"stage2_model": {"particle_type": {"target": target}},
}
def test_check_v030_onehot_support_allows_physical():
_check_v030_onehot_support(_nested_model_cfg(), "predict") # no raise
def test_check_v030_onehot_support_rejects_onehot_particle_conditioning():
cfg = _nested_model_cfg(particle_type="onehot")
with pytest.raises(typer.Exit):
_check_v030_onehot_support(cfg, "predict")
def test_check_v030_onehot_support_rejects_onehot_material_conditioning():
cfg = _nested_model_cfg(material_type="onehot")
with pytest.raises(typer.Exit):
_check_v030_onehot_support(cfg, "rollout")
def test_check_v030_onehot_support_rejects_onehot_particle_type_target():
cfg = _nested_model_cfg(target="onehot")
with pytest.raises(typer.Exit):
_check_v030_onehot_support(cfg, "predict")
def test_check_v030_onehot_support_rejects_embedding_particle_type_target():
cfg = _nested_model_cfg(
particle_type="embedding", material_type="embedding", target="embedding"
)
with pytest.raises(typer.Exit):
_check_v030_onehot_support(cfg, "predict")
def test_check_v030_onehot_support_is_noop_for_v02_flat_model_config():
"""A v0.2 checkpoint's flat model_config has conditioning as a plain
string, not a dict — never onehot/embedding-target, so this must be a
silent no-op rather than crash on `.get("particle")` against a string."""
cfg = {"conditioning": "embedding", "mode": "flow"}
_check_v030_onehot_support(cfg, "predict") # no raise
# ---------------------------------------------------------------------------
# _resolve_prediction_output
# ---------------------------------------------------------------------------
def test_non_ceph_path_goes_to_data_parent(tmp_path):
data = tmp_path / "pools" / "pbwo4" / "full.manifest"
out, dataset_path, pred_uuid = _resolve_prediction_output(data, None)
assert out.parent == data.parent
assert out.name == f"{pred_uuid}.parquet"
assert dataset_path == data.resolve()
def test_ceph_path_goes_to_central_store(tmp_path, monkeypatch):
# Patch resolve() so /ceph/... exists on any machine running the tests.
ceph_data = _CEPH_PREDICTIONS.parent / "pools" / "pbwo4" / "full.manifest"
monkeypatch.setattr(
"giant.cli.Path.resolve",
lambda self: ceph_data if self == ceph_data else self.absolute(),
)
out, _, pred_uuid = _resolve_prediction_output(ceph_data, None)
assert out.parent == _CEPH_PREDICTIONS
assert out.name == f"{pred_uuid}.parquet"
def test_explicit_out_is_used_as_is(tmp_path):
data = tmp_path / "data.parquet"
explicit = tmp_path / "my_output.parquet"
out, _, _ = _resolve_prediction_output(data, explicit)
assert out == explicit
def test_uuid_is_valid(tmp_path):
data = tmp_path / "data.parquet"
_, _, pred_uuid = _resolve_prediction_output(data, None)
parsed = uuid.UUID(pred_uuid)
assert parsed.version == 4
def test_each_call_produces_a_distinct_uuid(tmp_path):
data = tmp_path / "data.parquet"
_, _, uuid1 = _resolve_prediction_output(data, None)
_, _, uuid2 = _resolve_prediction_output(data, None)
assert uuid1 != uuid2
def test_dataset_path_is_resolved(tmp_path):
data = tmp_path / "data.parquet"
_, dataset_path, _ = _resolve_prediction_output(data, None)
assert dataset_path.is_absolute()
# ---------------------------------------------------------------------------
# _write_prediction_ref
# ---------------------------------------------------------------------------
def test_ref_file_created_in_checkpoint_dir(tmp_path):
ckpt_dir = tmp_path / "checkpoints" / "run1"
ckpt_dir.mkdir(parents=True)
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
out = tmp_path / "predictions" / "abc.parquet"
dataset = tmp_path / "pools" / "pbwo4" / "full.manifest"
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
assert ref_path == ckpt_dir / f"{pred_uuid}.yaml"
assert ref_path.exists()
def test_ref_yaml_contains_expected_fields(tmp_path):
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
out = tmp_path / "pred.parquet"
dataset = tmp_path / "full.manifest"
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
data = yaml.safe_load(ref_path.read_text())
assert data["prediction_id"] == pred_uuid
assert data["output"] == str(out)
assert data["dataset"] == str(dataset)
assert data["checkpoint"] == str(checkpoint.resolve())
assert "timestamp" in data
assert "comment" not in data
def test_ref_yaml_includes_comment_when_provided(tmp_path):
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
out = tmp_path / "pred.parquet"
dataset = tmp_path / "full.manifest"
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(
checkpoint, pred_uuid, out, dataset, comment="baseline sweep run 3"
)
data = yaml.safe_load(ref_path.read_text())
assert data["comment"] == "baseline sweep run 3"
def test_ref_timestamp_is_iso_format(tmp_path):
from datetime import datetime
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(
checkpoint, pred_uuid, tmp_path / "p.parquet", tmp_path / "d"
)
data = yaml.safe_load(ref_path.read_text())
# Must parse without error and be timezone-aware (UTC).
ts = datetime.fromisoformat(data["timestamp"])
assert ts.tzinfo is not None
def test_ref_checkpoint_path_is_absolute(tmp_path):
ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir()
checkpoint = ckpt_dir / "best.pt"
checkpoint.touch()
pred_uuid = str(uuid.uuid4())
ref_path = _write_prediction_ref(
checkpoint, pred_uuid, tmp_path / "p.parquet", tmp_path / "d"
)
data = yaml.safe_load(ref_path.read_text())
assert data["checkpoint"].startswith("/")