File size: 2,726 Bytes
9313a90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
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"])