8cebc4809d
First repo-wide ruff format pass, plus a note in CLAUDE.md to run ruff and ty periodically. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
68 lines
2.3 KiB
Python
68 lines
2.3 KiB
Python
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)
|