car/arm_plane_calib/predict_plane_cmd.py
2026-08-13 14:35:17 +08:00

152 lines
No EOL
4.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
import yaml
import cv2
import numpy as np
import rclpy
from rclpy.node import Node
from sensor_msgs.msg import Image, CameraInfo
from cv_bridge import CvBridge
def features(x, y):
return np.array([1.0, x, y, x * x, x * y, y * y], dtype=np.float64)
class PlaneCmdPredictor(Node):
def __init__(self):
super().__init__('plane_cmd_predictor')
self.declare_parameter('color_topic', '/camera/color/image_raw')
self.declare_parameter('depth_topic', '/camera/depth/image_raw')
self.declare_parameter('camera_info_topic', '/camera/depth/camera_info')
self.declare_parameter('model_path', 'model.yaml')
self.declare_parameter('depth_scale', 0.001)
self.color_topic = self.get_parameter('color_topic').value
self.depth_topic = self.get_parameter('depth_topic').value
self.camera_info_topic = self.get_parameter('camera_info_topic').value
self.model_path = self.get_parameter('model_path').value
self.depth_scale = float(self.get_parameter('depth_scale').value)
with open(self.model_path, 'r', encoding='utf-8') as f:
model = yaml.safe_load(f)
self.W = np.array(model['weights'], dtype=np.float64) # [6, cmd_dim]
self.bridge = CvBridge()
self.color_img = None
self.depth_img = None
self.K = None
# 你的模型输出顺序固定为:6号、5号、4号、3号
self.model_servo_order = [6, 5, 4, 3]
# 你想打印成:3号、4号、5号、6号
self.print_servo_order = [3, 4, 5, 6]
# 当前只处理 3/4/5/6 号舵机,范围都一致
self.servo_min = 125
self.servo_max = 875
self.create_subscription(Image, self.color_topic, self.on_color, 10)
self.create_subscription(Image, self.depth_topic, self.on_depth, 10)
self.create_subscription(CameraInfo, self.camera_info_topic, self.on_info, 10)
self.window_name = 'predict_plane_cmd'
cv2.namedWindow(self.window_name, cv2.WINDOW_NORMAL)
cv2.setMouseCallback(self.window_name, self.on_mouse)
def on_color(self, msg):
self.color_img = self.bridge.imgmsg_to_cv2(msg, desired_encoding='bgr8')
def on_depth(self, msg):
self.depth_img = self.bridge.imgmsg_to_cv2(msg, desired_encoding='passthrough')
def on_info(self, msg):
self.K = np.array(msg.k, dtype=np.float64).reshape(3, 3)
def pixel_to_xyz(self, u, v):
if self.depth_img is None or self.K is None:
return None
h, w = self.depth_img.shape[:2]
if not (0 <= u < w and 0 <= v < h):
return None
patch = self.depth_img[max(0, v - 2):min(h, v + 3), max(0, u - 2):min(w, u + 3)].astype(np.float32)
valid = patch[patch > 0]
if valid.size == 0:
return None
z = np.median(valid) * self.depth_scale
fx = self.K[0, 0]
fy = self.K[1, 1]
cx = self.K[0, 2]
cy = self.K[1, 2]
x = (u - cx) * z / fx
y = (v - cy) * z / fy
return float(x), float(y), float(z)
def clamp_servo(self, value):
value = int(round(value))
if value < self.servo_min:
value = self.servo_min
if value > self.servo_max:
value = self.servo_max
return value
def on_mouse(self, event, x, y, flags, param):
if event != cv2.EVENT_LBUTTONDOWN:
return
xyz = self.pixel_to_xyz(x, y)
if xyz is None:
print('[WARN] 无有效深度')
return
X, Y, Z = xyz
# 模型输出顺序:6,5,4,3
cmd_raw = features(X, Y) @ self.W
cmd_raw = cmd_raw.tolist()
# 先取整并限幅
cmd_int = [self.clamp_servo(v) for v in cmd_raw]
# 转成 {舵机ID: 值}
servo_map = {}
for sid, val in zip(self.model_servo_order, cmd_int):
servo_map[sid] = val
# 按你希望的顺序打印:3,4,5,6
ordered_pairs = [(sid, servo_map[sid]) for sid in self.print_servo_order]
print(f'\n点击像素: ({x}, {y})')
print(f'相机坐标: X={X:.4f}, Y={Y:.4f}, Z={Z:.4f}')
print('预测机械臂命令向量(按模型顺序 6,5,4,3):',
', '.join([str(servo_map[sid]) for sid in self.model_servo_order]))
print('舵机列表:', ','.join([f'{sid}:{val}' for sid, val in ordered_pairs]))
def spin_loop(self):
while rclpy.ok():
rclpy.spin_once(self, timeout_sec=0.03)
if self.color_img is not None:
cv2.imshow(self.window_name, self.color_img)
if (cv2.waitKey(1) & 0xFF) == ord('q'):
break
cv2.destroyAllWindows()
def main():
rclpy.init()
node = PlaneCmdPredictor()
try:
node.spin_loop()
finally:
node.destroy_node()
rclpy.shutdown()
if __name__ == '__main__':
main()