File size: 3,550 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
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)