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:
+24
-58
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user