Skip to content
Draft
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
23 changes: 15 additions & 8 deletions dlclivegui/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -507,10 +507,9 @@ class SkeletonColorMode(str, Enum):


class SkeletonStyle(BaseModel):
visible: bool = False
color_mode: SkeletonColorMode = SkeletonColorMode.SOLID
color_bgr: BGR = (0, 255, 255) # default if SOLID
thickness: int = Field(defalt=2, ge=1, le=20) # base thickness in pixels
thickness: int = Field(default=2, ge=1, le=20) # base thickness in pixels
gradient_steps: int = Field(default=16, ge=2, le=128) # segments per edge when gradient
scale_with_zoom: bool = True # scale thickness with (sx, sy)

Expand All @@ -521,15 +520,23 @@ def effective_thickness(self, sx: float, sy: float) -> int:


class VisualizationSettings(BaseModel):
p_cutoff: float = Field(default=0.6, ge=0.0, le=1.0)
p_cutoff: float = Field(
default=0.6,
ge=0.0,
le=1.0,
)
colormap: str = "hot"
bbox_color: tuple[int, int, int] = (0, 0, 255)
bbox_color: BGR = (0, 0, 255)

def get_bbox_color_bgr(self) -> tuple[int, int, int]:
"""Get bounding box color in BGR format"""
show_pose: bool = True
show_skeleton: bool = False
skeleton_style: SkeletonStyle = Field(default_factory=SkeletonStyle)

def get_bbox_color_bgr(self) -> BGR:
if isinstance(self.bbox_color, (list, tuple)) and len(self.bbox_color) == 3:
return tuple(int(c) for c in self.bbox_color)
return (0, 0, 255) # default red
return tuple(int(value) for value in self.bbox_color)

return (0, 0, 255)


class RecordingSettings(BaseModel):
Expand Down
23 changes: 22 additions & 1 deletion dlclivegui/display/display.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# dlclivegui/utils/display.py
# dlclivegui/display/display.py
from __future__ import annotations

import enum
Expand Down Expand Up @@ -215,6 +215,27 @@ def draw_keypoints(overlay, p_cutoff, sx, ox, sy, oy, radius, cmap, keypoints: n
cv2.drawMarker(overlay, (xs, ys), bgr, marker, radius * 2, 2)


def keypoint_colors_bgr(
colormap: str,
count: int,
) -> tuple[tuple[int, int, int], ...]:
"""Return the display color assigned to each keypoint."""
if count < 0:
raise ValueError(f"Keypoint count must be non-negative; received {count}.")

cmap = plt.get_cmap(colormap)
denominator = max(count - 1, 1)

return tuple(
(
round(cmap(index / denominator)[2] * 255),
round(cmap(index / denominator)[1] * 255),
round(cmap(index / denominator)[0] * 255),
)
for index in range(count)
)


def draw_pose(
frame: np.ndarray,
pose: np.ndarray,
Expand Down
188 changes: 188 additions & 0 deletions dlclivegui/display/overlays.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,188 @@
"""Composition of pose, skeleton, and bounding-box overlays."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Protocol

import numpy as np

from dlclivegui.config import (
BGR,
SkeletonColorMode,
SkeletonStyle,
)

from .display import (
draw_bbox,
draw_pose,
keypoint_colors_bgr,
)
from .skeleton import (
ResolvedSkeleton,
SkeletonRenderCode,
SkeletonResolutionError,
draw_skeleton,
resolve_packet_skeleton,
)


class PosePacketLike(Protocol):
keypoint_names: list[str] | None
skeleton_id: str | None
skeleton_edges: tuple[tuple[str, str], ...] | None


@dataclass(frozen=True, slots=True)
class PoseOverlaySettings:
visible: bool
p_cutoff: float
colormap: str


@dataclass(frozen=True, slots=True)
class BoundingBoxOverlaySettings:
visible: bool
coordinates: tuple[int, int, int, int]
color_bgr: BGR


@dataclass(frozen=True, slots=True)
class SkeletonOverlaySettings:
visible: bool
style: SkeletonStyle


@dataclass(frozen=True, slots=True)
class OverlaySettings:
pose: PoseOverlaySettings
bounding_box: BoundingBoxOverlaySettings
skeleton: SkeletonOverlaySettings


@dataclass(frozen=True, slots=True)
class OverlayRenderResult:
frame: np.ndarray
warning: str | None = None


class OverlayRenderer:
"""Render overlays and cache packet-derived skeleton resolution."""

def __init__(self) -> None:
self._skeleton_signature: object | None = None
self._resolved_skeleton: ResolvedSkeleton | None = None
self._last_warning: str | None = None

def clear_runtime_state(self) -> None:
self._skeleton_signature = None
self._resolved_skeleton = None
self._last_warning = None

def render(
self,
frame: np.ndarray,
*,
pose: np.ndarray | None,
packet: PosePacketLike | None,
overlay_settings: OverlaySettings,
offset: tuple[int, int] = (0, 0),
scale: tuple[float, float] = (1.0, 1.0),
) -> OverlayRenderResult:
output = frame.copy()
warning: str | None = None

if overlay_settings.pose.visible and pose is not None:
output = draw_pose(
output,
pose,
p_cutoff=overlay_settings.pose.p_cutoff,
colormap=overlay_settings.pose.colormap,
offset=offset,
scale=scale,
)

if overlay_settings.skeleton.visible and pose is not None:
try:
resolved = self._resolve_packet_skeleton(packet)
except SkeletonResolutionError as exc:
warning = self._new_warning(str(exc))
else:
if resolved is None:
warning = self._new_warning("Skeleton metadata is unavailable for this pose output.")
else:
style = overlay_settings.skeleton.style
colors = None

if style.color_mode == SkeletonColorMode.GRADIENT_KEYPOINTS:
colors = keypoint_colors_bgr(
overlay_settings.pose.colormap,
len(resolved.keypoint_names),
)

result = draw_skeleton(
output,
pose,
resolved,
style,
p_cutoff=overlay_settings.pose.p_cutoff,
offset=offset,
scale=scale,
keypoint_colors=colors,
)

if result.code in {
SkeletonRenderCode.RENDERED,
SkeletonRenderCode.NO_POSE,
}:
self._last_warning = None
else:
warning = self._new_warning(result.message)

if overlay_settings.bounding_box.visible:
output = draw_bbox(
output,
overlay_settings.bounding_box.coordinates,
color_bgr=overlay_settings.bounding_box.color_bgr,
offset=offset,
scale=scale,
)

return OverlayRenderResult(
frame=output,
warning=warning,
)

def _resolve_packet_skeleton(
self,
packet: PosePacketLike | None,
) -> ResolvedSkeleton | None:
if packet is None:
self._skeleton_signature = None
self._resolved_skeleton = None
return None

signature = (
packet.skeleton_id,
tuple(packet.keypoint_names or ()),
tuple(packet.skeleton_edges or ()),
)

if signature == self._skeleton_signature:
return self._resolved_skeleton

resolved = resolve_packet_skeleton(packet)

self._skeleton_signature = signature
self._resolved_skeleton = resolved
return resolved

def _new_warning(
self,
message: str,
) -> str | None:
if not message or message == self._last_warning:
return None

self._last_warning = message
return message
Loading
Loading