| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
| from pathlib import Path |
|
|
| import numpy as np |
| from sklearn.model_selection import GroupShuffleSplit |
|
|
|
|
| @dataclass(frozen=True) |
| class DatasetSplits: |
| train: np.ndarray |
| val: np.ndarray |
| test: np.ndarray |
|
|
|
|
| def load_pose_dataset(path: str | Path) -> dict[str, np.ndarray]: |
| """Load the compressed dataset produced by prepare_dataset.py.""" |
| with np.load(path, allow_pickle=False) as data: |
| required = {"poses", "labels", "groups", "sources"} |
| missing = required.difference(data.files) |
| if missing: |
| raise ValueError(f"Dataset is missing arrays: {sorted(missing)}") |
| return {key: data[key] for key in required} |
|
|
|
|
| def group_train_val_test_split( |
| labels: np.ndarray, |
| groups: np.ndarray, |
| test_size: float = 0.20, |
| val_size: float = 0.20, |
| seed: int = 42, |
| ) -> DatasetSplits: |
| """Split without putting windows from the same source video in two sets.""" |
| indices = np.arange(len(labels)) |
| relative_val_size = val_size / (1.0 - test_size) |
| selected = None |
| for attempt in range(100): |
| outer = GroupShuffleSplit( |
| n_splits=1, test_size=test_size, random_state=seed + attempt * 2 |
| ) |
| train_val_idx, test_idx = next(outer.split(indices, labels, groups)) |
| inner = GroupShuffleSplit( |
| n_splits=1, |
| test_size=relative_val_size, |
| random_state=seed + attempt * 2 + 1, |
| ) |
| inner_train, inner_val = next( |
| inner.split(train_val_idx, labels[train_val_idx], groups[train_val_idx]) |
| ) |
| train_idx = train_val_idx[inner_train] |
| val_idx = train_val_idx[inner_val] |
| if all(len(np.unique(labels[idx])) == 2 for idx in (train_idx, val_idx, test_idx)): |
| selected = (train_idx, val_idx, test_idx) |
| break |
| if selected is None: |
| raise ValueError("Could not create group splits containing both classes") |
| train_idx, val_idx, test_idx = selected |
|
|
| split_groups = [set(groups[idx].tolist()) for idx in (train_idx, val_idx, test_idx)] |
| if any(split_groups[i] & split_groups[j] for i in range(3) for j in range(i + 1, 3)): |
| raise RuntimeError("Group leakage detected") |
| return DatasetSplits(train=train_idx, val=val_idx, test=test_idx) |
|
|
|
|
| def save_splits(path: str | Path, splits: DatasetSplits) -> None: |
| path = Path(path) |
| path.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed(path, train=splits.train, val=splits.val, test=splits.test) |
|
|
|
|
| def load_splits(path: str | Path) -> DatasetSplits: |
| with np.load(path) as data: |
| return DatasetSplits(train=data["train"], val=data["val"], test=data["test"]) |
|
|