53fd2e4405
ruff removed unused imports across analysis.py and several test files. ty caught a wrong dict[int, int] annotation on StreamingStepsDataset's mat_map (materials are strings) and a real bug in steps_to_parquet.py where --compression none passed None to polars' write_parquet, which only accepts the literal "uncompressed". Also narrows a few Optional-typed attributes (ddpm_schedule, Normalizer.mean/std) with asserts and aligns __getitem__'s parameter name with torch's Dataset base class. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
62 lines
2.2 KiB
Python
62 lines
2.2 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)
|