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