FallKLTN / scripts /prepare_dataset.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw History Blame Contribute Delete
4.41 kB
#!/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()