Hello532/motion-capture
0
1"""Visualization module.2 3Generates five types of visual output:4- X-Y trajectory plots5- Kinematics curves (angle / velocity / acceleration vs time)6- Joint angle heatmaps7- Animated trajectory overlay8- 3D keypoint scatter (optional)9"""10 11from pathlib import Path12from typing import Dict, List, Optional, Tuple13 14import cv215import matplotlib16import matplotlib.pyplot as plt17import numpy as np18from matplotlib.animation import FuncAnimation19 20matplotlib.use("Agg") # non-interactive backend21 22# MediaPipe bone connections for stick-figure rendering23BONE_CONNECTIONS = [24 (11, 12), (11, 23), (12, 24), (23, 24), # torso25 (11, 13), (13, 15), (12, 14), (14, 16), # arms26 (23, 25), (25, 27), (24, 26), (26, 28), # legs27 (15, 17), (15, 19), (17, 19), # left hand28 (16, 18), (16, 20), (18, 20), # right hand29 (27, 29), (27, 31), (29, 31), # left foot30 (28, 30), (28, 32), (30, 32), # right foot31]32 33# Selected keypoints for trajectory overlay (major joints)34TRAJECTORY_KEYPOINTS = [11, 12, 13, 14, 15, 16, 23, 24, 25, 26, 27, 28]35TRAJECTORY_COLORS = [36 "#FF0000", "#0000FF", "#FF4444", "#4444FF", "#FF8888", "#8888FF",37 "#00AA00", "#AA00AA", "#00FF44", "#FF44FF", "#00AAAA", "#AAAAFF",38]39 40 41def plot_trajectory(42 positions: np.ndarray,43 keypoint_indices: List[int],44 keypoint_names: List[str],45 output_path: str,46 title: str = "Joint Trajectories (X-Y)",47):48 """Plot spatial (X-Y) trajectories of selected keypoints.49 50 Parameters51 ----------52 positions : np.ndarray53 Shape (T, 33, 2).54 keypoint_indices : list[int]55 Which keypoints to plot.56 keypoint_names : list[str]57 Names for the legend.58 output_path : str59 Path to save the figure (PNG/PDF).60 title : str61 Plot title.62 """63 fig, ax = plt.subplots(figsize=(10, 8))64 for idx, name in zip(keypoint_indices, keypoint_names):65 x = positions[:, idx, 0]66 y = positions[:, idx, 1]67 valid = ~np.isnan(x) & ~np.isnan(y)68 if valid.sum() > 0:69 ax.plot(x[valid], y[valid], linewidth=0.8, label=name, alpha=0.8)70 # Mark start and end71 if valid.sum() >= 2:72 start_t = np.argmax(valid)73 end_t = len(valid) - np.argmax(valid[::-1]) - 174 ax.scatter(x[start_t], y[start_t], s=20, marker="o", zorder=5)75 ax.scatter(x[end_t], y[end_t], s=20, marker="s", zorder=5)76 77 ax.invert_yaxis()78 ax.set_xlabel("X (pixels)")79 ax.set_ylabel("Y (pixels)")80 ax.set_title(title)81 ax.legend(loc="upper right", fontsize=7)82 ax.set_aspect("equal")83 fig.tight_layout()84 fig.savefig(output_path, dpi=150)85 plt.close(fig)86 87 88def plot_kinematics(89 time_sec: np.ndarray,90 signals: Dict[str, np.ndarray],91 output_path: str,92 title: str = "Kinematics Curves",93 xlabel: str = "Time (s)",94 ylabel: str = "Angle (deg)",95 overlay_signals: Optional[Dict[str, np.ndarray]] = None,96 overlay_ylabel: str = "Angular Velocity (deg/s)",97):98 """Plot kinematics curves, optionally with overlaid secondary signals.99 100 Parameters101 ----------102 time_sec : np.ndarray103 Time axis in seconds.104 signals : dict105 Primary signals {label: array}.106 output_path : str107 title, xlabel, ylabel : str108 overlay_signals : dict, optional109 Secondary signals plotted on a twin y-axis.110 overlay_ylabel : str111 """112 fig, ax1 = plt.subplots(figsize=(12, 5))113 114 for label, sig in signals.items():115 valid = ~np.isnan(sig)116 if valid.sum() > 0:117 ax1.plot(time_sec[valid], sig[valid], linewidth=1.2, label=label)118 ax1.set_xlabel(xlabel)119 ax1.set_ylabel(ylabel)120 ax1.legend(loc="upper left", fontsize=8)121 ax1.grid(True, alpha=0.3)122 123 if overlay_signals:124 ax2 = ax1.twinx()125 for label, sig in overlay_signals.items():126 valid = ~np.isnan(sig)127 if valid.sum() > 0:128 ax2.plot(time_sec[valid], sig[valid], linewidth=0.8,129 linestyle="--", alpha=0.6, label=label)130 ax2.set_ylabel(overlay_ylabel)131 ax2.legend(loc="upper right", fontsize=8)132 133 ax1.set_title(title)134 fig.tight_layout()135 fig.savefig(output_path, dpi=150)136 plt.close(fig)137 138 139def plot_angle_heatmap(140 time_sec: np.ndarray,141 angle_names: List[str],142 angles_matrix: np.ndarray,143 output_path: str,144 title: str = "Joint Angle Heatmap",145):146 """Plot joint angles as a heatmap (angles × time).147 148 Parameters149 ----------150 time_sec : np.ndarray151 Time axis.152 angle_names : list[str]153 Joint names (y-axis labels).154 angles_matrix : np.ndarray155 Shape (num_angles, T), each row is one angle's time series.156 output_path : str157 title : str158 """159 fig, ax = plt.subplots(figsize=(12, 5))160 im = ax.imshow(161 angles_matrix,162 aspect="auto",163 cmap="inferno",164 interpolation="bilinear",165 extent=[time_sec[0], time_sec[-1], len(angle_names) - 0.5, -0.5],166 )167 ax.set_yticks(range(len(angle_names)))168 ax.set_yticklabels(angle_names)169 ax.set_xlabel("Time (s)")170 ax.set_title(title)171 cbar = fig.colorbar(im, ax=ax)172 cbar.set_label("Angle (deg)")173 fig.tight_layout()174 fig.savefig(output_path, dpi=150)175 plt.close(fig)176 177 178def animate_with_trajectory(179 frames_bgr: np.ndarray,180 positions: np.ndarray,181 output_path: str,182 fps: float = 30.0,183 trail_length: int = 20,184 show_skeleton: bool = True,185):186 """Create an animation with trajectory overlay on original video.187 188 Parameters189 ----------190 frames_bgr : np.ndarray191 Shape (T, H, W, 3), original video frames.192 positions : np.ndarray193 Shape (T, 33, 2), interpolated keypoint positions.194 output_path : str195 Path for output MP4/GIF.196 fps : float197 Output animation FPS.198 trail_length : int199 Number of past frames to show as fading trail.200 show_skeleton : bool201 Whether to draw stick-figure skeleton.202 """203 T = min(len(frames_bgr), len(positions))204 output_path = Path(output_path)205 suffix = output_path.suffix.lower()206 207 if suffix == ".gif":208 _animate_gif(frames_bgr, positions, T, output_path, fps, trail_length, show_skeleton)209 else:210 _animate_mp4(frames_bgr, positions, T, output_path, fps, trail_length, show_skeleton)211 212 213def _draw_overlay(214 frame: np.ndarray,215 positions_t: np.ndarray,216 trail: List[np.ndarray],217 trail_length: int,218 show_skeleton: bool,219 alpha: float = 1.0,220) -> np.ndarray:221 """Draw trajectory trail and optional skeleton on a single frame."""222 canvas = frame.copy()223 h, w = canvas.shape[:2]224 225 # Draw fading trajectory trail226 for age, pos in enumerate(reversed(trail[-trail_length:])):227 fade = max(0.15, 1.0 - age / trail_length)228 for kp_idx, color_hex in zip(TRAJECTORY_KEYPOINTS, TRAJECTORY_COLORS):229 if kp_idx >= pos.shape[0]:230 continue231 x, y = pos[kp_idx]232 if np.isnan(x) or np.isnan(y):233 continue234 color = _hex_to_bgr(color_hex, fade)235 px, py = int(np.clip(x, 0, w - 1)), int(np.clip(y, 0, h - 1))236 cv2.circle(canvas, (px, py), 3, color, -1) # filled circle237 238 # Draw skeleton239 if show_skeleton:240 for p1, p2 in BONE_CONNECTIONS:241 if p1 >= positions_t.shape[0] or p2 >= positions_t.shape[0]:242 continue243 x1, y1 = positions_t[p1]244 x2, y2 = positions_t[p2]245 if np.isnan(x1) or np.isnan(y1) or np.isnan(x2) or np.isnan(y2):246 continue247 pt1 = (int(np.clip(x1, 0, w - 1)), int(np.clip(y1, 0, h - 1)))248 pt2 = (int(np.clip(x2, 0, w - 1)), int(np.clip(y2, 0, h - 1)))249 cv2.line(canvas, pt1, pt2, (0, 255, 0), 2)250 251 return canvas252 253 254def _animate_mp4(frames_bgr, positions, T, output_path, fps, trail_length, show_skeleton):255 """Write animation as MP4 using OpenCV VideoWriter."""256 h, w = frames_bgr[0].shape[:2]257 fourcc = cv2.VideoWriter_fourcc(*"mp4v")258 writer = cv2.VideoWriter(str(output_path), fourcc, fps, (w, h))259 260 trail: List[np.ndarray] = []261 for t in range(T):262 trail.append(positions[t].copy())263 canvas = _draw_overlay(frames_bgr[t], positions[t], trail, trail_length, show_skeleton)264 writer.write(canvas)265 writer.release()266 267 268def _animate_gif(frames_bgr, positions, T, output_path, fps, trail_length, show_skeleton):269 """Write animation as GIF using imageio (via Pillow)."""270 from PIL import Image271 272 trail: List[np.ndarray] = []273 pil_frames = []274 for t in range(T):275 trail.append(positions[t].copy())276 canvas = _draw_overlay(frames_bgr[t], positions[t], trail, trail_length, show_skeleton)277 canvas_rgb = cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB)278 pil_frames.append(Image.fromarray(canvas_rgb))279 280 duration = int(1000 / fps) if fps > 0 else 33281 pil_frames[0].save(282 str(output_path),283 save_all=True,284 append_images=pil_frames[1:],285 duration=duration,286 loop=0,287 )288 289 290def plot_landmarks_3d(291 positions: np.ndarray,292 output_path: str,293 frame_idx: int = 0,294):295 """Plot 3D scatter of keypoints (2D positions, z = keypoint index).296 297 Parameters298 ----------299 positions : np.ndarray300 Shape (T, 33, 2).301 output_path : str302 frame_idx : int303 Which frame to plot.304 """305 if frame_idx >= len(positions):306 frame_idx = len(positions) - 1307 308 pos = positions[frame_idx] # (33, 2)309 valid = ~np.isnan(pos).any(axis=1)310 311 fig = plt.figure(figsize=(10, 8))312 ax = fig.add_subplot(111, projection="3d")313 314 xs = pos[valid, 0]315 ys = pos[valid, 1]316 zs = np.arange(33)[valid]317 318 ax.scatter(xs, ys, zs, c=zs, cmap="viridis", s=30)319 ax.set_xlabel("X (pixels)")320 ax.set_ylabel("Y (pixels)")321 ax.set_zlabel("Keypoint Index")322 ax.set_title(f"3D Keypoint Scatter — Frame {frame_idx}")323 ax.invert_yaxis()324 325 fig.tight_layout()326 fig.savefig(output_path, dpi=150)327 plt.close(fig)328 329 330def _hex_to_bgr(hex_color: str, alpha: float = 1.0) -> Tuple[int, int, int]:331 """Convert hex color to BGR tuple with optional alpha blending on black."""332 hex_color = hex_color.lstrip("#")333 r, g, b = int(hex_color[0:2], 16), int(hex_color[2:4], 16), int(hex_color[4:6], 16)334 return (int(b * alpha), int(g * alpha), int(r * alpha))335 