93b19911f8
CI / Format (ruff format) (push) Successful in 36s
CI / Lint (ruff check) (push) Successful in 38s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Failing after 45s
CI / Lint (ruff check) (pull_request) Successful in 41s
CI / Tests (push) Has been skipped
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) Failing after 37s
CI / Tests (pull_request) Has been skipped
- giant/sample.py: fix every sampler's call convention against
Stage1Model/Stage2OneShot's actual forward signatures (was still
calling model(x, t, cond_cont, cond_cat) positionally); add
sample_secondaries_ar (free-running AR loop, unsnapped history feature)
and sample_stage1/sample_stage2/resolve_n_sec dispatch helpers that read
each stage's generator_kind/decoder off the model instance itself.
- giant/particles.py: decode_topn_class (argmax + other_policy) and
decode_embedding_nearest (L1-snap + distance) turn a secondary's
"onehot"/"embedding" type prediction into a concrete PDG.
- giant/rollout.py: decode_secondary_identity routes all three
particle_type.target values to real mass/charge; per-stage generator
dispatch (drops the single shared `mode` string, adds ddpm support);
L1DistCollector accumulates the §11.3 embedding-distance diagnostic.
- giant/cli.py: drop the onehot/embedding-target rejection gate (narrowed
to the still-unimplemented conditioning.particle/material.type=onehot
axis); fix the dead model_cfg.get("mode") bug in predict/rollout.
- giant/analysis/: new type_embedding_l1_distance PlotSpec, wired through
the rollout YAML sidecar (no live-model call needed, unlike
router_gating -- the histogram is already pre-aggregated at rollout
time).
- Un-xfail every test that was blocked on this step (test_rollout.py,
test_flow.py, test_wgan.py, test_phase2.py, test_router.py,
test_validate.py); add test_sample.py, test_type_embedding_distance.py.
Known follow-up: giant/validate.py still unpacks the training val-batch
as a stale 6-tuple and doesn't use the new per-stage dispatch, so
marginal validation during training degrades gracefully with a warning
rather than working -- not in this step's scope.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
218 lines
7.0 KiB
Python
218 lines
7.0 KiB
Python
import uuid
|
|
|
|
import pytest
|
|
import typer
|
|
import yaml
|
|
|
|
from giant.cli import (
|
|
_CEPH_PREDICTIONS,
|
|
_check_conditioning_onehot_support,
|
|
_resolve_prediction_output,
|
|
_write_prediction_ref,
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# _check_conditioning_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_conditioning_onehot_support_allows_physical():
|
|
_check_conditioning_onehot_support(_nested_model_cfg(), "predict") # no raise
|
|
|
|
|
|
def test_check_conditioning_onehot_support_rejects_onehot_particle_conditioning():
|
|
cfg = _nested_model_cfg(particle_type="onehot")
|
|
with pytest.raises(typer.Exit):
|
|
_check_conditioning_onehot_support(cfg, "predict")
|
|
|
|
|
|
def test_check_conditioning_onehot_support_rejects_onehot_material_conditioning():
|
|
cfg = _nested_model_cfg(material_type="onehot")
|
|
with pytest.raises(typer.Exit):
|
|
_check_conditioning_onehot_support(cfg, "rollout")
|
|
|
|
|
|
def test_check_conditioning_onehot_support_allows_onehot_particle_type_target():
|
|
"""stage2_model.particle_type.target="onehot" is implemented (v0.3.0
|
|
step 6, giant.rollout.decode_secondary_identity) — it's a separate axis
|
|
from conditioning.particle.type, which this guard doesn't gate at all."""
|
|
cfg = _nested_model_cfg(target="onehot")
|
|
_check_conditioning_onehot_support(cfg, "predict") # no raise
|
|
|
|
|
|
def test_check_conditioning_onehot_support_allows_embedding_particle_type_target():
|
|
cfg = _nested_model_cfg(
|
|
particle_type="embedding", material_type="embedding", target="embedding"
|
|
)
|
|
_check_conditioning_onehot_support(cfg, "predict") # no raise
|
|
|
|
|
|
def test_check_conditioning_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, so this must be a silent no-op rather
|
|
than crash on `.get("particle")` against a string."""
|
|
cfg = {"conditioning": "embedding", "mode": "flow"}
|
|
_check_conditioning_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("/")
|