import numpy as np from giant.data.dataset import StepsDataset, train_val_split def _dummy(N=500, n_events=20): 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 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 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) 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)