#!/usr/bin/env python3 from __future__ import annotations import argparse from pathlib import Path import numpy as np import pandas as pd from tqdm import tqdm from fall_detection.config import load_config from fall_detection.pose_extractor import extract_video, pad_window VIDEO_EXTENSIONS = {".mp4", ".avi", ".mov", ".mkv"} def load_urfd_annotations(video_dir: Path) -> dict[str, dict[int, int]]: annotations: dict[str, dict[int, int]] = {} for path in (video_dir / "annotations").glob("*.csv"): frame = pd.read_csv(path, header=None, usecols=[0, 1, 2]) frame.columns = ["sequence", "frame", "posture"] for sequence, rows in frame.groupby("sequence"): annotations[str(sequence)] = dict( zip(rows["frame"].astype(int), rows["posture"].astype(int)) ) return annotations def window_label( class_name: str, sequence_name: str, frame_numbers: np.ndarray, annotations: dict[str, dict[int, int]], ) -> int: if class_name == "normal": return 0 if sequence_name not in annotations: return 1 mapping = annotations[sequence_name] postures = np.asarray([mapping.get(int(frame), -1) for frame in frame_numbers]) # 0 is the falling transition and 1 is lying after the fall. return int(np.mean(postures >= 0) >= 0.15) def main() -> None: parser = argparse.ArgumentParser(description="Extract MediaPipe pose windows") parser.add_argument("--config", default="configs/default.yaml") parser.add_argument("--input", type=Path) parser.add_argument("--output", type=Path) args = parser.parse_args() config = load_config(args.config) data_config = config["data"] video_dir = args.input or Path(data_config["video_dir"]) output = args.output or Path(data_config["dataset_path"]) length = int(data_config["sequence_length"]) stride = int(data_config["stride"]) min_valid_ratio = float(data_config["min_valid_pose_ratio"]) target_fps = float(data_config["target_fps"]) annotations = load_urfd_annotations(video_dir) videos = sorted( path for path in video_dir.glob("*/*/*") if path.suffix.lower() in VIDEO_EXTENSIONS ) if not videos: raise SystemExit( f"No videos found in {video_dir}/{{fall,normal}}//. " "Run scripts/download_urfd.py first." ) poses: list[np.ndarray] = [] labels: list[int] = [] groups: list[str] = [] sources: list[str] = [] skipped = 0 for video_path in tqdm(videos, desc="Extracting poses"): class_name = video_path.relative_to(video_dir).parts[0] if class_name not in {"fall", "normal"}: continue group = video_path.parent.name sequence_name = group extracted = extract_video(video_path, target_fps=target_fps) max_start = max(len(extracted.poses) - length, 0) starts = list(range(0, max_start + 1, stride)) or [0] if starts[-1] != max_start: starts.append(max_start) for start in starts: stop = start + length window = pad_window(extracted.poses[start:stop], length) frame_numbers = extracted.frame_numbers[start:stop] valid_ratio = float(np.mean(window[:, :, 3].max(axis=1) > 0)) if valid_ratio < min_valid_ratio: skipped += 1 continue poses.append(window) labels.append( window_label( class_name, sequence_name, frame_numbers, annotations ) ) groups.append(group) sources.append(f"{video_path.as_posix()}#sample={start}:{stop}") if not poses: raise SystemExit("All windows were skipped because pose detection failed") output.parent.mkdir(parents=True, exist_ok=True) np.savez_compressed( output, poses=np.stack(poses).astype(np.float32), labels=np.asarray(labels, dtype=np.int64), groups=np.asarray(groups), sources=np.asarray(sources), ) unique, counts = np.unique(labels, return_counts=True) distribution = dict(zip(unique.tolist(), counts.tolist())) print(f"Saved {len(poses)} windows to {output}") print(f"Class distribution {distribution}; skipped {skipped} low-visibility windows") if __name__ == "__main__": main()