from __future__ import annotations
import argparse
import time
import cv2
import depthai as dai
import numpy as np
from pose_backends import (
create_pose_backend
)
from pose3d_utils import (
JOINT_NAMES,
SKELETON,
JointTracker3D,
pixel_to_xyz,
robust_joint_depth,
xyz_to_pixel,
)
# ============================================================
# Configuration
# ============================================================
FRAME_SIZE = (
640,
400
)
CAMERA_FPS = 30
DISPLAY_SCALE = 1.5
PERSON_CONF = 0.30
KEYPOINT_CONF = 0.35
MIN_DEPTH_M = 0.30
MAX_DEPTH_M = 5.00
DEPTH_ROI_RADIUS = 4
MIN_DEPTH_PIXELS = 4
PREVIOUS_DEPTH_GATE_M = 0.40
CENTER_DEPTH_GATE_M = 0.25
FILTER_ALPHA = 0.55
FILTER_BETA = 0.08
MAX_JOINT_SPEED_MPS = 6.0
BASE_POSITION_GATE_M = 0.10
MAX_MISSED_FRAMES = 6
SHOW_XYZ = True
# ============================================================
# Colors
# ============================================================
COLOR_MEASURED = (
0,
255,
0
)
COLOR_PREDICTED = (
0,
200,
255
)
COLOR_SKELETON = (
255,
180,
0
)
COLOR_TEXT = (
255,
255,
255
)
COLOR_BOX = (
255,
0,
255
)
# ============================================================
# Command line
# ============================================================
def parse_args():
parser = argparse.ArgumentParser(
description=(
"OAK-D RGB-D "
"3D pose estimation"
)
)
parser.add_argument(
"--backend",
choices=[
"ultralytics",
"mediapipe"
],
default="ultralytics",
help=(
"Pose estimation backend"
),
)
parser.add_argument(
"--mediapipe-model",
choices=[
"lite",
"full",
"heavy"
],
default="lite",
help=(
"MediaPipe model"
),
)
parser.add_argument(
"--yolo-model",
default=(
"yolo26n-pose.pt"
),
help=(
"Ultralytics pose model"
),
)
parser.add_argument(
"--yolo-imgsz",
type=int,
default=640,
help=(
"Ultralytics input size"
),
)
return parser.parse_args()
# ============================================================
# Trackers
# ============================================================
def create_joint_trackers():
return [
JointTracker3D(
nominal_fps=
CAMERA_FPS,
alpha=
FILTER_ALPHA,
beta=
FILTER_BETA,
max_speed_mps=
MAX_JOINT_SPEED_MPS,
base_gate_m=
BASE_POSITION_GATE_M,
max_missed_frames=
MAX_MISSED_FRAMES,
)
for _ in range(17)
]
# ============================================================
# Draw skeleton
# ============================================================
def draw_pose(
frame,
filtered_xyz,
measurement_accepted,
K,
):
h, w = frame.shape[:2]
display_uv = [
xyz_to_pixel(
xyz,
K
)
for xyz
in filtered_xyz
]
# --------------------------------------------------------
# Skeleton
# --------------------------------------------------------
for a, b in SKELETON:
pa = display_uv[a]
pb = display_uv[b]
if (
pa is None
or pb is None
):
continue
if not (
0 <= pa[0] < w
and
0 <= pa[1] < h
and
0 <= pb[0] < w
and
0 <= pb[1] < h
):
continue
cv2.line(
frame,
pa,
pb,
COLOR_SKELETON,
2,
cv2.LINE_AA,
)
# --------------------------------------------------------
# Joints
# --------------------------------------------------------
for joint_id, xyz in enumerate(
filtered_xyz
):
uv = display_uv[
joint_id
]
if (
xyz is None
or uv is None
):
continue
u, v = uv
if not (
0 <= u < w
and
0 <= v < h
):
continue
if measurement_accepted[
joint_id
]:
joint_color = (
COLOR_MEASURED
)
else:
joint_color = (
COLOR_PREDICTED
)
cv2.circle(
frame,
(u, v),
4,
joint_color,
-1,
cv2.LINE_AA,
)
# ----------------------------------------------------
# XYZ text
# ----------------------------------------------------
if SHOW_XYZ:
x, y, z = xyz
text = (
f"{joint_id}:"
f"{JOINT_NAMES[joint_id]} "
f"({x:+.2f},"
f"{y:+.2f},"
f"{z:.2f})"
)
dy = (
-7
if joint_id % 2 == 0
else 12
)
tx = min(
max(
u + 5,
0
),
max(
0,
w - 230
)
)
ty = min(
max(
v + dy,
10
),
h - 5
)
cv2.putText(
frame,
text,
(tx, ty),
cv2.FONT_HERSHEY_SIMPLEX,
0.30,
COLOR_TEXT,
1,
cv2.LINE_AA,
)
# ============================================================
# Main
# ============================================================
def main():
args = parse_args()
# --------------------------------------------------------
# Pose backend
# --------------------------------------------------------
pose_backend = (
create_pose_backend(
args.backend,
ultralytics_model=
args.yolo_model,
ultralytics_imgsz=
args.yolo_imgsz,
person_conf=
PERSON_CONF,
mediapipe_model=
args.mediapipe_model,
)
)
joint_trackers = (
create_joint_trackers()
)
try:
with dai.Pipeline() as pipeline:
# =================================================
# RGB camera
# =================================================
color = pipeline.create(
dai.node.Camera
).build(
dai.CameraBoardSocket.CAM_A,
sensorFps=
CAMERA_FPS,
)
# =================================================
# Stereo cameras
# =================================================
left = pipeline.create(
dai.node.Camera
).build(
dai.CameraBoardSocket.CAM_B,
sensorFps=
CAMERA_FPS,
)
right = pipeline.create(
dai.node.Camera
).build(
dai.CameraBoardSocket.CAM_C,
sensorFps=
CAMERA_FPS,
)
# =================================================
# StereoDepth
# =================================================
stereo = pipeline.create(
dai.node.StereoDepth
)
stereo.setDefaultProfilePreset(
dai.node.StereoDepth
.PresetMode.DEFAULT
)
stereo.setRectifyEdgeFillColor(
0
)
stereo.enableDistortionCorrection(
True
)
left.requestOutput(
FRAME_SIZE
).link(
stereo.left
)
right.requestOutput(
FRAME_SIZE
).link(
stereo.right
)
# =================================================
# RGBD
# =================================================
rgbd = pipeline.create(
dai.node.RGBD
).build(
color,
stereo,
FRAME_SIZE,
CAMERA_FPS,
)
rgbd_queue = (
rgbd.rgbd
.createOutputQueue(
maxSize=2,
blocking=False,
)
)
pipeline.start()
print()
print(
"OAK-D 3D Pose started"
)
print(
f"Pose backend: "
f"{pose_backend.name}"
)
print(
"Q / ESC : quit"
)
print()
intrinsic_printed = False
fps_value = 0.0
last_loop_time = (
time.perf_counter()
)
frame_counter = 0
# =================================================
# Main loop
# =================================================
while pipeline.isRunning():
# --------------------------------------------
# Synchronized RGB-D
# --------------------------------------------
rgbd_data = (
rgbd_queue.get()
)
rgb_msg = (
rgbd_data
.getRGBFrame()
)
depth_msg = (
rgbd_data
.getDepthFrame()
)
if (
rgb_msg is None
or depth_msg is None
):
continue
frame = (
rgb_msg.getCvFrame()
)
depth_mm = (
depth_msg.getCvFrame()
)
if (
depth_mm.shape[:2]
!= frame.shape[:2]
):
raise RuntimeError(
"RGB/Depth size mismatch: "
f"RGB={frame.shape[:2]}, "
f"Depth={depth_mm.shape[:2]}"
)
h, w = frame.shape[:2]
# --------------------------------------------
# RGB intrinsic matrix
# --------------------------------------------
K = np.asarray(
rgb_msg
.getTransformation()
.getIntrinsicMatrix(),
dtype=np.float64,
)
if not intrinsic_printed:
print(
f"RGB size: "
f"{w} x {h}"
)
print(
"Intrinsic matrix:"
)
print(K)
intrinsic_printed = True
# =============================================
# Pose inference
# =============================================
pose = (
pose_backend.infer(
frame
)
)
measurements = [
None
] * 17
now = (
time.perf_counter()
)
# =============================================
# 2D Pose -> Depth -> XYZ
# =============================================
if pose is not None:
# ----------------------------------------
# Bounding box
# ----------------------------------------
if pose.bbox is not None:
box = (
pose.bbox
.astype(int)
)
cv2.rectangle(
frame,
(
box[0],
box[1]
),
(
box[2],
box[3]
),
COLOR_BOX,
1,
)
# ----------------------------------------
# 17 joints
# ----------------------------------------
for joint_id in range(
17
):
conf = float(
pose.conf[
joint_id
]
)
if (
conf
< KEYPOINT_CONF
):
continue
u = float(
pose.xy[
joint_id,
0
]
)
v = float(
pose.xy[
joint_id,
1
]
)
if not (
0 <= u < w
and
0 <= v < h
):
continue
previous_z = (
joint_trackers[
joint_id
]
.get_predicted_z()
)
# ------------------------------------
# Robust ROI depth
# ------------------------------------
z_m = robust_joint_depth(
depth_mm,
u,
v,
previous_z_m=
previous_z,
min_depth_m=
MIN_DEPTH_M,
max_depth_m=
MAX_DEPTH_M,
roi_radius=
DEPTH_ROI_RADIUS,
min_pixels=
MIN_DEPTH_PIXELS,
previous_depth_gate_m=
PREVIOUS_DEPTH_GATE_M,
center_depth_gate_m=
CENTER_DEPTH_GATE_M,
)
if z_m is None:
continue
# ------------------------------------
# Pixel -> XYZ
# ------------------------------------
measurements[
joint_id
] = pixel_to_xyz(
u,
v,
z_m,
K,
)
# =============================================
# Temporal 3D filtering
# =============================================
filtered_xyz = [
None
] * 17
measurement_accepted = [
False
] * 17
for joint_id in range(
17
):
xyz, accepted = (
joint_trackers[
joint_id
].update(
measurements[
joint_id
],
now,
)
)
filtered_xyz[
joint_id
] = xyz
measurement_accepted[
joint_id
] = accepted
# =============================================
# Draw pose
# =============================================
draw_pose(
frame,
filtered_xyz,
measurement_accepted,
K,
)
# =============================================
# FPS
# =============================================
loop_now = (
time.perf_counter()
)
dt_loop = (
loop_now
- last_loop_time
)
last_loop_time = (
loop_now
)
if dt_loop > 0:
instant_fps = (
1.0 / dt_loop
)
if fps_value == 0.0:
fps_value = (
instant_fps
)
else:
fps_value = (
0.9 * fps_value
+
0.1 * instant_fps
)
# =============================================
# Status
# =============================================
if (
args.backend
== "mediapipe"
):
backend_detail = (
args.mediapipe_model
)
else:
backend_detail = (
args.yolo_model
)
status = (
f"{pose_backend.name}/"
f"{backend_detail} "
f"FPS {fps_value:.1f} "
f"Pose "
f"{pose_backend.last_inference_ms:.1f} ms "
f"XYZ[m]: "
f"X=right Y=down Z=forward"
)
cv2.rectangle(
frame,
(0, 0),
(w, 24),
(0, 0, 0),
-1,
)
cv2.putText(
frame,
status,
(8, 17),
cv2.FONT_HERSHEY_SIMPLEX,
0.40,
(
255,
255,
255
),
1,
cv2.LINE_AA,
)
# =============================================
# Display
# =============================================
if (
DISPLAY_SCALE
!= 1.0
):
display_frame = (
cv2.resize(
frame,
None,
fx=
DISPLAY_SCALE,
fy=
DISPLAY_SCALE,
interpolation=
cv2.INTER_LINEAR,
)
)
else:
display_frame = (
frame
)
cv2.imshow(
"OAK-D 3D Pose",
display_frame,
)
# =============================================
# Benchmark output
# =============================================
frame_counter += 1
if (
frame_counter
% 60
== 0
):
print(
f"{pose_backend.name}: "
f"pose="
f"{pose_backend.last_inference_ms:.1f} ms, "
f"total="
f"{fps_value:.1f} fps"
)
# =============================================
# Keyboard
# =============================================
key = (
cv2.waitKey(1)
& 0xFF
)
if (
key == ord("q")
or key == 27
):
break
pipeline.stop()
finally:
pose_backend.close()
cv2.destroyAllWindows()
if __name__ == "__main__":
main()