"""Tests for giant/sample.py's v0.3.0 stage-model sampling — the AR loop (`sample_secondaries_ar`) and non-"physical" `particle_type.target` coverage for the one-shot samplers.""" import pytest import torch from giant.constants import COND_DIM, CONT_SLOT_DIM, K_MAX, PARTICLE_PHYS_DIM, X_DIM from giant.model.network import ( Stage1Model, Stage2Autoregressive, Stage2OneShot, stage2_trunk_sec_dim, ) from giant.sample import ( sample_flow, sample_secondaries, sample_secondaries_ar, sample_secondaries_wgan, sample_wgan, ) _PHYS_CFG = {"type": "physical", "emb_dim": 8, "n_layers": 1} def _particle_material_cfg(conditioning: str, emb_dim: int) -> tuple[dict, dict]: cfg = {"type": conditioning, "emb_dim": emb_dim, "n_layers": 1} return dict(cfg), dict(cfg) def _cond(B: int, pdg: int = 3, mat: int = 2) -> tuple[torch.Tensor, torch.Tensor]: cond_cont = torch.randn(B, COND_DIM) cond_cat = torch.stack( [torch.randint(0, pdg, (B,)), torch.randint(0, mat, (B,))], dim=1 ) return cond_cont, cond_cat def _conditioning_for(target: str) -> str: # target="embedding" regresses against the conditioning's own embedding # table — only meaningful when the # conditioning axis is itself "embedding". return "embedding" if target == "embedding" else "physical" def _stage2_oneshot( target: str, generator: str, emb_dim: int = 6, pdg: int = 3, mat: int = 2 ) -> Stage2OneShot: particle_cfg, material_cfg = _particle_material_cfg( _conditioning_for(target), emb_dim ) particle_type_cfg = {"target": target} # build_models (giant/model/network.py) computes sec_dim this same way # before constructing Stage2OneShot — its own default (SEC_DIM, the # "physical" width) is only correct for target="physical". sec_dim = stage2_trunk_sec_dim(particle_type_cfg, generator, K_MAX, emb_dim) return Stage2OneShot( pdg_vocab=pdg, mat_vocab=mat, particle_cfg=particle_cfg, material_cfg=material_cfg, hidden_dim=32, n_res_blocks=2, generator=generator, time_dim=16, noise_dim=8, sec_dim=sec_dim, particle_type_cfg=particle_type_cfg, ).eval() def _stage2_ar( target: str, generator: str, emb_dim: int = 6, pdg: int = 3, mat: int = 2, k_max: int = 5, history: str = "markov", ) -> Stage2Autoregressive: particle_cfg, material_cfg = _particle_material_cfg( _conditioning_for(target), emb_dim ) return Stage2Autoregressive( pdg_vocab=pdg, mat_vocab=mat, particle_cfg=particle_cfg, material_cfg=material_cfg, hidden_dim=32, n_res_blocks=2, generator=generator, time_dim=16, noise_dim=8, k_max=k_max, particle_type_cfg={"target": target}, history=history, attn_n_heads=2, attn_n_layers=1, ).eval() def _expected_type_dim(target: str, emb_dim: int) -> int: return PARTICLE_PHYS_DIM if target == "physical" else emb_dim # ── Stage-1 n_sec ownership ────────────────────────────────────────────────── def test_sample_flow_returns_none_n_sec_when_stage1_owns_no_head(): model = Stage1Model( pdg_vocab=3, mat_vocab=2, particle_cfg=_PHYS_CFG, material_cfg=_PHYS_CFG, hidden_dim=16, n_res_blocks=1, ) cond_cont, cond_cat = _cond(4) sample, n_sec = sample_flow(model, cond_cont, cond_cat, steps=2) assert sample.shape == (4, X_DIM) assert n_sec is None def test_sample_wgan_returns_none_n_sec_when_stage1_owns_no_head(): model = Stage1Model( pdg_vocab=3, mat_vocab=2, particle_cfg=_PHYS_CFG, material_cfg=_PHYS_CFG, hidden_dim=16, n_res_blocks=1, generator="wgan", noise_dim=8, ) cond_cont, cond_cat = _cond(4) sample, n_sec = sample_wgan(model, cond_cont, cond_cat) assert sample.shape == (4, X_DIM) assert n_sec is None def test_sample_flow_returns_n_sec_for_legacy_stage1(): model = Stage1Model( pdg_vocab=3, mat_vocab=2, particle_cfg=_PHYS_CFG, material_cfg=_PHYS_CFG, hidden_dim=16, n_res_blocks=1, n_sec_head_k_max=K_MAX, ) cond_cont, cond_cat = _cond(5) _, n_sec = sample_flow(model, cond_cont, cond_cat, steps=2) assert n_sec is not None and n_sec.shape == (5,) # ── Stage2OneShot: non-"physical" particle_type.target ────────────────────── @pytest.mark.parametrize("target", ["physical", "onehot", "embedding"]) def test_sample_secondaries_flow_shapes_by_target(target): B, emb_dim = 5, 6 decoder = _stage2_oneshot(target, "flow", emb_dim=emb_dim) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.randint(0, K_MAX + 1, (B,)) sec_cont, sec_type, sec_valid = sample_secondaries( decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2 ) assert sec_cont.shape == (B, K_MAX, CONT_SLOT_DIM) assert sec_type.shape == (B, K_MAX, _expected_type_dim(target, emb_dim)) assert sec_valid.shape == (B, K_MAX) assert torch.isfinite(sec_cont).all() assert torch.isfinite(sec_type).all() @pytest.mark.parametrize("target", ["physical", "onehot", "embedding"]) def test_sample_secondaries_wgan_shapes_by_target(target): B, emb_dim = 5, 6 decoder = _stage2_oneshot(target, "wgan", emb_dim=emb_dim) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.randint(0, K_MAX + 1, (B,)) sec_cont, sec_type, sec_valid = sample_secondaries_wgan( decoder, cond_cont, cond_cat, stage1_out, n_sec_pred ) assert sec_cont.shape == (B, K_MAX, CONT_SLOT_DIM) assert sec_type.shape == (B, K_MAX, _expected_type_dim(target, emb_dim)) assert sec_valid.shape == (B, K_MAX) # ── Stage2Autoregressive ───────────────────────────────────────────────────── @pytest.mark.parametrize("history", ["markov", "attention"]) @pytest.mark.parametrize("generator", ["flow", "wgan"]) @pytest.mark.parametrize("target", ["physical", "onehot", "embedding"]) def test_sample_secondaries_ar_shapes(target, generator, history): B, k_max, emb_dim = 4, 5, 6 decoder = _stage2_ar( target, generator, emb_dim=emb_dim, k_max=k_max, history=history ) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.randint(0, k_max + 1, (B,)) sec_cont, sec_type, sec_valid = sample_secondaries_ar( decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2 ) assert sec_cont.shape == (B, k_max, CONT_SLOT_DIM) assert sec_type.shape == (B, k_max, _expected_type_dim(target, emb_dim)) assert sec_valid.shape == (B, k_max) assert torch.isfinite(sec_cont).all() assert torch.isfinite(sec_type).all() @pytest.mark.parametrize("generator", ["flow", "wgan"]) @pytest.mark.parametrize("target", ["physical", "onehot", "embedding"]) def test_sample_secondaries_ar_valid_mask_matches_n_sec(target, generator): B, k_max, emb_dim = 3, 5, 6 decoder = _stage2_ar(target, generator, emb_dim=emb_dim, k_max=k_max) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.tensor([0, 2, k_max]) _, _, sec_valid = sample_secondaries_ar( decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2 ) for i, n in enumerate(n_sec_pred.tolist()): assert sec_valid[i, :n].all() assert not sec_valid[i, n:].any() def test_sample_secondaries_ar_first_slot_has_no_history(): """Slot 0 always has has_prev=False internally — nothing to assert on the public API directly, but a k_max=1 run should not crash on the "previous token" path at all (has_prev never true).""" B, emb_dim = 3, 6 decoder = _stage2_ar("physical", "flow", emb_dim=emb_dim, k_max=1) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.tensor([0, 1, 1]) sec_cont, sec_type, sec_valid = sample_secondaries_ar( decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2 ) assert sec_cont.shape == (B, 1, CONT_SLOT_DIM) assert sec_valid.tolist() == [[False], [True], [True]]