FallKLTN / src /fall_detection /pose_extractor.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw History Blame Contribute Delete
3.55 kB
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import cv2
import numpy as np
@dataclass
class ExtractedVideo:
poses: np.ndarray
frame_numbers: np.ndarray
source_fps: float
class MediaPipePoseExtractor:
"""Thin wrapper so training code stays independent from MediaPipe."""
def __init__(
self,
model_complexity: int = 1,
min_detection_confidence: float = 0.5,
min_tracking_confidence: float = 0.5,
) -> None:
import mediapipe as mp
self._mp = mp
self._pose = mp.solutions.pose.Pose(
static_image_mode=False,
model_complexity=model_complexity,
smooth_landmarks=True,
enable_segmentation=False,
min_detection_confidence=min_detection_confidence,
min_tracking_confidence=min_tracking_confidence,
)
def process_rgb(self, rgb_frame: np.ndarray) -> np.ndarray | None:
result = self._pose.process(rgb_frame)
if result.pose_landmarks is None:
return None
return np.asarray(
[[point.x, point.y, point.z, point.visibility] for point in result.pose_landmarks.landmark],
dtype=np.float32,
)
def close(self) -> None:
self._pose.close()
def __enter__(self) -> "MediaPipePoseExtractor":
return self
def __exit__(self, *_: object) -> None:
self.close()
def extract_video(
path: str | Path, target_fps: float = 10.0, crop_wide_right_half: bool = True
) -> ExtractedVideo:
path = Path(path)
capture = cv2.VideoCapture(str(path))
if not capture.isOpened():
raise ValueError(f"Cannot open video: {path}")
source_fps = capture.get(cv2.CAP_PROP_FPS)
if not np.isfinite(source_fps) or source_fps <= 0:
source_fps = 25.0
sample_every = max(1, int(round(source_fps / target_fps)))
poses: list[np.ndarray] = []
frame_numbers: list[int] = []
with MediaPipePoseExtractor() as extractor:
frame_index = 0
while True:
ok, frame = capture.read()
if not ok:
break
if frame_index % sample_every == 0:
# Official URFD preview videos concatenate depth (left) and RGB
# (right) into a 640x240 frame. Cropping makes the person large
# enough for reliable landmark detection.
if crop_wide_right_half and frame.shape[1] / frame.shape[0] > 2.2:
frame = frame[:, frame.shape[1] // 2 :]
rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
pose = extractor.process_rgb(rgb)
if pose is None:
pose = np.zeros((33, 4), dtype=np.float32)
poses.append(pose)
frame_numbers.append(frame_index + 1) # URFD annotations are 1-based.
frame_index += 1
capture.release()
if not poses:
raise ValueError(f"No frames read from video: {path}")
return ExtractedVideo(
poses=np.stack(poses),
frame_numbers=np.asarray(frame_numbers, dtype=np.int32),
source_fps=float(source_fps),
)
def pad_window(window: np.ndarray, length: int) -> np.ndarray:
if len(window) >= length:
return window[:length]
if len(window) == 0:
return np.zeros((length, 33, 4), dtype=np.float32)
padding = np.repeat(window[-1][None, ...], length - len(window), axis=0)
return np.concatenate([window, padding], axis=0)