File size: 2,154 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
from __future__ import annotations

from pathlib import Path

import numpy as np

from fall_detection.config import load_config
from fall_detection.data import (
    DatasetSplits,
    group_train_val_test_split,
    load_pose_dataset,
    load_splits,
    save_splits,
)
from fall_detection.features import featurize_dataset


def prepare_experiment_data(
    dataset_path: str | Path,
    config_path: str | Path,
    output_root: str | Path,
    seed: int | None = None,
) -> tuple[dict, dict[str, np.ndarray], np.ndarray, DatasetSplits]:
    config = load_config(config_path)
    if seed is not None:
        config["seed"] = int(seed)
    dataset = load_pose_dataset(dataset_path)
    features = featurize_dataset(
        dataset["poses"], config["data"]["visibility_threshold"]
    )
    split_path = Path(output_root) / "splits.npz"
    if split_path.exists():
        splits = load_splits(split_path)
        all_indices = np.concatenate([splits.train, splits.val, splits.test])
        if len(all_indices) != len(dataset["labels"]) or all_indices.max() >= len(dataset["labels"]):
            raise ValueError(
                f"Existing {split_path} does not match this dataset; remove it or use another output directory"
            )
    else:
        splits = group_train_val_test_split(
            dataset["labels"],
            dataset["groups"],
            test_size=config["data"]["test_size"],
            val_size=config["data"]["val_size"],
            seed=config["seed"],
        )
        save_splits(split_path, splits)
    return config, dataset, features, splits


def split_summary(
    labels: np.ndarray, groups: np.ndarray, splits: DatasetSplits
) -> dict[str, dict[str, int]]:
    result: dict[str, dict[str, int]] = {}
    for name, indices in (
        ("train", splits.train),
        ("validation", splits.val),
        ("test", splits.test),
    ):
        result[name] = {
            "samples": int(len(indices)),
            "groups": int(len(np.unique(groups[indices]))),
            "normal": int(np.sum(labels[indices] == 0)),
            "fall": int(np.sum(labels[indices] == 1)),
        }
    return result