minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw
History Blame Contribute Delete
2.73 kB
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"])