File size: 4,406 Bytes
9313a90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
#!/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}}/<group>/. "
            "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()