Team Ai
Apppublic

Hello532/motion-capture

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
visualizer.py335 linesDownload Raw Back to src
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