Implement Phase 2: secondary particle prediction

Two-stage factorisation: Stage 1 predicts 9D primary kinematics + n_sec
classification head (COND_DIM reduced to 8, dropping n_sec/e_sec inputs);
Stage 2 (SecondaryDecoder) generates K_MAX=15 secondary slots via masked
flow matching over (stick_logit, local_dir, type_emb) conditioned on Stage 1
output. Joint training with combined loss L_s1 + λ_nsec*L_nsec + λ_s2*L_s2.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-29 11:34:31 +02:00
parent c627142135
commit e6e0eb22bf
18 changed files with 1174 additions and 234 deletions
+24 -58
View File
@@ -1,67 +1,33 @@
import numpy as np
from giant.data.dataset import StepsDataset, train_val_split
from giant.data.dataset import make_event_split
def _dummy(N=500, n_events=20):
def test_make_event_split_sizes():
rng = np.random.default_rng(42)
data = {"event_id": rng.integers(0, n_events, size=N)}
cond_cont = rng.standard_normal((N, 9)).astype(np.float32)
cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64)
target = rng.standard_normal((N, 6)).astype(np.float32)
return data, cond_cont, cond_cat, target
event_ids = rng.integers(0, 50, size=1000)
train_set, val_set = make_event_split(event_ids, val_fraction=0.2)
unique = np.unique(event_ids)
assert len(train_set) + len(val_set) == len(unique)
def test_dataset_length():
data, cond_cont, cond_cat, target = _dummy()
assert len(StepsDataset(cond_cont, cond_cat, target)) == len(target)
def test_dataset_item_shapes():
data, cond_cont, cond_cat, target = _dummy()
c, k, t = StepsDataset(cond_cont, cond_cat, target)[0]
assert c.shape == (9,)
assert k.shape == (2,)
assert t.shape == (6,)
def test_split_sizes_sum_to_total():
data, cond_cont, cond_cat, target = _dummy(N=500)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=0.2
)
assert len(train_ds) + len(val_ds) == 500
def test_split_no_empty_sets():
data, cond_cont, cond_cat, target = _dummy(N=500, n_events=20)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=0.2
)
assert len(val_ds) > 0
assert len(train_ds) > 0
def test_split_event_leakage():
"""Train and val must not share any event_id."""
N = 1000
n_events = 50
def test_make_event_split_no_overlap():
rng = np.random.default_rng(7)
event_ids = rng.integers(0, n_events, size=N)
data = {"event_id": event_ids}
cond_cont = rng.standard_normal((N, 9)).astype(np.float32)
cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64)
target = rng.standard_normal((N, 6)).astype(np.float32)
event_ids = rng.integers(0, 50, size=1000)
train_set, val_set = make_event_split(event_ids, val_fraction=0.2)
assert train_set.isdisjoint(val_set)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=0.2
)
# Recover which event_ids ended up in each split via the indices
# (The dataset doesn't store event_ids, so we check via the original mask logic)
unique_events = np.unique(event_ids)
rng2 = np.random.default_rng(42)
rng2.shuffle(unique_events)
n_val = max(1, int(len(unique_events) * 0.2))
val_events = set(unique_events[:n_val].tolist())
train_events = set(unique_events[n_val:].tolist())
assert val_events.isdisjoint(train_events)
def test_make_event_split_no_empty_sets():
rng = np.random.default_rng(0)
event_ids = rng.integers(0, 20, size=500)
train_set, val_set = make_event_split(event_ids, val_fraction=0.2)
assert len(train_set) > 0
assert len(val_set) > 0
def test_make_event_split_reproducible():
event_ids = np.arange(100)
a_tr, a_val = make_event_split(event_ids, val_fraction=0.1, seed=42)
b_tr, b_val = make_event_split(event_ids, val_fraction=0.1, seed=42)
assert a_tr == b_tr
assert a_val == b_val