"""Tests for Phase 2: secondary particle prediction.""" import numpy as np import pytest import torch from giant.constants import COND_DIM, EMB_DIM, K_MAX, SEC_DIM, X_DIM from giant.model.network import DenoisingMLP, SecondaryDecoder from giant.model.schedule import flow_matching_loss_secondary from giant.sample import sample_secondaries, snap_type_to_pdg_idx # ── helpers ────────────────────────────────────────────────────────────────── def _stage1(pdg=3, mat=2): return DenoisingMLP(pdg_vocab=pdg, mat_vocab=mat, hidden_dim=32, n_blocks=2) def _sec_decoder(pdg=3, mat=2): return SecondaryDecoder(pdg_vocab=pdg, mat_vocab=mat, hidden_dim=32, n_blocks=2) def _cond(B=8, pdg=3, mat=2): 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 # ── DenoisingMLP Phase-2 additions ─────────────────────────────────────────── def test_predict_n_sec_shape(): B = 8 model = _stage1() cond_cont, cond_cat = _cond(B) logits = model.predict_n_sec(cond_cont, cond_cat) assert logits.shape == (B, K_MAX + 1) def test_predict_n_sec_no_nan(): B = 8 model = _stage1() cond_cont, cond_cat = _cond(B) logits = model.predict_n_sec(cond_cont, cond_cat) assert torch.isfinite(logits).all() def test_pdg_embedding_weight_shape(): model = _stage1(pdg=5, mat=2) w = model.pdg_embedding_weight() assert w.shape == (5, EMB_DIM) # ── SecondaryDecoder ────────────────────────────────────────────────────────── def test_sec_decoder_output_shape(): B = 8 decoder = _sec_decoder() x_t = torch.randn(B, SEC_DIM) t = torch.rand(B) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) out = decoder(x_t, t, cond_cont, cond_cat, stage1_out) assert out.shape == (B, SEC_DIM) def test_sec_decoder_no_nan(): B = 4 decoder = _sec_decoder() x_t = torch.randn(B, SEC_DIM) t = torch.rand(B) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) out = decoder(x_t, t, cond_cont, cond_cat, stage1_out) assert torch.isfinite(out).all() def test_sec_decoder_gradients(): B = 4 decoder = _sec_decoder() x_t = torch.randn(B, SEC_DIM) t = torch.rand(B) cond_cont, cond_cat = _cond(B) stage1_out = torch.randn(B, X_DIM) decoder(x_t, t, cond_cont, cond_cat, stage1_out).sum().backward() for name, p in decoder.named_parameters(): assert p.grad is not None, f"no grad for {name}" # ── masked flow matching loss ───────────────────────────────────────────────── def test_flow_matching_loss_secondary_scalar(): B, pdg, mat = 8, 3, 2 decoder = _sec_decoder(pdg, mat) x1 = torch.randn(B, SEC_DIM) cond_cont, cond_cat = _cond(B, pdg, mat) stage1_out = torch.randn(B, X_DIM) sec_mask = torch.ones(B, K_MAX, dtype=torch.bool) loss = flow_matching_loss_secondary( decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask ) assert loss.shape == () assert loss.item() >= 0.0 def test_flow_matching_loss_secondary_mask_zeros_padding(): """Loss with all-zero mask (no valid secondaries) should be 0.""" B, pdg, mat = 4, 3, 2 decoder = _sec_decoder(pdg, mat) x1 = torch.randn(B, SEC_DIM) cond_cont, cond_cat = _cond(B, pdg, mat) stage1_out = torch.randn(B, X_DIM) sec_mask = torch.zeros(B, K_MAX, dtype=torch.bool) loss = flow_matching_loss_secondary( decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask ) assert loss.item() == pytest.approx(0.0, abs=1e-6) def test_flow_matching_loss_secondary_has_grad(): B, pdg, mat = 4, 3, 2 decoder = _sec_decoder(pdg, mat) x1 = torch.randn(B, SEC_DIM) cond_cont, cond_cat = _cond(B, pdg, mat) stage1_out = torch.randn(B, X_DIM) sec_mask = torch.ones(B, K_MAX, dtype=torch.bool) flow_matching_loss_secondary( decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask ).backward() assert any(p.grad is not None for p in decoder.parameters()) # ── sampling ────────────────────────────────────────────────────────────────── def test_sample_secondaries_shapes(): B, pdg, mat = 6, 3, 2 decoder = _sec_decoder(pdg, mat) cond_cont, cond_cat = _cond(B, pdg, mat) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.randint(0, K_MAX + 1, (B,)) sec_cont, sec_type_emb, sec_valid = sample_secondaries( decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=3 ) assert sec_cont.shape == (B, K_MAX, 4) assert sec_type_emb.shape == (B, K_MAX, EMB_DIM) assert sec_valid.shape == (B, K_MAX) assert sec_valid.dtype == torch.bool def test_sample_secondaries_valid_mask_matches_n_sec(): B, pdg, mat = 4, 3, 2 decoder = _sec_decoder(pdg, mat) cond_cont, cond_cat = _cond(B, pdg, mat) stage1_out = torch.randn(B, X_DIM) n_sec_pred = torch.tensor([0, 1, 3, K_MAX]) _, _, sec_valid = sample_secondaries( 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_snap_type_to_pdg_idx_shape(): B, pdg_vocab = 4, 5 emb_weight = torch.randn(pdg_vocab, EMB_DIM) sec_type_emb = torch.randn(B, K_MAX, EMB_DIM) idx = snap_type_to_pdg_idx(sec_type_emb, emb_weight) assert idx.shape == (B, K_MAX) assert idx.dtype == torch.int64 assert (idx >= 0).all() and (idx < pdg_vocab).all() # ── encode_secondaries round-trip ───────────────────────────────────────────── def test_encode_secondaries_energy_conservation(): """Decoded stick-breaking fractions must sum to ≈ e_sec.""" from giant.data.transforms import encode_secondaries rng = np.random.default_rng(42) N = 50 n_sec = rng.integers(1, 5, size=N) e_sec = rng.uniform(0.1, 10.0, size=N).astype(np.float32) sec_E_list = np.zeros((N, K_MAX), dtype=np.float32) sec_dir_list = np.zeros((N, K_MAX, 3), dtype=np.float32) sec_dir_list[:, :, 2] = 1.0 sec_valid = np.zeros((N, K_MAX), dtype=bool) for i in range(N): k = n_sec[i] energies = rng.dirichlet(np.ones(k)) * e_sec[i] energies = np.sort(energies)[::-1] sec_E_list[i, :k] = energies.astype(np.float32) sec_valid[i, :k] = True pre_dir = rng.standard_normal((N, 3)).astype(np.float32) pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True) sec_cont = encode_secondaries(sec_E_list, sec_dir_list, sec_valid, e_sec, pre_dir) assert sec_cont.shape == (N, K_MAX, 4) assert np.isfinite(sec_cont).all() def test_encode_secondaries_direction_encoding(): """Local-frame secondary directions should be unit vectors for valid slots.""" from giant.data.transforms import encode_secondaries rng = np.random.default_rng(7) N = 20 e_sec = np.ones(N, dtype=np.float32) * 5.0 sec_E_list = np.zeros((N, K_MAX), dtype=np.float32) sec_E_list[:, 0] = 3.0 sec_E_list[:, 1] = 2.0 sec_dir_list = rng.standard_normal((N, K_MAX, 3)).astype(np.float32) norms = np.linalg.norm(sec_dir_list, axis=-1, keepdims=True) sec_dir_list /= np.where(norms > 0, norms, 1.0) sec_valid = np.zeros((N, K_MAX), dtype=bool) sec_valid[:, :2] = True pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32) sec_cont = encode_secondaries(sec_E_list, sec_dir_list, sec_valid, e_sec, pre_dir) # dir columns are sec_cont[:, :, 1:4] local_dirs = sec_cont[:, :2, 1:4] # (N, 2, 3) — valid slots only norms_out = np.linalg.norm(local_dirs, axis=-1) np.testing.assert_allclose(norms_out, 1.0, atol=1e-5)