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"])