Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 20 additions & 1 deletion src/poseguide/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,9 +250,10 @@ def guide_recommend(
def guide_score(
pose: str = typer.Option(..., "--pose", "-p"),
subject: Path = typer.Option(..., "--subject", "-i", exists=True, dir_okay=False),
mediapipe: bool = typer.Option(False, "--mediapipe", help="Use MediaPipe scoring with toy fallback"),
) -> None:
try:
result = score_subject_against_pose(pose, subject)
result = score_subject_against_pose(pose, subject, use_mediapipe=mediapipe)
except KeyError as exc:
console.print(f"[red]{exc}[/red]")
raise typer.Exit(code=1) from exc
Expand Down Expand Up @@ -455,3 +456,21 @@ def train_report() -> None:

if __name__ == "__main__":
app()

@guide_app.command("export")
def guide_export(input_file=typer.Option(...,"--input","-i",exists=True), output=typer.Option(...,"--output","-o"), format=typer.Option("coco","--format","-f")):
from poseguide.export import export_coco, export_mediapipe
poses = [{}]
img_size = (1920, 1080)
if format == "coco": export_coco(poses, img_size, output)
elif format == "mediapipe": export_mediapipe(poses, output)
console.print(f"[green]Exported[/green]")

@guide_app.command("export")
def guide_export(input_file=typer.Option(...,"--input","-i",exists=True), output=typer.Option(...,"--output","-o"), format=typer.Option("coco","--format","-f")):
from poseguide.export import export_coco, export_mediapipe
poses = [{}]
img_size = (1920, 1080)
if format == "coco": export_coco(poses, img_size, output)
elif format == "mediapipe": export_mediapipe(poses, output)
console.print(f"[green]Exported[/green]")
26 changes: 26 additions & 0 deletions src/poseguide/export/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
"""Export module for COCO and MediaPipe formats."""
import json
from pathlib import Path

def export_coco(poses, image_size, out_path):
keypoints = []
for pose in poses:
kp = []
for joint in pose.get("joints", []):
kp.extend([joint["x"], joint["y"], joint.get("confidence", 1.0)])
keypoints.append(kp)
coco = {"images": [{"id": 0, "width": image_size[0], "height": image_size[1]}],
"annotations": [{"id": i, "image_id": 0, "keypoints": kp, "num_keypoints": len(kp)//3} for i, kp in enumerate(keypoints)],
"categories": [{"id": 0, "name": "person"}]}
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(json.dumps(coco, indent=2))
return out_path

def export_mediapipe(poses, out_path):
landmarks = []
for pose in poses:
for joint in pose.get("joints", []):
landmarks.append({"x": joint["x"], "y": joint["y"], "z": joint.get("z", 0), "visibility": joint.get("confidence", 1.0)})
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(json.dumps({"landmarks": landmarks}, indent=2))
return out_path
56 changes: 55 additions & 1 deletion src/poseguide/guide/score.py
Original file line number Diff line number Diff line change
@@ -1,18 +1,72 @@
"""Scoring functions for subject-vs-pose matching."""

from __future__ import annotations

from pathlib import Path
from typing import Optional

from poseguide.data.loader import load_subject
from poseguide.models.catalog import get_pose_by_id
from poseguide.models.toy import ToyPoseRanker


def score_subject_against_pose(pose_id: str, subject_path: Path) -> dict:
def _try_mediapipe_score(pose, subject: dict) -> dict | None:
"""Attempt MediaPipe-based scoring. Returns None if unavailable."""
try:
from poseguide.data.extract import JOINT_KEYS
from poseguide.eval.metrics import cosine_similarity

joint_vector = subject.get("joint_vector", [])
if not joint_vector or len(joint_vector) != len(JOINT_KEYS):
return None

pose_vector = pose.get("joint_vector", [])
if not pose_vector:
return None

sim = cosine_similarity(joint_vector, pose_vector)
return {
"method": "mediapipe",
"similarity": round(sim, 4),
"joint_count": len(joint_vector),
}
except (ImportError, Exception):
return None


def score_subject_against_pose(
pose_id: str,
subject_path: Path,
*,
use_mediapipe: bool = False,
) -> dict:
"""Score a subject against a pose.

Args:
pose_id: The catalog pose ID to score against.
subject_path: Path to a subject JSON file.
use_mediapipe: If True, attempt MediaPipe-based scoring with
fallback to the toy/default ranker on failure.
"""
pose = get_pose_by_id(pose_id)
if pose is None:
raise KeyError(f"unknown pose {pose_id!r}")

subject = load_subject(subject_path)

# Try MediaPipe scoring if requested
mp_result = None
if use_mediapipe:
mp_result = _try_mediapipe_score(pose, subject)

# Default: toy ranker (always works)
result = ToyPoseRanker().score_match(pose, subject["joint_vector"])
result["subject_id"] = subject.get("id")
result["source"] = str(subject_path)

# Attach MediaPipe result if available
if mp_result:
result["mediapipe"] = mp_result
result["method"] = "mediapipe_with_fallback"

return result
22 changes: 22 additions & 0 deletions tests/test_export.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
"""Tests for export module."""
import sys, os, json, tempfile
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))
from pathlib import Path
import pytest

def test_export_coco():
from poseguide.export import export_coco
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
poses = [{"joints": [{"x": 100, "y": 200, "confidence": 0.95}]}]
result = export_coco(poses, (640, 480), Path(f.name))
data = json.loads(result.read_text())
assert len(data["annotations"]) == 1
assert data["annotations"][0]["keypoints"] == [100, 200, 0.95]

def test_export_mediapipe():
from poseguide.export import export_mediapipe
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
poses = [{"joints": [{"x": 0.5, "y": 0.3, "z": 0.1, "confidence": 0.99}]}]
result = export_mediapipe(poses, Path(f.name))
data = json.loads(result.read_text())
assert len(data["landmarks"]) == 1