Download src/fall_detection/experiment.py from minhy112/FallKLTN: direct link, hf CLI and curl.
- Browser
- Download file 2.15 kB
-
https://huggingface.co/minhy112/FallKLTN/resolve/main/src/fall_detection/experiment.py
- Command line
-
hf download hf://minhy112/FallKLTN/src/fall_detection/experiment.py
-
curl -L -o experiment.py https://huggingface.co/minhy112/FallKLTN/resolve/main/src/fall_detection/experiment.py
2.15 kB
| 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 | |