Files
giant/tests/test_dataset.py
T
lars 53fd2e4405 Add ruff and ty as dev dependencies, fix lint/type findings
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>
2026-06-18 17:38:40 +02:00

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)