FallKLTN / src /fall_detection /experiment.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw History Blame Contribute Delete
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