内容主要是获取目标位姿,步骤如下:

1.获取相机的数据,然后通过相机内参将像素坐标转换为相机坐标;

2.yolov8的分割模型获取抓取物体的烟码;

3.把深度图的深度值和分割的彩图值结合构成3维数据;

4.foundationpose对3维数据找位姿。

1.分割环境搭建

项目用conda搭建环境。

系统:ubuntu22.04  ROS2  cuda12.3

相机环境搭建参考https://blog.csdn.net/weixin_71719718/article/details/160549145?spm=1011.2124.3001.6209

一、基础环境说明

  1. 已安装:ROS2 HumbleMiniconda
  2. 显卡:RTX 4060 Laptop
  3. 系统 CUDA:12.3(nvcc)
  4. 核心原则:所有 CUDA 相关依赖必须统一为 12.3版本

二、Conda 环境创建与基础配置

# 1. 创建conda环境(Python3.10)
conda create -n ros2_tensorrt python=3.10 -y
conda activate ros2_tensorrt

# 2. 基础pip升级
pip install --upgrade pip setuptools wheel ninja

三、核心 CUDA/PyTorch 安装(关键版本)

# 卸载冲突包
pip uninstall -y torch torchvision torchaudio triton

# 安装指定版本(CUDA12.4 匹配系统)
python -m pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu124

# 验证(必须输出 12.4)
python -c "import torch; print(torch.version.cuda, torch.cuda.is_available())"

四、NumPy 降级(修复兼容性问题)

pip install numpy==1.24.4

五、一键安装 YOLO 分割环境

numpy==1.24.0
opencv-python
open3d 
pyyaml
scikit-learn==1.2.2
ultralytics==8.2.100
numpy==1.24.0

六、分割代码ros2_sender.py

#!/usr/bin/env python3
import os
import sys
import json
import base64
import time
import threading
import socket
import struct  # ✅ 添加缺失的导入
import rclpy
import cv2
import numpy as np
import torch
from rclpy.node import Node
from rclpy.action import ActionServer, GoalResponse, CancelResponse
from rclpy.action.server import ServerGoalHandle
from rclpy.executors import MultiThreadedExecutor
from rclpy.callback_groups import ReentrantCallbackGroup
from rclpy.qos import QoSProfile
from sensor_msgs.msg import Image, CameraInfo, PointCloud2, PointField
from geometry_msgs.msg import PoseStamped, Pose, Point, Quaternion
from std_msgs.msg import Header
from ultralytics import YOLO
from scipy.spatial.transform import Rotation
from vision_detection_action.action import VisionDetection


# 屏蔽无关警告
os.environ["QT_LOGGING_RULES"] = "qt.fonts.warning=false"
os.environ["OPENCV_LOG_LEVEL"] = "FATAL"
os.environ["CV_LOG_LEVEL"] = "FATAL"

# 全局QoS配置
QOS_PROFILE = QoSProfile(depth=10)


# ==========================
# 图像解析工具
# ==========================
def imgmsg_to_bgr8(msg: Image) -> np.ndarray:
    height = msg.height
    width = msg.width
    step = msg.step
    encoding = msg.encoding

    img_data = np.frombuffer(msg.data, dtype=np.uint8)
    img_data = img_data.reshape(height, step)
    valid_width = width * 3
    img = img_data[:, :valid_width].reshape(height, width, 3)

    if encoding == "rgb8":
        img = img[:, :, [2, 1, 0]]
    elif encoding != "bgr8":
        raise RuntimeError(f"仅支持 bgr8 / rgb8 格式, 当前格式: {encoding}")
    return img


def imgmsg_to_depth_16uc1(msg: Image) -> np.ndarray:
    height = msg.height
    width = msg.width
    step = msg.step
    encoding = msg.encoding

    if encoding != "16UC1":
        raise RuntimeError(f"仅支持 16UC1 格式, 当前格式: {encoding}")

    img_data = np.frombuffer(msg.data, dtype=np.uint16)
    img_data = img_data.reshape(height, step // 2)
    return img_data[:, :width]


def cv2_to_imgmsg(cv_img, encoding="bgr8"):
    """将OpenCV图像转换为ROS Image消息"""
    if encoding == "bgr8":
        img_msg = Image()
        img_msg.height = cv_img.shape[0]
        img_msg.width = cv_img.shape[1]
        img_msg.encoding = "bgr8"
        img_msg.is_bigendian = False
        img_msg.step = cv_img.shape[1] * 3
        img_msg.data = cv_img.tobytes()
        return img_msg
    else:
        raise NotImplementedError(f"不支持的编码格式: {encoding}")


# ==========================
# 自定义 Action 类型
# ==========================
#from ros2_vision_action.action import DetectionPoseCloud
#from vision_detection_action.action import VisionDetection

# ==========================
# 全局配置
# ==========================
RGB_TOPIC = '/camera/color/image_raw'
DEPTH_TOPIC = '/camera/depth/image_raw'
CAM_INFO_TOPIC = '/camera/color/camera_info'

YOLO_MODEL_PATH = "/home/wyq/ros2_ws/best_v3.pt"
YOLO_DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
YOLO_CONF = 0.5

DEPTH_MIN = 0.3
DEPTH_MAX = 3.0
MIN_MASK_PIXELS = 10

PLY_FILENAME = "phone.ply"
voxel_size = 0.005
eps = 0.015
min_samples = 5

os.chdir("/home/wyq/ros2_ws/sam2-main")
SAM2_CFG_PATH = "configs/sam2.1/sam2.1_hiera_b+.yaml"
SAM2_CKPT = "weights/sam2.1_hiera_base_plus.pt"
SAM2_DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

# 帧率控制
RUN_INTERVAL = 1.0

# ==========================
# FoundationPose TCP配置
# ==========================
FOUNDATIONPOSE_CONFIG = {
    "host": "127.0.0.1",
    "port": 12345,
    "input_width": 640,
    "input_height": 400,
    "timeout": 10.0,
}

# ==========================
# 加载 SAM2
# ==========================
try:
    from sam2.build_sam import build_sam2
    from sam2.sam2_image_predictor import SAM2ImagePredictor

    sam2_model = build_sam2(SAM2_CFG_PATH, SAM2_CKPT, device=SAM2_DEVICE)
    sam2_predictor = SAM2ImagePredictor(sam2_model)
    print("✅ SAM2 模型加载成功!")
except Exception as e:
    print(f"❌ SAM2 模型加载失败: {e}")
    sys.exit(1)


# ==========================
# 工具函数
# ==========================
def yolo_sam_refine_segment(rgb_img, yolo_model, conf_thresh=0.3):
    H, W = rgb_img.shape[:2]
    img_rgb = cv2.cvtColor(rgb_img, cv2.COLOR_BGR2RGB)

    results = yolo_model(rgb_img, conf=conf_thresh)[0]
    if results.boxes is None or len(results.boxes) == 0:
        print("[YOLO状态] ❌ 未检测到目标")
        return None, None

    bbox = results.boxes[0].xyxy.cpu().numpy()[0]
    x1, y1, x2, y2 = map(int, bbox)
    x1 = max(0, x1)
    y1 = max(0, y1)
    x2 = min(W, x2)
    y2 = min(H, y2)

    if x2 - x1 < 10 or y2 - y1 < 10:
        print("[YOLO状态] ❌ 检测框尺寸过小,无效")
        return None, None
    print("[YOLO状态] ✅ 检测到目标")

    # SAM2 分割
    sam2_predictor.set_image(img_rgb)
    masks, scores, _ = sam2_predictor.predict(box=bbox, multimask_output=True)
    best_mask = masks[np.argmax(scores)]
    refined_mask = (best_mask > 0).astype(np.uint8) * 255
    print("[SAM2状态] ✅ 图像分割完成")
    return refined_mask, (x1, y1, x2, y2)


def draw_detection_result(rgb_img, mask, bbox, pose=None):
    """绘制检测结果图片"""
    img_copy = rgb_img.copy()

    # 绘制掩码(半透明绿色覆盖)
    if mask is not None:
        # 确保mask是二值图像
        if len(mask.shape) == 3:
            mask = mask[:, :, 0]
        mask_colored = np.zeros_like(img_copy)
        mask_colored[:, :, 1] = mask  # 绿色通道
        alpha = 0.3
        img_copy = cv2.addWeighted(img_copy, 1 - alpha, mask_colored, alpha, 0)

    # 绘制边界框
    if bbox is not None:
        x1, y1, x2, y2 = bbox
        cv2.rectangle(img_copy, (x1, y1), (x2, y2), (0, 255, 0), 2)
        cv2.putText(img_copy, "Target", (x1, y1 - 10),
                    cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2)

    # 如果有点姿信息,绘制位姿信息
    if pose is not None:
        h, w = img_copy.shape[:2]
        # 显示位姿矩阵的前3行
        if isinstance(pose, np.ndarray) and pose.shape == (4, 4):
            pos_text = f"Pos: ({pose[0,3]:.3f}, {pose[1,3]:.3f}, {pose[2,3]:.3f})"
            cv2.putText(img_copy, pos_text, (10, h - 30),
                        cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 255), 1)

    # 添加信息文字
    h, w = img_copy.shape[:2]
    cv2.putText(img_copy, "Detection Result", (10, 30),
                cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 255, 0), 2)
    cv2.putText(img_copy, f"Size: {w}x{h}", (10, 60),
                cv2.FONT_HERSHEY_SIMPLEX, 0.6, (255, 255, 255), 1)

    # ✅ 修复:确保图像正确显示
    try:
        # 确保图像数据有效
        if img_copy is not None and img_copy.size > 0:
            # 创建或更新窗口
            window_name = "Segmentation Result"
            cv2.namedWindow(window_name, cv2.WINDOW_NORMAL)
            # 调整窗口大小以适应屏幕
            cv2.resizeWindow(window_name, 800, 600)
            cv2.imshow(window_name, img_copy)
            cv2.waitKey(1)  # 1ms延迟,允许窗口刷新
        else:
            print("[显示状态] ❌ 图像数据无效")
    except Exception as e:
        print(f"[显示状态] ❌ 显示失败: {e}")

    return img_copy


def fast_denoise_point_clouds(points, colors=None):
    if len(points) == 0:
        return points, colors
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    pts = torch.tensor(points, device=device, dtype=torch.float32)
    cols = torch.tensor(colors, device=device) if colors is not None else None

    coords = torch.floor(pts / voxel_size).to(torch.int32)
    unique_coords, inverse = torch.unique(coords, dim=0, return_inverse=True)
    idx = torch.unique(inverse)
    pts = pts[idx]
    if cols is not None:
        cols = cols[idx]

    try:
        dist = torch.cdist(pts, pts)
        count = (dist < eps).sum(dim=1)
        mask = count >= min_samples
        pts = pts[mask]
        if cols is not None:
            cols = cols[mask]
    except Exception as e:
        print(f"点云滤波异常: {e}")

    pts_np = pts.cpu().numpy()
    cols_np = cols.cpu().numpy() if cols is not None else None
    return pts_np, cols_np


def save_point_cloud_to_ply(points_3d, colors, filename=PLY_FILENAME):
    if len(points_3d) == 0:
        print("[点云状态] ❌ 无有效点云,跳过保存")
        return False
    try:
        with open(filename, 'w') as f:
            f.write("ply\nformat ascii 1.0\n")
            f.write(f"element vertex {len(points_3d)}\n")
            f.write("property float x\nproperty float y\nproperty float z\n")
            f.write("property uchar red\nproperty uchar green\nproperty uchar blue\n")
            f.write("end_header\n")
            for (x, y, z), (r, g, b) in zip(points_3d, colors):
                f.write(f"{x:.4f} {y:.4f} {z:.4f} {int(r)} {int(g)} {int(b)}\n")
        print(f"[点云状态] ✅ PLY文件保存成功,总点数: {len(points_3d)}")
        return True
    except Exception as e:
        print(f"[点云状态] ❌ PLY保存失败: {e}")
        return False


def create_point_cloud(points_3d, colors, frame_id, clock):
    """创建PointCloud2消息"""
    if len(points_3d) == 0:
        return PointCloud2()

    cloud_msg = PointCloud2()
    cloud_msg.header = Header()
    cloud_msg.header.stamp = clock.now().to_msg()
    cloud_msg.header.frame_id = frame_id

    cloud_msg.height = 1
    cloud_msg.width = len(points_3d)
    cloud_msg.is_bigendian = False
    cloud_msg.is_dense = True

    # 定义字段
    cloud_msg.fields = [
        PointField(name='x', offset=0, datatype=PointField.FLOAT32, count=1),
        PointField(name='y', offset=4, datatype=PointField.FLOAT32, count=1),
        PointField(name='z', offset=8, datatype=PointField.FLOAT32, count=1),
        PointField(name='rgb', offset=12, datatype=PointField.UINT32, count=1),
    ]

    cloud_msg.point_step = 16  # 4 * 4 bytes
    cloud_msg.row_step = cloud_msg.point_step * len(points_3d)

    # 序列化数据
    data = []
    for pt, col in zip(points_3d, colors):
        # 将RGB打包为uint32
        rgb = (int(col[2]) << 16) | (int(col[1]) << 8) | int(col[0])
        data.append(struct.pack('ffff', pt[0], pt[1], pt[2], float(rgb)))
    cloud_msg.data = b''.join(data)

    return cloud_msg


def send_to_foundationpose(rgb, depth, mask, K, config):
    """发送数据到FoundationPose服务器并接收位姿"""
    target_w = config["input_width"]  # 640
    target_h = config["input_height"]  # 400

    # 调整图像尺寸
    orig_h, orig_w = rgb.shape[:2]

    # ✅ 确保depth是有效的
    if depth is None or depth.size == 0:
        print("[FoundationPose] ❌ depth数据为空")
        return None

    # ✅ 确保depth是float32类型
    if depth.dtype != np.float32:
        depth = depth.astype(np.float32)

    # 调整尺寸
    rgb_resized = cv2.resize(rgb, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
    depth_resized = cv2.resize(depth, (target_w, target_h), interpolation=cv2.INTER_NEAREST)
    mask_resized = cv2.resize(mask, (target_w, target_h), interpolation=cv2.INTER_NEAREST)

    # ✅ 确保数据类型正确
    rgb_resized = rgb_resized.astype(np.uint8)
    depth_resized = depth_resized.astype(np.float32)  # depth保持float32
    mask_resized = mask_resized.astype(np.uint8)

    print(f"[调试] 原始尺寸: {orig_w}x{orig_h}, 目标尺寸: {target_w}x{target_h}")
    print(f"[调试] rgb_resized shape: {rgb_resized.shape}, dtype: {rgb_resized.dtype}")
    print(f"[调试] depth_resized shape: {depth_resized.shape}, dtype: {depth_resized.dtype}")
    print(f"[调试] mask_resized shape: {mask_resized.shape}, dtype: {mask_resized.dtype}")

    # ✅ 验证数据大小
    rgb_bytes = rgb_resized.tobytes()
    depth_bytes = depth_resized.tobytes()
    mask_bytes = mask_resized.tobytes()

    print(f"[调试] rgb数据大小: {len(rgb_bytes)} bytes, 期望: {target_h * target_w * 3}")
    print(f"[调试] depth数据大小: {len(depth_bytes)} bytes, 期望: {target_h * target_w * 4} (float32)")
    print(f"[调试] mask数据大小: {len(mask_bytes)} bytes, 期望: {target_h * target_w}")

    # ✅ 检查数据大小是否正确
    expected_rgb_size = target_h * target_w * 3
    expected_depth_size = target_h * target_w * 4  # float32 = 4 bytes
    expected_mask_size = target_h * target_w

    if len(rgb_bytes) != expected_rgb_size:
        print(f"[错误] rgb数据大小错误: {len(rgb_bytes)} != {expected_rgb_size}")
        return None
    if len(depth_bytes) != expected_depth_size:
        print(f"[错误] depth数据大小错误: {len(depth_bytes)} != {expected_depth_size}")
        return None
    if len(mask_bytes) != expected_mask_size:
        print(f"[错误] mask数据大小错误: {len(mask_bytes)} != {expected_mask_size}")
        return None

    # 调整相机内参
    K_resized = K.copy()
    scale_x = target_w / orig_w
    scale_y = target_h / orig_h
    K_resized[0, 0] *= scale_x
    K_resized[0, 2] *= scale_x
    K_resized[1, 1] *= scale_y
    K_resized[1, 2] *= scale_y

    # ✅ 编码数据 - 使用正确的dtype
    data = {
        'rgb_b64': base64.b64encode(rgb_bytes).decode(),
        'rgb_shape': [target_h, target_w, 3],
        'depth_b64': base64.b64encode(depth_bytes).decode(),
        'depth_shape': [target_h, target_w],
        'mask_b64': base64.b64encode(mask_bytes).decode(),
        'mask_shape': [target_h, target_w],
        'K': K_resized.tolist(),
    }

    # 发送并接收
    try:
        sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        sock.settimeout(config["timeout"])
        sock.connect((config["host"], config["port"]))
        print("[FoundationPose] ✅ 连接成功")

        message = json.dumps(data) + '\n'
        sock.sendall(message.encode('utf-8'))
        print("[FoundationPose] 📤 数据已发送")

        response = sock.recv(65536).decode('utf-8')
        sock.close()

        lines = response.strip().split('\n')
        for line in lines:
            if line:
                resp_data = json.loads(line)
                if 'pose' in resp_data:
                    pose_matrix = np.array(resp_data['pose'])
                    print("[FoundationPose] ✅ 收到6D位姿")
                    return pose_matrix
        return None
    except socket.timeout:
        print("[FoundationPose] ❌ 连接超时")
        return None
    except ConnectionRefusedError:
        print("[FoundationPose] ❌ 连接被拒绝,请确保FoundationPose服务器已启动")
        return None
    except Exception as e:
        print(f"[FoundationPose] ❌ 通信失败: {e}")
        return None


# ==========================
# 主节点
# ==========================
class DetectionPoseCloudServer(Node):
    def __init__(self):
        super().__init__('detection_pose_cloud_server')
        self.latest_points_3d = None
        self.latest_colors = None
        self.latest_pose_matrix = None
        self.latest_bbox = None
        self.rgb_img = None
        self.depth_img = None
        self.K = None
        self.rgb_err_cnt = 0
        self.depth_err_cnt = 0
        self.callback_group = ReentrantCallbackGroup()
        self.last_run_ts = 0.0
        self.latest_detection_img = None
        self.pipeline_lock = threading.Lock()

        # 加载 YOLO
        self.yolo_model = YOLO(YOLO_MODEL_PATH)
        self.yolo_model.to(YOLO_DEVICE)

        # 订阅话题
        self.create_subscription(
            Image,
            RGB_TOPIC,
            self.rgb_callback,
            QOS_PROFILE,
            callback_group=self.callback_group
        )
        self.create_subscription(
            Image,
            DEPTH_TOPIC,
            self.depth_callback,
            QOS_PROFILE,
            callback_group=self.callback_group
        )
        self.create_subscription(
            CameraInfo,
            CAM_INFO_TOPIC,
            self.caminfo_callback,
            QOS_PROFILE,
            callback_group=self.callback_group
        )

        # 发布检测结果图片
        self.result_image_pub = self.create_publisher(
            Image,
            '/vision/detection_result_image',
            QOS_PROFILE
        )

        # ✅ Action Server - 只创建一次
        self.action_server = ActionServer(
            node=self,
            action_name='/vision/detection_pose_cloud',
            action_type=VisionDetection,
            execute_callback=self.execute_callback,
            goal_callback=self.goal_callback,
            cancel_callback=self.cancel_callback,
        )

        print("=" * 60)
        print("✅ Action Server 创建信息:")
        print(f"   Action 名称: /vision/detection_pose_cloud")
        print(f"   execute_callback: {self.execute_callback}")
        print("=" * 60)
        self.get_logger().info("✅ 节点初始化完成")

    # def goal_callback(self, goal_request: VisionDetection.Goal) -> GoalResponse:
    #     print("=" * 60)
    #     print("🔵🔵🔵 GOAL_CALLBACK 被调用!🔵🔵🔵")
    #     print("=" * 60)
    #     self.get_logger().info(f"📥 收到目标: {goal_request.target_obj_name}")
    #     return GoalResponse.ACCEPT
    def goal_callback(self, goal_request: VisionDetection.Goal) -> GoalResponse:
        # 只拒绝负数
        if goal_request.priority < 0:
            self.get_logger().warn(f"Invalid priority: {goal_request.priority}. Rejecting goal.")
            return GoalResponse.REJECT

        # 或者完全移除优先级检查
        self.get_logger().info(f"✅ Accepting goal with priority: {goal_request.priority}")
        return GoalResponse.ACCEPT

    def cancel_callback(self, goal_handle: ServerGoalHandle) -> CancelResponse:
        self.get_logger().info("❌ 任务取消")
        return CancelResponse.ACCEPT

    def run_pipeline(self, force=False):

        with self.pipeline_lock:
            if self.rgb_img is None or self.depth_img is None or self.K is None:
                self.get_logger().warn("相机数据不全,跳过计算")
                return False, None, None, None, None

            # 只在非强制模式下检查帧率
            if not force:
                now = time.time()
                if now - self.last_run_ts < RUN_INTERVAL:
                    return False, None, None, None, None

            # 更新最后执行时间
            self.last_run_ts = time.time()
        # YOLO + SAM 分割
        mask, bbox = yolo_sam_refine_segment(self.rgb_img, self.yolo_model, YOLO_CONF)
        if mask is None:
            return False, None, None, None, None

        # 点云计算 - ✅ 修复相机内参提取
        depth = self.depth_img.astype(np.float32) / 1000.0
        depth = np.clip(depth, DEPTH_MIN, DEPTH_MAX)
        h, w = self.rgb_img.shape[:2]

        # ✅ 正确提取相机内参
        fx = self.K[0, 0]  # X方向焦距
        fy = self.K[1, 1]  # Y方向焦距
        cx = self.K[0, 2]  # 主点X
        cy = self.K[1, 2]  # 主点Y

        ys, xs = np.where(mask > 0)

        if len(ys) < MIN_MASK_PIXELS:
            print(f"[点云状态] ❌ 有效掩码像素不足,当前数量: {len(ys)}")
            return False, None, None, None, None

        # 生成原始点云
        points_3d = []
        colors = []
        for v, u in zip(ys, xs):
            z = depth[v, u]
            if z <= 0 or z > DEPTH_MAX:
                continue
            x = (u - cx) * z / fx
            y = (v - cy) * z / fy
            b, g, r = self.rgb_img[v, u]
            points_3d.append([x, y, z])
            colors.append([r, g, b])

        if len(points_3d) == 0:
            print("[点云状态] ❌ 无有效点云数据")
            return False, None, None, None, None

        points_3d = np.array(points_3d, dtype=np.float32)
        colors = np.array(colors, dtype=np.uint8)
        print(f"[点云状态] ✅ 原始点云生成完成,点数: {len(points_3d)}")

        # 点云去噪
        points_3d, colors = fast_denoise_point_clouds(points_3d, colors)
        print(f"[点云状态] ✅ 点云去噪完成,剩余点数: {len(points_3d)}")

        # 保存PLY文件
        save_point_cloud_to_ply(points_3d, colors)

        # ============================================================
        # 调用 FoundationPose 获取6D位姿
        # ============================================================
        pose_matrix = None
        try:
            print("[FoundationPose] 📡 发送数据到FoundationPose服务器...")
            pose_matrix = send_to_foundationpose(
                self.rgb_img,
                self.depth_img,
                mask,
                self.K,
                FOUNDATIONPOSE_CONFIG
            )
            if pose_matrix is not None:
                print(f"[FoundationPose] ✅ 收到位姿矩阵: {pose_matrix[:3, :4]}")
            else:
                print("[FoundationPose] ❌ 未收到位姿数据")
        except Exception as e:
            print(f"[FoundationPose] ❌ 调用失败: {e}")

        # 绘制检测结果图片(包含位姿信息)
        result_img = draw_detection_result(self.rgb_img, mask, bbox, pose_matrix)
        # ✅ 保存结果图像到文件(用于调试)
        cv2.imwrite("/tmp/debug_result.jpg", result_img)
        print("[调试] 结果图像已保存到 /tmp/debug_result.jpg")
        self.latest_detection_img = result_img

        # 发布检测结果图片
        img_msg = cv2_to_imgmsg(result_img, "bgr8")
        img_msg.header.stamp = self.get_clock().now().to_msg()
        img_msg.header.frame_id = "camera_link"
        self.result_image_pub.publish(img_msg)
        self.get_logger().info("📷 检测结果图片已发布")
        # ✅ 保存到缓存
        self.latest_points_3d = points_3d
        self.latest_colors = colors
        self.latest_pose_matrix = pose_matrix
        self.latest_bbox = bbox
        return True, points_3d, colors, pose_matrix, bbox

    def rgb_callback(self, msg: Image):
        try:
            self.rgb_img = imgmsg_to_bgr8(msg)
            now = time.time()
            if now - self.last_run_ts > RUN_INTERVAL:
                # ✅ 直接调用,不要用线程(避免阻塞Action Server)
                self.run_pipeline(force=False)
        except Exception as e:
            self.rgb_err_cnt += 1
            if self.rgb_err_cnt % 10 == 0:
                import traceback
                self.get_logger().error(f"RGB解析错误:\n{traceback.format_exc()}")
    def depth_callback(self, msg: Image):
        try:
            self.depth_img = imgmsg_to_depth_16uc1(msg)
        except Exception as e:
            self.depth_err_cnt += 1
            if self.depth_err_cnt % 10 == 0:
                import traceback
                self.get_logger().error(f"深度解析错误:\n{traceback.format_exc()}")

    def caminfo_callback(self, msg: CameraInfo):
        self.K = np.array(msg.k).reshape(3, 3)

    def execute_callback(self, goal_handle: ServerGoalHandle):
        """同步执行回调 - 直接返回缓存结果"""
        print("=" * 60)
        print("🔴🔴🔴 EXECUTE_CALLBACK 被调用!🔴🔴🔴")
        print("=" * 60)
        self.get_logger().info("🔴🔴🔴 EXECUTE_CALLBACK 被调用!")

        goal = goal_handle.request
        res = VisionDetection.Result()

        try:
            # ✅ 直接使用缓存的最新结果
            points_3d = self.latest_points_3d
            colors = self.latest_colors
            pose_matrix = self.latest_pose_matrix

            self.get_logger().info(f"📊 缓存状态: points={points_3d is not None}, pose={pose_matrix is not None}")

            if points_3d is None or len(points_3d) == 0:
                self.get_logger().error("❌ 缓存中没有有效数据")
                res.success = False
                res.status_code = 1
                res.error_message = "没有可用的检测结果"
                goal_handle.abort()  # ✅ 不传参数
                return res

            self.get_logger().info(f"✅ 使用缓存结果,点数: {len(points_3d)}")

            # 设置结果
            res.success = True
            res.status_code = 0
            res.error_message = ""
            res.execution_time = time.time() - self.last_run_ts

            # 设置header
            res.header = Header()
            res.header.stamp = self.get_clock().now().to_msg()
            res.header.frame_id = goal.expected_frame_id if goal.expected_frame_id else "camera_link"

            # 处理6D位姿
            if goal.need_6d_pose:
                res.target_6d_pose = PoseStamped()
                res.target_6d_pose.header = Header()
                res.target_6d_pose.header.stamp = self.get_clock().now().to_msg()
                res.target_6d_pose.header.frame_id = goal.expected_frame_id if goal.expected_frame_id else "camera_link"

                if pose_matrix is not None and pose_matrix.shape == (4, 4):
                    self.get_logger().info(f"✅ 使用FoundationPose位姿")

                    res.target_6d_pose.pose.position = Point()
                    res.target_6d_pose.pose.position.x = float(pose_matrix[0, 3])
                    res.target_6d_pose.pose.position.y = float(pose_matrix[1, 3])
                    res.target_6d_pose.pose.position.z = float(pose_matrix[2, 3])

                    rot_matrix = pose_matrix[:3, :3]
                    rotation = Rotation.from_matrix(rot_matrix)
                    quat = rotation.as_quat()

                    res.target_6d_pose.pose.orientation = Quaternion()
                    res.target_6d_pose.pose.orientation.x = quat[0]
                    res.target_6d_pose.pose.orientation.y = quat[1]
                    res.target_6d_pose.pose.orientation.z = quat[2]
                    res.target_6d_pose.pose.orientation.w = quat[3]

                    res.pose_confidence = 0.95
                    self.get_logger().info(
                        f"  位置: ({pose_matrix[0, 3]:.3f}, {pose_matrix[1, 3]:.3f}, {pose_matrix[2, 3]:.3f})"
                    )
                else:
                    centroid = np.mean(points_3d, axis=0)
                    res.target_6d_pose.pose.position = Point()
                    res.target_6d_pose.pose.position.x = float(centroid[0])
                    res.target_6d_pose.pose.position.y = float(centroid[1])
                    res.target_6d_pose.pose.position.z = float(centroid[2])
                    res.target_6d_pose.pose.orientation = Quaternion()
                    res.target_6d_pose.pose.orientation.w = 1.0
                    res.pose_confidence = 0.5

            # 创建点云
            if goal.need_env_point_cloud and points_3d is not None:
                cloud_frame = goal.expected_frame_id if goal.expected_frame_id else "camera_link"
                res.env_point_cloud = create_point_cloud(
                    points_3d,
                    colors,
                    cloud_frame,
                    self.get_clock()
                )
                res.point_cloud_frame_id = cloud_frame
                self.get_logger().info(f"☁️ 点云已创建,点数: {len(points_3d)}")

            self.get_logger().info("✅ 检测完成,返回结果")
            goal_handle.succeed()  # ✅ 不传参数
            return res

        except Exception as e:
            self.get_logger().error(f"❌ execute_callback 异常: {e}")
            import traceback
            traceback.print_exc()
            res.success = False
            res.status_code = 3
            res.error_message = str(e)
            goal_handle.abort()  # ✅ 不传参数
            return res


def main(args=None):
    rclpy.init(args=args)
    node = DetectionPoseCloudServer()
    executor = MultiThreadedExecutor(num_threads=4)
    executor.add_node(node)

    try:
        executor.spin()
    except KeyboardInterrupt:
        node.get_logger().info("接收到退出信号")
    finally:
        # 关闭所有OpenCV窗口
        cv2.destroyAllWindows()
        node.action_server.destroy()
        node.destroy_node()


if __name__ == '__main__':
    main()

2. foundationpose环境搭建

一、克隆官方仓库(用国内镜像克隆)

git clone https://gitcode.com/gh_mirrors/fo/FoundationPose.git
mv FoundationPose ~/ros2_ws/

二、安装FoundationPose 的依赖

cd ~/ros2_ws/FoundationPose
pip install -r requirements.txt

三、安装PyTorch3D

cd ~/ros2_ws
git clone https://github.com/facebookresearch/pytorch3d.git
cd pytorch3d

# 环境变量(强制CUDA11.8 + 4线程编译)
export CUDA_HOME=/usr/local/cuda-11.8
export FORCE_CUDA=1
export MAX_JOBS=4

# 编译安装
pip install . --no-build-isolation

# 验证
python -c "import pytorch3d; print('✅ pytorch3d 安装成功')"

四、安装nvdiffrast

打开这个链接,浏览器会自动开始下载 nvdiffrast 压缩包: 👉 https://github.com/NVlabs/nvdiffrast/archive/refs/heads/master.zip

回到终端

# 1. 把下载好的 zip 包挪到当前目录(你下载的文件一般在 ~/下载/)
mv ~/下载/nvdiffrast-master.zip .

# 2. 解压
unzip nvdiffrast-master.zip

# 3. 把文件夹重命名成项目需要的名字 nvdiffrast
mv nvdiffrast-master nvdiffrast

切到 nvdiffrast-main 文件夹里

cd ~/ros2_ws/FoundationPose/nvdiffrast-main
ls

修改setup.py文件成下面

# Copyright (c) 2020, NVIDIA CORPORATION.  All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto.  Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.

import setuptools
import os
import torch
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
# Print an error message if there's no PyTorch installed.
#try:
    #from torch.utils.cpp_extension import BuildExtension, CUDAExtension
#except ImportError:
    # This happens if the user runs 'pip install' with default build isolation
    # OR if they simply don't have torch installed at all.
    #print("\n\n" + "*" * 70)
    #print("ERROR! Cannot compile nvdiffrast CUDA extension. Please ensure that:\n")
    #print("1. You have PyTorch installed")
    #print("2. You run 'pip install' with --no-build-isolation flag")
    #print("*" * 70 + "\n\n")
    #exit(1)

setuptools.setup(
    ext_modules=[
        CUDAExtension(
            "_nvdiffrast_c",
            sources=[
                "csrc/common/antialias.cu",
                "csrc/common/common.cpp",
                "csrc/common/cudaraster/impl/Buffer.cpp",
                "csrc/common/cudaraster/impl/CudaRaster.cpp",
                "csrc/common/cudaraster/impl/RasterImpl.cpp",
                "csrc/common/cudaraster/impl/RasterImpl_kernel.cu",
                "csrc/common/interpolate.cu",
                "csrc/common/rasterize.cu",
                "csrc/common/texture.cpp",
                "csrc/common/texture_kernel.cu",
                "csrc/torch/torch_antialias.cpp",
                "csrc/torch/torch_bindings.cpp",
                "csrc/torch/torch_interpolate.cpp",
                "csrc/torch/torch_rasterize.cpp",
                "csrc/torch/torch_texture.cpp",
            ],
            extra_compile_args={
                "cxx": ["-DNVDR_TORCH"]
                # Disable warnings in torch headers.
                + (["/wd4067", "/wd4624", "/wd4996"] if os.name == "nt" else []),
                "nvcc": ["-DNVDR_TORCH", "-lineinfo"],
            },
        )
    ],
    cmdclass={"build_ext": BuildExtension},
)

降级setuptools

conda activate ros2_tensorrt

# 降级到 81.x(最后一个带 pkg_resources 的版本)
pip install "setuptools<82" --force-reinstall

验证 pkg_resources
python -c "import pkg_resources; print('✅ pkg_resources OK')"

回去编译nvdiffrast

cd ~/ros2_ws/FoundationPose/nvdiffrast-main
rm -rf build dist *.egg-info

python setup.py build_ext --inplace
pip install . --no-build-isolation

#最后测试
python -c "import nvdiffrast; print('✅ nvdiffrast 安装成功')"

最后在 ~/ros2_ws/FoundationPose 下验证

python -c "from estimater import FoundationPose; print('✅ FoundationPose 导入成功!可以跑6D位姿了!')"

会看到

✅ FoundationPose 导入成功!可以跑6D位姿了!

五、运行foundation的代码fp_receiver_debug.py

#!/usr/bin/env python3
import socket
import json
import base64
import numpy as np
import cv2
import trimesh
import torch
import sys
import time
import os

sys.path.insert(0, "/home/wyq/ros2_ws/FoundationPose")

try:
    from estimater import FoundationPose

    print("✅ FoundationPose 导入成功")
except ImportError as e:
    print(f"❌ FoundationPose 导入失败: {e}")
    sys.exit(1)

# -----------------------------------------------------------------------------
# ✅ FoundationPose 统一可修改参数
# -----------------------------------------------------------------------------
CONFIG = {
    "host": "0.0.0.0",
    "port": 12345,
    "mesh_path": "/home/wyq/ros2_ws/phones.obj",
    "device": "cuda",
    "input_width": 640,
    "input_height": 400,
    "zfar": 2.0,
    "min_mask_area": 500,
    "iteration_register": 30,
    "scale_mesh": 1.0,
    "debug": 1,
}

# -----------------------------------------------------------------------------
# 加载3D模型
# -----------------------------------------------------------------------------
print("📦 加载3D模型...")
if not os.path.exists(CONFIG["mesh_path"]):
    print(f"❌ 模型文件不存在: {CONFIG['mesh_path']}")
    sys.exit(1)

mesh = trimesh.load(CONFIG["mesh_path"])
mesh.vertices /= CONFIG["scale_mesh"]
model_pts = mesh.vertices.astype(np.float32)
model_normals = mesh.vertex_normals.astype(np.float32)
print(f"✅ 模型加载成功,顶点数: {len(model_pts)}")
print(f"[调试] 模型范围: x [{model_pts[:, 0].min():.3f}, {model_pts[:, 0].max():.3f}]")
print(f"[调试] 模型范围: y [{model_pts[:, 1].min():.3f}, {model_pts[:, 1].max():.3f}]")
print(f"[调试] 模型范围: z [{model_pts[:, 2].min():.3f}, {model_pts[:, 2].max():.3f}]")

# 初始化FoundationPose
print("🔧 初始化FoundationPose...")
try:
    fp = FoundationPose(
        model_pts=model_pts,
        model_normals=model_normals,
        mesh=mesh,
        symmetry_tfs=None,
        debug=CONFIG["debug"]
    )
    print("✅ FoundationPose 初始化成功")
except Exception as e:
    print(f"❌ FoundationPose 初始化失败: {e}")
    sys.exit(1)

device = torch.device(CONFIG["device"])
torch.set_float32_matmul_precision('medium')
print(f"✅ 使用设备: {device}")


# -----------------------------------------------------------------------------
# Base64解码工具
# -----------------------------------------------------------------------------
def from_b64(s, shape, dtype=np.uint8):
    try:
        decoded = base64.b64decode(s)
        expected_size = 1
        for dim in shape:
            expected_size *= dim
        if dtype == np.float32:
            expected_size *= 4
        elif dtype == np.uint16:
            expected_size *= 2

        if len(decoded) != expected_size:
            print(f"[警告] 数据大小不匹配: {len(decoded)} != {expected_size}")
            if dtype == np.float32 and len(decoded) == expected_size // 4:
                arr = np.frombuffer(decoded, dtype=np.uint8).reshape(shape).astype(np.float32)
                return arr
            elif dtype == np.float32 and len(decoded) == expected_size // 2:
                arr = np.frombuffer(decoded, dtype=np.uint16).reshape(shape).astype(np.float32)
                return arr

        return np.frombuffer(decoded, dtype=dtype).reshape(shape).copy()
    except Exception as e:
        print(f"[错误] from_b64失败: {e}")
        raise


# -----------------------------------------------------------------------------
# 检查点云有效性
# -----------------------------------------------------------------------------
def check_point_cloud(rgb, depth, mask, K):
    """检查从depth和mask提取的点云是否有效"""
    ys, xs = np.where(mask > 0)
    if len(ys) == 0:
        print("[错误] mask为空,没有有效像素")
        return False

    # ✅ 判断depth单位
    depth_max = depth.max()
    if depth_max > 10:
        print(f"[调试] depth最大值 {depth_max:.0f},识别为毫米单位")
        depth_min_valid = 100  # 0.1米 = 100毫米
        depth_max_valid = 2000  # 2.0米 = 2000毫米
        is_mm = True
    else:
        print(f"[调试] depth最大值 {depth_max:.2f},识别为米单位")
        depth_min_valid = 0.1
        depth_max_valid = 2.0
        is_mm = False

    # 采样检查
    sample_indices = np.random.choice(len(ys), min(200, len(ys)), replace=False)
    valid_points = 0

    for idx in sample_indices:
        v, u = ys[idx], xs[idx]
        z = depth[v, u]
        if z > depth_min_valid and z < depth_max_valid:
            valid_points += 1

    print(f"[调试] 采样点数: {len(sample_indices)}, 有效深度点: {valid_points}")

    if valid_points < 50:
        print("[错误] 有效深度点太少,请检查depth数据")
        return False

    # 计算点云范围
    pts_3d = []
    for idx in sample_indices[:50]:
        v, u = ys[idx], xs[idx]
        z = depth[v, u]
        if z > depth_min_valid and z < depth_max_valid:
            # ✅ 转换为米
            if is_mm:
                z_m = z / 1000.0
            else:
                z_m = z
            x = (u - K[0, 2]) * z_m / K[0, 0]
            y = (v - K[1, 2]) * z_m / K[1, 1]
            pts_3d.append([x, y, z_m])

    if pts_3d:
        pts_3d = np.array(pts_3d)
        print(f"[调试] 点云范围: x [{pts_3d[:, 0].min():.3f}, {pts_3d[:, 0].max():.3f}]")
        print(f"[调试] 点云范围: y [{pts_3d[:, 1].min():.3f}, {pts_3d[:, 1].max():.3f}]")
        print(f"[调试] 点云范围: z [{pts_3d[:, 2].min():.3f}, {pts_3d[:, 2].max():.3f}]")
        return True
    else:
        print("[错误] 无法提取有效3D点")
        return False


# -----------------------------------------------------------------------------
# 主TCP接收+推理循环
# -----------------------------------------------------------------------------
def main():
    device = torch.device(CONFIG["device"])
    target_w = CONFIG["input_width"]
    target_h = CONFIG["input_height"]

    print(f"\n🚀 FoundationPose 6D位姿 TCP 服务器")
    print(f"   监听地址: {CONFIG['host']}:{CONFIG['port']}")
    print(f"   输入尺寸: {target_w}x{target_h}")
    print(f"   模式: 🔄 每次都是 Register")
    print(f"   迭代次数: {CONFIG['iteration_register']}")
    print("   等待连接...\n")

    while True:
        sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)

        try:
            sock.bind((CONFIG["host"], CONFIG["port"]))
            sock.listen(1)
            print("✅ 服务器已启动,等待客户端连接...")

            conn, addr = sock.accept()
            print(f"✅ 客户端连接: {addr}")
            conn.settimeout(10.0)

            buf = b''
            frame_count = 0

            while True:
                try:
                    data = conn.recv(262144)
                    if not data:
                        print("⚠️ 客户端断开连接")
                        break

                    buf += data
                    while b'\n' in buf:
                        line, buf = buf.split(b'\n', 1)
                        start_time = time.time()

                        try:
                            d = json.loads(line)

                            rgb_shape = d['rgb_shape']
                            depth_shape = d['depth_shape']
                            mask_shape = d['mask_shape']

                            rgb = from_b64(d['rgb_b64'], rgb_shape)

                            try:
                                depth = from_b64(d['depth_b64'], depth_shape, np.float32)
                            except ValueError as e:
                                print(f"[警告] depth解码失败: {e}")
                                try:
                                    depth = from_b64(d['depth_b64'], depth_shape, np.uint16).astype(np.float32)
                                except:
                                    depth = from_b64(d['depth_b64'], depth_shape, np.uint8).astype(np.float32)

                            # ✅ 检查depth单位并转换为米
                            if depth.max() > 10:
                                print(f"[调试] depth为毫米单位 (max={depth.max():.0f}mm),转换为米")
                                depth = depth / 1000.0
                                print(f"[调试] 转换后 depth range: {depth.min():.3f} - {depth.max():.3f}m")

                            mask = from_b64(d['mask_b64'], mask_shape)
                            K = np.array(d['K'])

                            print(f"[调试] rgb shape: {rgb.shape}, dtype: {rgb.dtype}")
                            print(f"[调试] depth shape: {depth.shape}, dtype: {depth.dtype}")
                            print(f"[调试] mask sum: {np.sum(mask)}, unique: {np.unique(mask)}")
                            print(f"[调试] depth range: {depth.min():.3f} - {depth.max():.3f}m")

                            # 检查并调整尺寸
                            if rgb.shape[0] != target_h or rgb.shape[1] != target_w:
                                print(f"⚠️ 尺寸不匹配,重新resize")
                                rgb = cv2.resize(rgb, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
                                depth = cv2.resize(depth, (target_w, target_h), interpolation=cv2.INTER_NEAREST)
                                mask = cv2.resize(mask, (target_w, target_h), interpolation=cv2.INTER_NEAREST)

                            mask_sum = np.sum(mask)
                            if mask_sum < CONFIG["min_mask_area"]:
                                print(f"[警告] mask面积太小: {mask_sum} < {CONFIG['min_mask_area']}")
                                torch.cuda.empty_cache()
                                continue

                            # ============================================================
                            # ✅ 检查点云有效性
                            # ============================================================
                            if not check_point_cloud(rgb, depth, mask, K):
                                print("[警告] 点云检查失败,跳过此帧")
                                torch.cuda.empty_cache()
                                continue

                            frame_count += 1
                            pose_data = None

                            # ============================================================
                            # ✅ 执行 Register
                            # ============================================================
                            try:
                                print(f"\n⚡ [帧 {frame_count}] Register (6D位姿注册)...")
                                print(f"[调试] 迭代次数: {CONFIG['iteration_register']}")

                                # 保存调试图像
                                cv2.imwrite(f"/tmp/debug_rgb_{frame_count}.jpg", rgb)
                                cv2.imwrite(f"/tmp/debug_depth_{frame_count}.jpg",
                                            (depth / depth.max() * 255).astype(np.uint8))
                                cv2.imwrite(f"/tmp/debug_mask_{frame_count}.jpg", mask)
                                print(f"[调试] 已保存图像到 /tmp/debug_rgb_{frame_count}.jpg")

                                # ✅ 调用 register
                                pose_1d = fp.register(
                                    K=K,
                                    rgb=rgb,
                                    depth=depth,
                                    ob_mask=mask,
                                    iteration=CONFIG["iteration_register"]
                                )

                                pose_mat = torch.tensor(pose_1d, device=device).float()

                                # ✅ 检查是否为有效位姿
                                pose_np = pose_mat.cpu().numpy()
                                if np.allclose(pose_np, np.eye(4)):
                                    print("[警告] ⚠️ 位姿为单位矩阵,计算失败!")
                                    pose_data = None
                                else:
                                    print("✅ Register OK")
                                    print(f"   6D位姿矩阵:\n{pose_np}")
                                    pose_data = pose_np.tolist()

                            except Exception as e:
                                print(f"❌ 6D位姿计算失败: {e}")
                                import traceback
                                traceback.print_exc()
                                pose_data = None

                            # ============================================================
                            # ✅ 回传6D位姿数据
                            # ============================================================
                            if pose_data is not None:
                                try:
                                    resp = json.dumps({
                                        "pose": pose_data,
                                        "ts": time.time(),
                                        "frame": frame_count
                                    }) + "\n"
                                    conn.sendall(resp.encode("utf-8"))
                                    elapsed = time.time() - start_time
                                    print(f"   ✅ 6D位姿已回传,总耗时: {elapsed:.3f}s")
                                except Exception as e:
                                    print(f"⚠️ 回传位姿失败: {e}")
                                    break

                            torch.cuda.empty_cache()

                        except json.JSONDecodeError as e:
                            print(f"JSON解析错误: {e}")
                            continue
                        except Exception as e:
                            print(f"帧处理异常: {e}")
                            import traceback
                            traceback.print_exc()
                            continue

                except socket.timeout:
                    print("⏰ 接收超时,继续等待...")
                    continue
                except Exception as e:
                    print(f"连接异常: {e}")
                    break

            conn.close()
            sock.close()
            print("连接已关闭,等待新连接...\n")

        except KeyboardInterrupt:
            print("\n\n👋 服务器已停止")
            sys.exit(0)
        except Exception as e:
            print(f"服务器异常: {e}")
            sock.close()
            time.sleep(1)


if __name__ == '__main__':
    print("\n" + "=" * 60)
    print("FoundationPose 6D位姿 TCP 服务器 (每次Register)")
    print("=" * 60)
    try:
        main()
    except KeyboardInterrupt:
        print("\n\n👋 服务器已停止")
        sys.exit(0)

运行以上代码就可以接收分割代码的运行的数据,并调用foundationpose的功能,运行前注意修改参数,终端运行代码如下:

CUDA_LAUNCH_BLOCKING=1 TORCH_USE_CUDA_DSA=1 python fp_receiver_debug.py

六、若运行onnx模型,需要修改foundationpose源码

        由于我这里用的foundationpose模型是下载的onnx模型(pth模型无法下载),对源代码做了一些修改。把以下的py文件直接用这里的代码替换即可:

1.estimater.py

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto.  Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.


from Utils import *
from datareader import *
import itertools
from learning.training.predict_score import *
from learning.training.predict_pose_refine import *
import yaml


class FoundationPose:
  def __init__(self, model_pts, model_normals, symmetry_tfs=None, mesh=None, scorer:ScorePredictor=None, refiner:PoseRefinePredictor=None, glctx=None, debug=0, debug_dir='/home/wyq/ros2_ws/debug/novel_pose_debug/'):
    self.gt_pose = None
    self.ignore_normal_flip = True
    self.debug = debug
    self.debug_dir = debug_dir
    os.makedirs(debug_dir, exist_ok=True)

    self.reset_object(model_pts, model_normals, symmetry_tfs=symmetry_tfs, mesh=mesh)
    self.make_rotation_grid(min_n_views=40, inplane_step=60)

    self.glctx = glctx

    if scorer is not None:
      self.scorer = scorer
    else:
      self.scorer = ScorePredictor()

    if refiner is not None:
      self.refiner = refiner
    else:
      self.refiner = PoseRefinePredictor()

    self.pose_last = None   # Used for tracking; per the centered mesh


  def reset_object(self, model_pts, model_normals, symmetry_tfs=None, mesh=None):
    max_xyz = mesh.vertices.max(axis=0)
    min_xyz = mesh.vertices.min(axis=0)
    self.model_center = (min_xyz+max_xyz)/2
    if mesh is not None:
      self.mesh_ori = mesh.copy()
      mesh = mesh.copy()
      mesh.vertices = mesh.vertices - self.model_center.reshape(1,3)

    model_pts = mesh.vertices
    self.diameter = compute_mesh_diameter(model_pts=mesh.vertices, n_sample=10000)
    self.vox_size = max(self.diameter/20.0, 0.003)
    logging.info(f'self.diameter:{self.diameter}, vox_size:{self.vox_size}')
    self.dist_bin = self.vox_size/2
    self.angle_bin = 20  # Deg
    pcd = toOpen3dCloud(model_pts, normals=model_normals)
    pcd = pcd.voxel_down_sample(self.vox_size)
    self.max_xyz = np.asarray(pcd.points).max(axis=0)
    self.min_xyz = np.asarray(pcd.points).min(axis=0)
    self.pts = torch.tensor(np.asarray(pcd.points), dtype=torch.float32, device='cuda')
    self.normals = F.normalize(torch.tensor(np.asarray(pcd.normals), dtype=torch.float32, device='cuda'), dim=-1)
    logging.info(f'self.pts:{self.pts.shape}')
    self.mesh_path = None
    self.mesh = mesh
    if self.mesh is not None:
      self.mesh_path = f'/tmp/{uuid.uuid4()}.obj'
      self.mesh.export(self.mesh_path)
    self.mesh_tensors = make_mesh_tensors(self.mesh)

    if symmetry_tfs is None:
      self.symmetry_tfs = torch.eye(4).float().cuda()[None]
    else:
      self.symmetry_tfs = torch.as_tensor(symmetry_tfs, device='cuda', dtype=torch.float)

    logging.info("reset done")



  def get_tf_to_centered_mesh(self):
    tf_to_center = torch.eye(4, dtype=torch.float, device='cuda')
    tf_to_center[:3,3] = -torch.as_tensor(self.model_center, device='cuda', dtype=torch.float)
    return tf_to_center


  def to_device(self, s='cuda:0'):
    for k in self.__dict__:
      self.__dict__[k] = self.__dict__[k]
      if torch.is_tensor(self.__dict__[k]) or isinstance(self.__dict__[k], nn.Module):
        logging.info(f"Moving {k} to device {s}")
        self.__dict__[k] = self.__dict__[k].to(s)
    for k in self.mesh_tensors:
      logging.info(f"Moving {k} to device {s}")
      self.mesh_tensors[k] = self.mesh_tensors[k].to(s)
    if self.refiner is not None:
      self.refiner.model.to(s)
    if self.scorer is not None:
      self.scorer.model.to(s)
    if self.glctx is not None:
      self.glctx = dr.RasterizeCudaContext(s)



  def make_rotation_grid(self, min_n_views=40, inplane_step=60):
    cam_in_obs = sample_views_icosphere(n_views=min_n_views)
    logging.info(f'cam_in_obs:{cam_in_obs.shape}')
    rot_grid = []
    for i in range(len(cam_in_obs)):
      for inplane_rot in np.deg2rad(np.arange(0, 360, inplane_step)):#(0, 360, 360)
        cam_in_ob = cam_in_obs[i]
        R_inplane = euler_matrix(0,0,inplane_rot)
        cam_in_ob = cam_in_ob@R_inplane
        ob_in_cam = np.linalg.inv(cam_in_ob)
        rot_grid.append(ob_in_cam)

    rot_grid = np.asarray(rot_grid)
    logging.info(f"rot_grid:{rot_grid.shape}")
    rot_grid = mycpp.cluster_poses(8, 99999, rot_grid, self.symmetry_tfs.data.cpu().numpy())#30, 99999
    rot_grid = np.asarray(rot_grid)
    # 🔥🔥🔥 暴力强行只保留前 8 个姿态,必解决显存!
    rot_grid = rot_grid[:8]
    logging.info(f"after cluster, rot_grid:{rot_grid.shape}")
    self.rot_grid = torch.as_tensor(rot_grid, device='cuda', dtype=torch.float)
    logging.info(f"self.rot_grid: {self.rot_grid.shape}")


  def generate_random_pose_hypo(self, K, rgb, depth, mask, scene_pts=None):
    '''
    @scene_pts: torch tensor (N,3)
    '''
    ob_in_cams = self.rot_grid.clone()
    center = self.guess_translation(depth=depth, mask=mask, K=K)
    ob_in_cams[:,:3,3] = torch.tensor(center, device='cuda', dtype=torch.float).reshape(1,3)
    return ob_in_cams


  def guess_translation(self, depth, mask, K):
    vs,us = np.where(mask>0)
    if len(us)==0:
      logging.info(f'mask is all zero')
      return np.zeros((3))
    uc = (us.min()+us.max())/2.0
    vc = (vs.min()+vs.max())/2.0
    valid = mask.astype(bool) & (depth>=0.001)
    if not valid.any():
      logging.info(f"valid is empty")
      return np.zeros((3))

    zc = np.median(depth[valid])
    center = (np.linalg.inv(K)@np.asarray([uc,vc,1]).reshape(3,1))*zc

    if self.debug>=2:
      pcd = toOpen3dCloud(center.reshape(1,3))
      o3d.io.write_point_cloud(f'{self.debug_dir}/init_center.ply', pcd)

    return center.reshape(3)


  def register(self, K, rgb, depth, ob_mask, ob_id=None, glctx=None, iteration=5):
    '''Copmute pose from given pts to self.pcd
    @pts: (N,3) np array, downsampled scene points
    '''
    set_seed(0)
    logging.info('Welcome')

    if self.glctx is None:
      if glctx is None:
        self.glctx = dr.RasterizeCudaContext()
        # self.glctx = dr.RasterizeGLContext()
      else:
        self.glctx = glctx

    depth = erode_depth(depth, radius=2, device='cuda')
    depth = bilateral_filter_depth(depth, radius=2, device='cuda')

    if self.debug>=2:
      xyz_map = depth2xyzmap(depth, K)
      valid = xyz_map[...,2]>=0.001
      pcd = toOpen3dCloud(xyz_map[valid], rgb[valid])
      o3d.io.write_point_cloud(f'{self.debug_dir}/scene_raw.ply',pcd)
      cv2.imwrite(f'{self.debug_dir}/ob_mask.png', (ob_mask*255.0).clip(0,255))

    normal_map = None
    valid = (depth>=0.001) & (ob_mask>0)
    if valid.sum()<4:
      logging.info(f'valid too small, return')
      pose = np.eye(4)
      pose[:3,3] = self.guess_translation(depth=depth, mask=ob_mask, K=K)
      return pose

    if self.debug>=2:
      imageio.imwrite(f'{self.debug_dir}/color.png', rgb)
      cv2.imwrite(f'{self.debug_dir}/depth.png', (depth*1000).astype(np.uint16))
      valid = xyz_map[...,2]>=0.001
      pcd = toOpen3dCloud(xyz_map[valid], rgb[valid])
      o3d.io.write_point_cloud(f'{self.debug_dir}/scene_complete.ply',pcd)

    self.H, self.W = depth.shape[:2]
    self.K = K
    self.ob_id = ob_id
    self.ob_mask = ob_mask

    poses = self.generate_random_pose_hypo(K=K, rgb=rgb, depth=depth, mask=ob_mask, scene_pts=None)
    poses = poses.data.cpu().numpy()
    logging.info(f'poses:{poses.shape}')
    center = self.guess_translation(depth=depth, mask=ob_mask, K=K)

    poses = torch.as_tensor(poses, device='cuda', dtype=torch.float)
    poses[:,:3,3] = torch.as_tensor(center.reshape(1,3), device='cuda')

    add_errs = self.compute_add_err_to_gt_pose(poses)
    logging.info(f"after viewpoint, add_errs min:{add_errs.min()}")

    xyz_map = depth2xyzmap(depth, K)
    poses, vis = self.refiner.predict(mesh=self.mesh, mesh_tensors=self.mesh_tensors, rgb=rgb, depth=depth, K=K, ob_in_cams=poses.data.cpu().numpy(), normal_map=normal_map, xyz_map=xyz_map, glctx=self.glctx, mesh_diameter=self.diameter, iteration=iteration, get_vis=self.debug>=2)
    if vis is not None:
      imageio.imwrite(f'{self.debug_dir}/vis_refiner.png', vis)

    scores, vis = self.scorer.predict(mesh=self.mesh, rgb=rgb, depth=depth, K=K, ob_in_cams=poses.data.cpu().numpy(), normal_map=normal_map, mesh_tensors=self.mesh_tensors, glctx=self.glctx, mesh_diameter=self.diameter, get_vis=self.debug>=2)
    if vis is not None:
      imageio.imwrite(f'{self.debug_dir}/vis_score.png', vis)

    add_errs = self.compute_add_err_to_gt_pose(poses)
    logging.info(f"final, add_errs min:{add_errs.min()}")

    ids = torch.as_tensor(scores).argsort(descending=True)
    logging.info(f'sort ids:{ids}')
    scores = scores[ids]
    poses = poses[ids]

    logging.info(f'sorted scores:{scores}')

    best_pose = poses[0]@self.get_tf_to_centered_mesh()
    self.pose_last = poses[0]
    self.best_id = ids[0]

    self.poses = poses
    self.scores = scores

    return best_pose.data.cpu().numpy()


  def compute_add_err_to_gt_pose(self, poses):
    '''
    @poses: wrt. the centered mesh
    '''
    return -torch.ones(len(poses), device='cuda', dtype=torch.float)


  def track_one(self, rgb, depth, K, iteration, extra={}):
    if self.pose_last is None:
      logging.info("Please init pose by register first")
      raise RuntimeError
    logging.info("Welcome")

    depth = torch.as_tensor(depth, device='cuda', dtype=torch.float)
    depth = erode_depth(depth, radius=2, device='cuda')
    depth = bilateral_filter_depth(depth, radius=2, device='cuda')
    logging.info("depth processing done")

    xyz_map = depth2xyzmap_batch(depth[None], torch.as_tensor(K, dtype=torch.float, device='cuda')[None], zfar=np.inf)[0]

    pose, vis = self.refiner.predict(mesh=self.mesh, mesh_tensors=self.mesh_tensors, rgb=rgb, depth=depth, K=K, ob_in_cams=self.pose_last.reshape(1,4,4).data.cpu().numpy(), normal_map=None, xyz_map=xyz_map, mesh_diameter=self.diameter, glctx=self.glctx, iteration=iteration, get_vis=self.debug>=2)
    logging.info("pose done")
    if self.debug>=2:
      extra['vis'] = vis
    self.pose_last = pose
    return (pose@self.get_tf_to_centered_mesh()).data.cpu().numpy().reshape(4,4)

2.predict_pose_refine.py

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto.  Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.


import functools
import os,sys,kornia
import time
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../../')
import numpy as np
import torch
from omegaconf import OmegaConf
from learning.datasets.h5_dataset import *
from Utils import *
from datareader import *
import onnxruntime as ort


@torch.inference_mode()
def make_crop_data_batch(render_size, ob_in_cams, mesh, rgb, depth, K, crop_ratio, xyz_map, normal_map=None, mesh_diameter=None, cfg=None, glctx=None, mesh_tensors=None, dataset:PoseRefinePairH5Dataset=None):
  logging.info("Welcome make_crop_data_batch")
  H,W = depth.shape[:2]
  args = []
  method = 'box_3d'
  tf_to_crops = compute_crop_window_tf_batch(pts=mesh.vertices, H=H, W=W, poses=ob_in_cams, K=K, crop_ratio=crop_ratio, out_size=(render_size[1], render_size[0]), method=method, mesh_diameter=mesh_diameter)

  logging.info("make tf_to_crops done")

  B = len(ob_in_cams)
  poseA = torch.as_tensor(ob_in_cams, dtype=torch.float, device='cuda')

  bs = 4
  rgb_rs = []
  depth_rs = []
  normal_rs = []
  xyz_map_rs = []

  bbox2d_crop = torch.as_tensor(np.array([0, 0, cfg['input_resize'][0]-1, cfg['input_resize'][1]-1]).reshape(2,2), device='cuda', dtype=torch.float)
  bbox2d_ori = transform_pts(bbox2d_crop, tf_to_crops.inverse()).reshape(-1,4)

  for b in range(0,len(poseA),bs):
    extra = {}
    rgb_r, depth_r, normal_r = nvdiffrast_render(K=K, H=H, W=W, ob_in_cams=poseA[b:b+bs], context='cuda', get_normal=cfg['use_normal'], glctx=glctx, mesh_tensors=mesh_tensors, output_size=cfg['input_resize'], bbox2d=bbox2d_ori[b:b+bs], use_light=True, extra=extra)
    rgb_rs.append(rgb_r)
    depth_rs.append(depth_r[...,None])
    normal_rs.append(normal_r)
    xyz_map_rs.append(extra['xyz_map'])
  rgb_rs = torch.cat(rgb_rs, dim=0).permute(0,3,1,2) * 255
  depth_rs = torch.cat(depth_rs, dim=0).permute(0,3,1,2)  #(B,1,H,W)
  xyz_map_rs = torch.cat(xyz_map_rs, dim=0).permute(0,3,1,2)  #(B,3,H,W)
  Ks = torch.as_tensor(K, device='cuda', dtype=torch.float).reshape(1,3,3)
  if cfg['use_normal']:
    normal_rs = torch.cat(normal_rs, dim=0).permute(0,3,1,2)  #(B,3,H,W)

  logging.info("render done")

  rgbBs = kornia.geometry.transform.warp_perspective(torch.as_tensor(rgb, dtype=torch.float, device='cuda').permute(2,0,1)[None].expand(B,-1,-1,-1), tf_to_crops, dsize=render_size, mode='bilinear', align_corners=False)
  if rgb_rs.shape[-2:]!=cfg['input_resize']:
    rgbAs = kornia.geometry.transform.warp_perspective(rgb_rs, tf_to_crops, dsize=render_size, mode='bilinear', align_corners=False)
  else:
    rgbAs = rgb_rs
  if xyz_map_rs.shape[-2:]!=cfg['input_resize']:
    xyz_mapAs = kornia.geometry.transform.warp_perspective(xyz_map_rs, tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)
  else:
    xyz_mapAs = xyz_map_rs
  xyz_mapBs = kornia.geometry.transform.warp_perspective(torch.as_tensor(xyz_map, device='cuda', dtype=torch.float).permute(2,0,1)[None].expand(B,-1,-1,-1), tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)  #(B,3,H,W)

  if cfg['use_normal']:
    normalAs = kornia.geometry.transform.warp_perspective(normal_rs, tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)
    normalBs = kornia.geometry.transform.warp_perspective(torch.as_tensor(normal_map, dtype=torch.float, device='cuda').permute(2,0,1)[None].expand(B,-1,-1,-1), tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)
  else:
    normalAs = None
    normalBs = None

  logging.info("warp done")

  mesh_diameters = torch.ones((len(rgbAs)), dtype=torch.float, device='cuda')*mesh_diameter
  pose_data = BatchPoseData(rgbAs=rgbAs, rgbBs=rgbBs, depthAs=None, depthBs=None, normalAs=normalAs, normalBs=normalBs, poseA=poseA, poseB=None, xyz_mapAs=xyz_mapAs, xyz_mapBs=xyz_mapBs, tf_to_crops=tf_to_crops, Ks=Ks, mesh_diameters=mesh_diameters)
  pose_data = dataset.transform_batch(batch=pose_data, H_ori=H, W_ori=W, bound=1)

  logging.info("pose batch data done")

  return pose_data


class PoseRefinePredictor:
  def __init__(self,):
    logging.info("welcome")
    self.amp = True
    self.run_name = "2023-10-28-18-33-37"
    code_dir = os.path.dirname(os.path.realpath(__file__))

    self.cfg = OmegaConf.load(f'{code_dir}/../../weights/{self.run_name}/config.yml')
    self.cfg['enable_amp'] = True

    ########## Defaults, to be backward compatible
    if 'use_normal' not in self.cfg:
      self.cfg['use_normal'] = False
    if 'use_mask' not in self.cfg:
      self.cfg['use_mask'] = False
    if 'use_BN' not in self.cfg:
      self.cfg['use_BN'] = False
    if 'c_in' not in self.cfg:
      self.cfg['c_in'] = 4
    if 'crop_ratio' not in self.cfg or self.cfg['crop_ratio'] is None:
      self.cfg['crop_ratio'] = 1.2
    if 'n_view' not in self.cfg:
      self.cfg['n_view'] = 1
    if 'trans_rep' not in self.cfg:
      self.cfg['trans_rep'] = 'tracknet'
    if 'rot_rep' not in self.cfg:
      self.cfg['rot_rep'] = 'axis_angle'
    if 'zfar' not in self.cfg:
      self.cfg['zfar'] = 3
    if 'normalize_xyz' not in self.cfg:
      self.cfg['normalize_xyz'] = False
    if isinstance(self.cfg['zfar'], str) and 'inf' in self.cfg['zfar'].lower():
      self.cfg['zfar'] = np.inf
    if 'normal_uint8' not in self.cfg:
      self.cfg['normal_uint8'] = False
    if 'trans_normalizer' not in self.cfg:
      self.cfg['trans_normalizer'] = 0.1
    if 'rot_normalizer' not in self.cfg:
      self.cfg['rot_normalizer'] = 0.1
    if 'input_resize' not in self.cfg:
      self.cfg['input_resize'] = (160, 160)
    
    logging.info(f"self.cfg: \n {OmegaConf.to_yaml(self.cfg)}")

    self.dataset = PoseRefinePairH5Dataset(cfg=self.cfg, h5_file='', mode='test')

    # ====================== 🔥 ONNX 模型加载 ======================
    self.model = ort.InferenceSession(
        '/home/wyq/ros2_ws/FoundationPose/weights/2023-10-28-18-33-37/refine_model.onnx',
        providers=['CUDAExecutionProvider']
    )
    # =============================================================

    logging.info("✅ PoseRefine ONNX 模型加载成功!")
    self.last_trans_update = None
    self.last_rot_update = None


  @torch.inference_mode()
  def predict(self, rgb, depth, K, ob_in_cams, xyz_map, normal_map=None, get_vis=False, mesh=None, mesh_tensors=None, glctx=None, mesh_diameter=None, iteration=5):
    '''
    @rgb: np array (H,W,3)
    @ob_in_cams: np array (N,4,4)
    '''
    torch.set_default_tensor_type('torch.cuda.FloatTensor')
    logging.info(f'ob_in_cams:{ob_in_cams.shape}')
    tf_to_center = np.eye(4)
    ob_centered_in_cams = ob_in_cams
    mesh_centered = mesh

    logging.info(f'self.cfg.use_normal:{self.cfg.use_normal}')
    if not self.cfg.use_normal:
      normal_map = None

    crop_ratio = self.cfg['crop_ratio']
    logging.info(f"trans_normalizer:{self.cfg['trans_normalizer']}, rot_normalizer:{self.cfg['rot_normalizer']}")
    bs = 1024

    B_in_cams = torch.as_tensor(ob_centered_in_cams, device='cuda', dtype=torch.float)

    if mesh_tensors is None:
      mesh_tensors = make_mesh_tensors(mesh_centered)

    rgb_tensor = torch.as_tensor(rgb, device='cuda', dtype=torch.float)
    depth_tensor = torch.as_tensor(depth, device='cuda', dtype=torch.float)
    xyz_map_tensor = torch.as_tensor(xyz_map, device='cuda', dtype=torch.float)
    trans_normalizer = self.cfg['trans_normalizer']
    if not isinstance(trans_normalizer, float):
      trans_normalizer = torch.as_tensor(list(trans_normalizer), device='cuda', dtype=torch.float).reshape(1,3)

    for _ in range(iteration):
      logging.info("making cropped data")
      pose_data = make_crop_data_batch(self.cfg.input_resize, B_in_cams, mesh_centered, rgb_tensor, depth_tensor, K, crop_ratio=crop_ratio, normal_map=normal_map, xyz_map=xyz_map_tensor, cfg=self.cfg, glctx=glctx, mesh_tensors=mesh_tensors, dataset=self.dataset, mesh_diameter=mesh_diameter)
      B_in_cams = []
      for b in range(0, pose_data.rgbAs.shape[0], bs):
        A = torch.cat([pose_data.rgbAs[b:b+bs].cuda(), pose_data.xyz_mapAs[b:b+bs].cuda()], dim=1).float()
        B = torch.cat([pose_data.rgbBs[b:b+bs].cuda(), pose_data.xyz_mapBs[b:b+bs].cuda()], dim=1).float()
        
        A = torch.nn.functional.interpolate(A, size=(160, 160), mode='bilinear', align_corners=False)
        B = torch.nn.functional.interpolate(B, size=(160, 160), mode='bilinear', align_corners=False)
        # 形状从 [B,6,160,160] → [B,160,160,6]
        A = A.permute(0, 2, 3, 1).contiguous()
        B = B.permute(0, 2, 3, 1).contiguous()

        
        logging.info("forward start")

        # ====================== 🔥 ONNX 推理 ======================
        A_np = A.cpu().numpy().astype(np.float32)
        B_np = B.cpu().numpy().astype(np.float32)

        outs = self.model.run(None, {
            "input1": A_np,
            "input2": B_np
        })

        # 输出转回 torch
        output = {
            "trans": torch.from_numpy(outs[0]).cuda(),
            "rot": torch.from_numpy(outs[1]).cuda()
        }
        # ===========================================================

        for k in output:
          output[k] = output[k].float()
        logging.info("forward done")

        if self.cfg['trans_rep']=='tracknet':
          if not self.cfg['normalize_xyz']:
            trans_delta = torch.tanh(output["trans"])*trans_normalizer
          else:
            trans_delta = output["trans"]

        elif self.cfg['trans_rep']=='deepim':
          def project_and_transform_to_crop(centers):
            uvs = (pose_data.Ks[b:b+bs]@centers.reshape(-1,3,1)).reshape(-1,3)
            uvs = uvs/uvs[:,2:3]
            uvs = (pose_data.tf_to_crops[b:b+bs]@uvs.reshape(-1,3,1)).reshape(-1,3)
            return uvs[:,:2]

          rot_delta = output["rot"]
          z_pred = output['trans'][:,2]*pose_data.poseA[b:b+bs][...,2,3]
          uvA_crop = project_and_transform_to_crop(pose_data.poseA[b:b+bs][...,:3,3])
          uv_pred_crop = uvA_crop + output['trans'][:,:2]*self.cfg['input_resize'][0]
          uv_pred = transform_pts(uv_pred_crop, pose_data.tf_to_crops[b:b+bs].inverse().cuda())
          center_pred = torch.cat([uv_pred, torch.ones((len(rot_delta),1), dtype=torch.float, device='cuda')], dim=-1)
          center_pred = (pose_data.Ks[b:b+bs].inverse().cuda()@center_pred.reshape(len(rot_delta),3,1)).reshape(len(rot_delta),3) * z_pred.reshape(len(rot_delta),1)
          trans_delta = center_pred-pose_data.poseA[b:b+bs][...,:3,3]

        else:
          trans_delta = output["trans"]

        if self.cfg['rot_rep']=='axis_angle':
          rot_mat_delta = torch.tanh(output["rot"])*self.cfg['rot_normalizer']
          rot_mat_delta = so3_exp_map(rot_mat_delta).permute(0,2,1)
        elif self.cfg['rot_rep']=='6d':
          rot_mat_delta = rotation_6d_to_matrix(output['rot']).permute(0,2,1)
        else:
          raise RuntimeError

        if self.cfg['normalize_xyz']:
          trans_delta *= (mesh_diameter/2)

        B_in_cam = egocentric_delta_pose_to_pose(pose_data.poseA[b:b+bs], trans_delta=trans_delta, rot_mat_delta=rot_mat_delta)
        B_in_cams.append(B_in_cam)

      B_in_cams = torch.cat(B_in_cams, dim=0).reshape(len(ob_in_cams),4,4)

    B_in_cams_out = B_in_cams@torch.tensor(tf_to_center[None], device='cuda', dtype=torch.float)
    torch.cuda.empty_cache()
    self.last_trans_update = trans_delta
    self.last_rot_update = rot_mat_delta

    if get_vis:
      logging.info("get_vis...")
      canvas = []
      padding = 2
      pose_data = make_crop_data_batch(self.cfg.input_resize, torch.as_tensor(ob_centered_in_cams), mesh_centered, rgb, depth, K, crop_ratio=crop_ratio, normal_map=normal_map, xyz_map=xyz_map_tensor, cfg=self.cfg, glctx=glctx, mesh_tensors=mesh_tensors, dataset=self.dataset, mesh_diameter=mesh_diameter)
      for id in range(0, len(B_in_cams)):
        rgbA_vis = (pose_data.rgbAs[id]*255).permute(1,2,0).data.cpu().numpy()
        rgbB_vis = (pose_data.rgbBs[id]*255).permute(1,2,0).data.cpu().numpy()
        row = [rgbA_vis, rgbB_vis]
        H,W = rgbA_vis.shape[:2]
        if pose_data.depthAs is not None:
          depthA = pose_data.depthAs[id].data.cpu().numpy().reshape(H,W)
          depthB = pose_data.depthBs[id].data.cpu().numpy().reshape(H,W)
        elif pose_data.xyz_mapAs is not None:
          depthA = pose_data.xyz_mapAs[id][2].data.cpu().numpy().reshape(H,W)
          depthB = pose_data.xyz_mapBs[id][2].data.cpu().numpy().reshape(H,W)
        zmin = min(depthA.min(), depthB.min())
        zmax = max(depthA.max(), depthB.max())
        depthA_vis = depth_to_vis(depthA, zmin=zmin, zmax=zmax, inverse=False)
        depthB_vis = depth_to_vis(depthB, zmin=zmin, zmax=zmax, inverse=False)
        row += [depthA_vis, depthB_vis]
        if pose_data.normalAs is not None:
          pass
        row = make_grid_image(row, nrow=len(row), padding=padding, pad_value=255)
        row = cv_draw_text(row, text=f'id:{id}', uv_top_left=(10,10), color=(0,255,0), fontScale=0.5)
        canvas.append(row)
      canvas = make_grid_image(canvas, nrow=1, padding=padding, pad_value=255)

      pose_data = make_crop_data_batch(self.cfg.input_resize, B_in_cams, mesh_centered, rgb, depth, K, crop_ratio=crop_ratio, normal_map=normal_map, xyz_map=xyz_map_tensor, cfg=self.cfg, glctx=glctx, mesh_tensors=mesh_tensors, dataset=self.dataset, mesh_diameter=mesh_diameter)
      canvas_refined = []
      for id in range(0, len(B_in_cams)):
        rgbA_vis = (pose_data.rgbAs[id]*255).permute(1,2,0).data.cpu().numpy()
        rgbB_vis = (pose_data.rgbBs[id]*255).permute(1,2,0).data.cpu().numpy()
        row = [rgbA_vis, rgbB_vis]
        H,W = rgbA_vis.shape[:2]
        if pose_data.depthAs is not None:
          depthA = pose_data.depthAs[id].data.cpu().numpy().reshape(H,W)
          depthB = pose_data.depthBs[id].data.cpu().numpy().reshape(H,W)
        elif pose_data.xyz_mapAs is not None:
          depthA = pose_data.xyz_mapAs[id][2].data.cpu().numpy().reshape(H,W)
          depthB = pose_data.xyz_mapBs[id][2].data.cpu().numpy().reshape(H,W)
        zmin = min(depthA.min(), depthB.min())
        zmax = max(depthA.max(), depthB.max())
        depthA_vis = depth_to_vis(depthA, zmin=zmin, zmax=zmax, inverse=False)
        depthB_vis = depth_to_vis(depthB, zmin=zmin, zmax=zmax, inverse=False)
        row += [depthA_vis, depthB_vis]
        row = make_grid_image(row, nrow=len(row), padding=padding, pad_value=255)
        canvas_refined.append(row)

      canvas_refined = make_grid_image(canvas_refined, nrow=1, padding=padding, pad_value=255)
      canvas = make_grid_image([canvas, canvas_refined], nrow=2, padding=padding, pad_value=255)
      torch.cuda.empty_cache()
      return B_in_cams_out, canvas

    return B_in_cams_out, None

3.predict_score.py

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto.  Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.


import functools
import os,sys,kornia
import time
import numpy as np
import torch
import torch.distributed as dist
from omegaconf import OmegaConf
from tqdm import tqdm
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../../../')
from learning.datasets.h5_dataset import *
from learning.datasets.pose_dataset import *
from Utils import *
from datareader import *
import onnxruntime as ort


def vis_batch_data_scores(pose_data, ids, scores, pad_margin=5):
  assert len(scores)==len(ids)
  canvas = []
  for id in ids:
    rgbA_vis = (pose_data.rgbAs[id]*255).permute(1,2,0).data.cpu().numpy()
    rgbB_vis = (pose_data.rgbBs[id]*255).permute(1,2,0).data.cpu().numpy()
    H,W = rgbA_vis.shape[:2]
    zmin = pose_data.depthAs[id].data.cpu().numpy().reshape(H,W).min()
    zmax = pose_data.depthAs[id].data.cpu().numpy().reshape(H,W).max()
    depthA_vis = depth_to_vis(pose_data.depthAs[id].data.cpu().numpy().reshape(H,W), zmin=zmin, zmax=zmax, inverse=False)
    depthB_vis = depth_to_vis(pose_data.depthBs[id].data.cpu().numpy().reshape(H,W), zmin=zmin, zmax=zmax, inverse=False)
    if pose_data.normalAs is not None:
      pass
    pad = np.ones((rgbA_vis.shape[0],pad_margin,3))*255
    if pose_data.normalAs is not None:
      pass
    else:
      row = np.concatenate([rgbA_vis, pad, depthA_vis, pad, rgbB_vis, pad, depthB_vis], axis=1)
    s = 100/row.shape[0]
    row = cv2.resize(row, fx=s, fy=s, dsize=None)
    row = cv_draw_text(row, text=f'id:{id}, score:{scores[id]:.3f}', uv_top_left=(10,10), color=(0,255,0), fontScale=0.5)
    canvas.append(row)
    pad = np.ones((pad_margin, row.shape[1], 3))*255
    canvas.append(pad)
  canvas = np.concatenate(canvas, axis=0).astype(np.uint8)
  return canvas



@torch.no_grad()
def make_crop_data_batch(render_size, ob_in_cams, mesh, rgb, depth, K, crop_ratio, normal_map=None, mesh_diameter=None, glctx=None, mesh_tensors=None, dataset:TripletH5Dataset=None, cfg=None):
  logging.info("Welcome make_crop_data_batch")
  H,W = depth.shape[:2]

  args = []
  method = 'box_3d'
  tf_to_crops = compute_crop_window_tf_batch(pts=mesh.vertices, H=H, W=W, poses=ob_in_cams, K=K, crop_ratio=crop_ratio, out_size=(render_size[1], render_size[0]), method=method, mesh_diameter=mesh_diameter)
  logging.info("make tf_to_crops done")

  B = len(ob_in_cams)
  poseAs = torch.as_tensor(ob_in_cams, dtype=torch.float, device='cuda')

  bs = 4
  rgb_rs = []
  depth_rs = []
  xyz_map_rs = []

  bbox2d_crop = torch.as_tensor(np.array([0, 0, cfg['input_resize'][0]-1, cfg['input_resize'][1]-1]).reshape(2,2), device='cuda', dtype=torch.float)
  bbox2d_ori = transform_pts(bbox2d_crop, tf_to_crops.inverse()[:,None]).reshape(-1,4)

  for b in range(0,len(ob_in_cams),bs):
    extra = {}
    rgb_r, depth_r, normal_r = nvdiffrast_render(K=K, H=H, W=W, ob_in_cams=poseAs[b:b+bs], context='cuda', get_normal=cfg['use_normal'], glctx=glctx, mesh_tensors=mesh_tensors, output_size=cfg['input_resize'], bbox2d=bbox2d_ori[b:b+bs], use_light=True, extra=extra)
    rgb_rs.append(rgb_r)
    depth_rs.append(depth_r[...,None])
    xyz_map_rs.append(extra['xyz_map'])

  rgb_rs = torch.cat(rgb_rs, dim=0).permute(0,3,1,2) * 255
  depth_rs = torch.cat(depth_rs, dim=0).permute(0,3,1,2)
  xyz_map_rs = torch.cat(xyz_map_rs, dim=0).permute(0,3,1,2)  #(B,3,H,W)
  logging.info("render done")

  rgbBs = kornia.geometry.transform.warp_perspective(torch.as_tensor(rgb, dtype=torch.float, device='cuda').permute(2,0,1)[None].expand(B,-1,-1,-1), tf_to_crops, dsize=render_size, mode='bilinear', align_corners=False)
  depthBs = kornia.geometry.transform.warp_perspective(torch.as_tensor(depth, dtype=torch.float, device='cuda')[None,None].expand(B,-1,-1,-1), tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)
  if rgb_rs.shape[-2:]!=cfg['input_resize']:
    rgbAs = kornia.geometry.transform.warp_perspective(rgb_rs, tf_to_crops, dsize=render_size, mode='bilinear', align_corners=False)
    depthAs = kornia.geometry.transform.warp_perspective(depth_rs, tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)
  else:
    rgbAs = rgb_rs
    depthAs = depth_rs

  if xyz_map_rs.shape[-2:]!=cfg['input_resize']:
    xyz_mapAs = kornia.geometry.transform.warp_perspective(xyz_map_rs, tf_to_crops, dsize=render_size, mode='nearest', align_corners=False)
  else:
    xyz_mapAs = xyz_map_rs

  normalAs = None
  normalBs = None

  Ks = torch.as_tensor(K, dtype=torch.float).reshape(1,3,3).expand(len(rgbAs),3,3)
  mesh_diameters = torch.ones((len(rgbAs)), dtype=torch.float, device='cuda')*mesh_diameter

  pose_data = BatchPoseData(rgbAs=rgbAs, rgbBs=rgbBs, depthAs=depthAs, depthBs=depthBs, normalAs=normalAs, normalBs=normalBs, poseA=poseAs, xyz_mapAs=xyz_mapAs, tf_to_crops=tf_to_crops, Ks=Ks, mesh_diameters=mesh_diameters)
  pose_data = dataset.transform_batch(pose_data, H_ori=H, W_ori=W, bound=1)

  logging.info("pose batch data done")

  return pose_data


class ScorePredictor:
  def __init__(self, amp=True):
    self.amp = amp
    self.run_name = "2024-01-11-20-02-45"
    
    code_dir = os.path.dirname(os.path.realpath(__file__))
    self.cfg = OmegaConf.load(f'{code_dir}/../../weights/{self.run_name}/config.yml')
    
    self.cfg['enable_amp'] = True
    ########## Defaults, to be backward compatible
    if 'use_normal' not in self.cfg:
      self.cfg['use_normal'] = False
    if 'use_BN' not in self.cfg:
      self.cfg['use_BN'] = False
    if 'zfar' not in self.cfg:
      self.cfg['zfar'] = np.inf
    if 'c_in' not in self.cfg:
      self.cfg['c_in'] = 4
    if 'normalize_xyz' not in self.cfg:
      self.cfg['normalize_xyz'] = False
    if 'crop_ratio' not in self.cfg or self.cfg['crop_ratio'] is None:
      self.cfg['crop_ratio'] = 1.2
    if 'input_resize' not in self.cfg:
      self.cfg['input_resize'] = (160, 160)

    logging.info(f"self.cfg: \n {OmegaConf.to_yaml(self.cfg)}")

    self.dataset = ScoreMultiPairH5Dataset(cfg=self.cfg, mode='test', h5_file=None, max_num_key=1)

    # ====================== 🔥 纯 ONNX 加载 ======================
    # 删掉 PyTorch 模型
    # self.model = ScoreNetMultiPair(cfg=self.cfg, c_in=self.cfg['c_in']).cuda()

    # 直接加载 ONNX
    self.model = ort.InferenceSession(
        '/home/wyq/ros2_ws/FoundationPose/weights/2024-01-11-20-02-45/score_model.onnx',
        providers=["CUDAExecutionProvider"]
    )
    self.input_names = [i.name for i in self.model.get_inputs()]
    self.output_names = [o.name for o in self.model.get_outputs()]

    logging.info("✅ Score ONNX model loaded (CUDA)")


  @torch.inference_mode()
  def predict(self, rgb, depth, K, ob_in_cams, normal_map=None, get_vis=False, mesh=None, mesh_tensors=None, glctx=None, mesh_diameter=None):
    '''
    @rgb: np array (H,W,3)
    '''
    logging.info(f"ob_in_cams:{ob_in_cams.shape}")
    ob_in_cams = torch.as_tensor(ob_in_cams, dtype=torch.float, device='cuda')

    logging.info(f'self.cfg.use_normal:{self.cfg.use_normal}')
    if not self.cfg.use_normal:
      normal_map = None

    logging.info("making cropped data")

    if mesh_tensors is None:
      mesh_tensors = make_mesh_tensors(mesh)

    rgb = torch.as_tensor(rgb, device='cuda', dtype=torch.float)
    depth = torch.as_tensor(depth, device='cuda', dtype=torch.float)

    pose_data = make_crop_data_batch(self.cfg.input_resize, ob_in_cams, mesh, rgb, depth, K, crop_ratio=self.cfg['crop_ratio'], glctx=glctx, mesh_tensors=mesh_tensors, dataset=self.dataset, cfg=self.cfg, mesh_diameter=mesh_diameter)

    def find_best_among_pairs(pose_data:BatchPoseData):
      logging.info(f'pose_data.rgbAs.shape[0]: {pose_data.rgbAs.shape[0]}')
      ids = []
      scores = []
      bs = pose_data.rgbAs.shape[0]
      for b in range(0, pose_data.rgbAs.shape[0], bs):
        A = torch.cat([pose_data.rgbAs[b:b+bs].cuda(), pose_data.xyz_mapAs[b:b+bs].cuda()], dim=1).float()
        B = torch.cat([pose_data.rgbBs[b:b+bs].cuda(), pose_data.xyz_mapBs[b:b+bs].cuda()], dim=1).float()
        
        if pose_data.normalAs is not None:
          A = torch.cat([A, pose_data.normalAs.cuda().float()], dim=1)
          B = torch.cat([B, pose_data.normalBs.cuda().float()], dim=1)
        
        A = torch.nn.functional.interpolate(A, size=(160, 160), mode='bilinear', align_corners=False)
        B = torch.nn.functional.interpolate(B, size=(160, 160), mode='bilinear', align_corners=False)
        
        # 形状从 [B,6,160,160] → [B,160,160,6]
        A = A.permute(0, 2, 3, 1).contiguous()
        B = B.permute(0, 2, 3, 1).contiguous()
        
        # ONNX Runtime 推理
        A_np = A.cpu().numpy().astype(np.float32)
        B_np = B.cpu().numpy().astype(np.float32)
        
        outputs = self.model.run(
            None, 
            {
                "input1": A_np,
                "input2": B_np
            }
        )
        
        scores_cur = torch.from_numpy(outputs[0]).cuda().float().reshape(-1)

        ids.append(scores_cur.argmax()+b)
        scores.append(scores_cur)
      ids = torch.stack(ids, dim=0).reshape(-1)
      scores = torch.cat(scores, dim=0).reshape(-1)
      return ids, scores

    pose_data_iter = pose_data
    global_ids = torch.arange(len(ob_in_cams), device='cuda', dtype=torch.long)
    scores_global = torch.zeros((len(ob_in_cams)), dtype=torch.float, device='cuda')

    while 1:
      ids, scores = find_best_among_pairs(pose_data_iter)
      if len(ids)==1:
        scores_global[global_ids] = scores + 100
        break
      global_ids = global_ids[ids]
      pose_data_iter = pose_data.select_by_indices(global_ids)

    scores = scores_global

    logging.info(f'forward done')
    torch.cuda.empty_cache()

    if get_vis:
      logging.info("get_vis...")
      canvas = []
      ids = scores.argsort(descending=True)
      canvas = vis_batch_data_scores(pose_data, ids=ids, scores=scores)
      return scores, canvas

    return scores, None

4.training_config.py

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto.  Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.


import os,sys
from dataclasses import dataclass, field
from typing import List, Optional, Tuple,Union
import numpy as np
import omegaconf
import torch


@dataclass
class TrainingConfig(omegaconf.dictconfig.DictConfig):
    input_resize: tuple = (160, 160)
    normalize_xyz:Optional[bool] = True
    use_mask:Optional[bool] = False
    crop_ratio:Optional[float] = None
    split_objects_across_gpus: bool = True
    max_num_key: Optional[int] = None
    use_normal:bool = False
    n_view:int = 1
    zfar:float = np.inf
    c_in:int = 6
    train_num_pair:Optional[int] = None
    make_pair_online:Optional[bool] = False
    render_backend:Optional[str] = 'nvdiffrast'

    # Run management
    run_id: Optional[str] = None
    exp_name:Optional[str] = None
    resume_run_id: Optional[str] = None
    save_dir: Optional[str] = None
    batch_size: int = 64
    epoch_size: int = 115200
    val_size: int = 1280
    n_epochs: int = 25
    save_epoch_interval: int = 100
    n_dataloader_workers: int = 20
    n_rendering_workers: int = 1
    gradient_max_norm:float = np.inf
    max_step_per_epoch: Optional[int] = 25000

    # Network
    use_BN:bool = True
    loss_type:Optional[str] = 'pairwise_valid'

    # Optimizer
    optimizer: str = "adam"
    weight_decay: float = 0.0
    clip_grad_norm: float = np.inf
    lr: float = 0.0001
    warmup_step: int = -1   # -1 means disable
    n_epochs_warmup: int = 1

    # Visualization
    vis_interval: Optional[int] = 1000

    debug: Optional[bool] = None



@dataclass
class TrainRefinerConfig:
    # Datasets
    input_resize: tuple = (160, 160)  #(W,H)
    crop_ratio:Optional[float] = None
    max_num_key: Optional[int] = None
    use_normal:bool = False
    use_mask:Optional[bool] = False
    normal_uint8:bool = False
    normalize_xyz:Optional[bool] = True
    trans_normalizer:Optional[list] = None
    rot_normalizer:Optional[float] = None
    c_in:int = 6
    n_view:int = 1
    zfar:float = np.inf
    trans_rep:str = 'tracknet'  # tracknet/deepim
    rot_rep:Optional[str] = 'axis_angle'  # 6d/axis_angle
    save_dir: Optional[str] = None

    # Run management
    run_id: Optional[str] = None
    exp_name:Optional[str] = None
    batch_size: int = 64
    use_BN:bool = True
    optimizer: str = "adam"
    weight_decay: float = 0.0
    clip_grad_norm: float = np.inf
    lr: float = 0.0001
    warmup_step: int = -1
    loss_type:str = 'l2'   # l1/l2/add

    vis_interval: Optional[int] = 1000
    debug: Optional[bool] = None

5.代码运行

在foundationpose环境下运行

python3 fp_receiver_debug.py

在seg_pose环境下运行

source install/setup.bash
python3 src/vision_py_node/vision_py_node/ros2_sender.py

在第3个终端运行测试

ros2 action send_goal /vision/detection_pose_cloud vision_detection_action/action/VisionDetection "{need_6d_pose: true, need_env_point_cloud: false, target_obj_name: 'object', expected_frame_id: 'camera_link', priority: 1}"

3.bundlesdf环境搭建

一、下载源码

资源类型地址
论文 PDFhttps://arxiv.org/abs/2303.14158
项目主页https://bundlesdf.github.io/
GitCode 镜像https://gitcode.com/gh_mirrors/bu/BundleSDF

二、环境搭建

系统:ubuntu22.04

1.Docker环境搭建

进入docker文件夹加

cd ~/public/BundleSDF-master/docker

在该文件内下载pybind11.zip、pytorch3d.zip、opencv-4.11.0.tar.gz、opencv_contrib-4.11.0.tar.gz、yaml-cpp.zip,各自的版本在Docker里面有注释,下载后修改Docker如下:

FROM nvcr.io/nvidia/deepstream:7.1-triton-multiarch

ARG CMAKE_VERSION_MAJOR=3
ARG CMAKE_VERSION_MINOR=25
ARG CMAKE_VERSION_PATCH=3

ARG EIGEN_VERSION_MAJOR=3
ARG EIGEN_VERSION_MINOR=4
ARG EIGEN_VERSION_PATCH=0

ARG OPENCV_VERSION_MAJOR=4
ARG OPENCV_VERSION_MINOR=11
ARG OPENCV_VERSION_PATCH=0

ARG PCL_VERSION_MAJOR=1
ARG PCL_VERSION_MINOR=10
ARG PCL_VERSION_PATCH=0

ARG PYBIND11_VERSION_MAJOR=2
ARG PYBIND11_VERSION_MINOR=13
ARG PYBIND11_VERSION_PATCH=0

ARG YAML_CPP_VERSION_MAJOR=0
ARG YAML_CPP_VERSION_MINOR=8
ARG YAML_CPP_VERSION_PATCH=0

ENV TZ=US/Pacific
#RUN ln -snf /usr/share/zoneinfo/$TZ /etc/localtime && echo $TZ > /etc/timezone
ENV TZ=US/Pacific
RUN apt-get update && apt-get install -y --no-install-recommends tzdata \
    && ln -snf /usr/share/zoneinfo/$TZ /etc/localtime \
    && echo $TZ > /etc/timezone

# Install dependencies
RUN apt-get update --fix-missing && apt-get install -y --no-install-recommends \
    # Python and development tools
    python3-pip \
    python3-dev \
    # Build tools and compilers
    build-essential \
    cmake \
    cmake-curses-gui \
    checkinstall \
    g++ \
    gcc \
    gfortran \
    # Version control and utilities
    git \
    vim \
    tmux \
    wget \
    curl \
    bzip2 \
    ca-certificates \
    gnupg \
    software-properties-common \
    # Graphics and GUI libraries
    libglib2.0-0 \
    libsm6 \
    libxext6 \
    libxrender-dev \
    libgtk2.0-dev \
    qtbase5-dev \
    # Math and science libraries
    libblas-dev \
    liblapack-dev \
    libatlas-base-dev \
    # Security and networking
    libssl-dev \
    libzmq3-dev \
    # Boost libraries
    libboost-filesystem-dev \
    libboost-date-time-dev \
    libboost-iostreams-dev \
    libboost-system-dev \
    libboost-program-options-dev \
    libboost-all-dev \
    # Point cloud and 3D processing
    libflann-dev \
    # Image and video processing
    libjpeg8-dev \
    libtiff5-dev \
    pkg-config \
    yasm \
    libavcodec-dev \
    libavformat-dev \
    libswscale-dev \
    libdc1394-dev \
    libxine2-dev \
    libv4l-dev \
    libtbb-dev \
    ffmpeg \
    # Audio/video codecs
    libfaac-dev \
    libmp3lame-dev \
    libtheora-dev \
    libvorbis-dev \
    libxvidcore-dev \
    libopencore-amrnb-dev \
    libopencore-amrwb-dev \
    x264 \
    v4l-utils \
    # Protocol buffers and logging
    libprotobuf-dev \
    protobuf-compiler \
    libgoogle-glog-dev \
    libgflags-dev \
    # Additional libraries
    libgphoto2-dev \
    libhdf5-dev \
    doxygen \
    proj-data \
    libproj-dev \
    libyaml-cpp-dev \
    libzmq3-dev \
    freeglut3-dev \
    && rm -rf /var/lib/apt/lists/*


# Install cmake
RUN cd / &&\
wget http://www.cmake.org/files/v${CMAKE_VERSION_MAJOR}.${CMAKE_VERSION_MINOR}/cmake-${CMAKE_VERSION_MAJOR}.${CMAKE_VERSION_MINOR}.${CMAKE_VERSION_PATCH}.tar.gz &&\
tar xf cmake-${CMAKE_VERSION_MAJOR}.${CMAKE_VERSION_MINOR}.${CMAKE_VERSION_PATCH}.tar.gz &&\
cd cmake-${CMAKE_VERSION_MAJOR}.${CMAKE_VERSION_MINOR}.${CMAKE_VERSION_PATCH} &&\
./configure &&\
make &&\
make install


# Install Eigen
RUN cd / && \
    wget https://gitlab.com/libeigen/eigen/-/archive/${EIGEN_VERSION_MAJOR}.${EIGEN_VERSION_MINOR}.${EIGEN_VERSION_PATCH}/eigen-${EIGEN_VERSION_MAJOR}.${EIGEN_VERSION_MINOR}.${EIGEN_VERSION_PATCH}.tar.gz && \
    tar xf eigen-${EIGEN_VERSION_MAJOR}.${EIGEN_VERSION_MINOR}.${EIGEN_VERSION_PATCH}.tar.gz && \
    cd eigen-${EIGEN_VERSION_MAJOR}.${EIGEN_VERSION_MINOR}.${EIGEN_VERSION_PATCH} && \
    mkdir build && \
    cd build && \
    cmake .. && \
    make install && \
    cd / && \
    rm -rf eigen-${EIGEN_VERSION_MAJOR}.${EIGEN_VERSION_MINOR}.${EIGEN_VERSION_PATCH}.tar.gz eigen-${EIGEN_VERSION_MAJOR}.${EIGEN_VERSION_MINOR}.${EIGEN_VERSION_PATCH}



SHELL ["/bin/bash", "--login", "-c"]
# 1. 系统依赖
#RUN apt-get update
RUN apt-key adv --keyserver hkp://keyserver.ubuntu.com:80 --recv-keys FB0B24895113F120 || true
RUN apt-get update --fix-missing
RUN apt-get remove -y python3-blinker || true
RUN apt-get install -y unzip
# 2. 安装 PyTorch + 核心库
RUN pip3 install --upgrade pip setuptools wheel -i https://pypi.tuna.tsinghua.edu.cn/simple
RUN pip3 install torch==2.5.0 torchvision==0.20.0 torchaudio --index-url https://download.pytorch.org/whl/cu124
RUN pip3 install --no-cache-dir --ignore-installed blinker kaolin==0.17.0 -f https://nvidia-kaolin.s3.us-east-2.amazonaws.com/torch-2.5.0_cu124.html -i https://pypi.tuna.tsinghua.edu.cn/simple

# 3. 安装 pytorch3d + 全部Python库
COPY pytorch3d.zip /tmp/pytorch3d.zip
RUN cd /tmp && unzip -q pytorch3d.zip && cd pytorch3d-main && MAX_JOBS=4 pip3 install --no-cache-dir --no-build-isolation . -i https://pypi.tuna.tsinghua.edu.cn/simple

RUN cd / && rm -rf /tmp/pytorch3d*
RUN pip3 install --force-reinstall blinker -i https://pypi.tuna.tsinghua.edu.cn/simple
RUN pip3 install trimesh wandb matplotlib imageio tqdm -i https://pypi.tuna.tsinghua.edu.cn/simple
RUN pip3 install open3d ruamel.yaml sacred kornia -i https://pypi.tuna.tsinghua.edu.cn/simple
RUN pip3 install pymongo pyrender jupyterlab ninja -i https://pypi.tuna.tsinghua.edu.cn/simple
RUN pip3 install "Cython>=0.29.37" yacs -i https://pypi.tuna.tsinghua.edu.cn/simple
RUN pip3 install scipy scikit-learn -i https://pypi.tuna.tsinghua.edu.cn/simple



# Install OpenCV
COPY opencv-4.11.0.tar.gz /
COPY opencv_contrib-4.11.0.tar.gz /


# 解压源码包
RUN cd / && \
    tar -xzf opencv-4.11.0.tar.gz && mv opencv-4.11.0 opencv && \
    tar -xzf opencv_contrib-4.11.0.tar.gz && mv opencv_contrib-4.11.0 opencv_contrib && \
    rm -rf *.tar.gz
# Build OpenCV
RUN cd /opencv && \
    mkdir build && \
    cd build && \
    cmake ..  -DCMAKE_BUILD_TYPE=Release \
        -DBUILD_CUDA_STUBS=OFF \
        -DBUILD_DOCS=OFF \
        -DWITH_MATLAB=OFF \
        -DCUDA_FAST_MATH=ON \
        -DMKL_WITH_OPENMP=ON \
        -DOPENCV_ENABLE_NONFREE=ON \
        -DWITH_OPENMP=ON \
        -DWITH_QT=ON \
        -DWITH_OPENEXR=ON \
        -DENABLE_PRECOMPILED_HEADERS=OFF \
        -DBUILD_opencv_cudacodec=OFF \
        -DINSTALL_PYTHON_EXAMPLES=OFF \
        -DWITH_TIFF=OFF \
        -DWITH_WEBP=OFF \
        -DWITH_FFMPEG=ON \
        -DOPENCV_EXTRA_MODULES_PATH=../../opencv_contrib/modules \
        -DCMAKE_CXX_FLAGS=-std=c++17 \
        -DENABLE_CXX11=OFF \
        -DBUILD_opencv_xfeatures2d=OFF \
        -DOPENCV_DNN_OPENCL=OFF \
        -DWITH_CUDA=ON \
        -DWITH_OPENCL=OFF \
        -DBUILD_opencv_wechat_qrcode=OFF \
        -DCMAKE_CXX_STANDARD=17 \
        -DCMAKE_CXX_STANDARD_REQUIRED=ON \
        -DOPENCV_CUDA_OPTIONS_opencv_test_cudev=-std=c++17 \
        -DCUDA_ARCH_BIN="8.6" \
        -DCMAKE_INSTALL_PREFIX=/usr/local \
        -DCMAKE_INSTALL_LIBDIR=lib \
        -DINSTALL_PKGCONFIG=ON \
        -DBUILD_opencv_python3=ON \
        -DPYTHON3_EXECUTABLE=$(which python3) \
        -DOPENCV_GENERATE_PKGCONFIG=ON \
        -DPKG_CONFIG_PATH=/usr/local/lib/pkgconfig \
        -DINSTALL_PYTHON_EXAMPLES=OFF \
        -DINSTALL_C_EXAMPLES=OFF && \
    make -j4 && \
    make install && \
    cd / && \
    rm -rf /opencv /opencv_contrib


# Install PCL
RUN apt-get update && apt-get install -y --no-install-recommends \
    libeigen3-dev \
    libflann-dev \
    libboost-all-dev \
    libusb-1.0-0-dev \
    libvtk7-dev \
    && rm -rf /var/lib/apt/lists/*

RUN apt-get update && apt-get install -y libpcl-dev

# Install Pybind11
COPY pybind11.zip /pybind11.zip
RUN cd / && \
    unzip -q pybind11.zip && \
    mv pybind11-* pybind11 && \
    mkdir -p /pybind11/build && \
    cd /pybind11/build && \
    cmake .. -DCMAKE_BUILD_TYPE=Release -DPYBIND11_INSTALL=ON -DPYBIND11_TEST=OFF && \
    make -j$(nproc) && \
    make install && \
    cd / && \
    rm -rf /pybind11

# Install YAML-CPP
COPY yaml-cpp.zip /yaml-cpp.zip
RUN cd / && \
    unzip -q yaml-cpp.zip && \
    mv yaml-cpp-* yaml-cpp && \
    mkdir -p /yaml-cpp/build && \
    cd /yaml-cpp/build && \
    cmake .. \
        -DCMAKE_POLICY_VERSION_MINIMUM=3.5 \
        -DBUILD_TESTING=OFF \
        -DCMAKE_BUILD_TYPE=Release \
        -DINSTALL_GTEST=OFF \
        -DYAML_CPP_BUILD_TESTS=OFF \
        -DYAML_BUILD_SHARED_LIBS=ON && \
    make -j$(nproc) && \
    make install && \
    cd / && \
    rm -rf /yaml-cpp



ENV CUDA_HOME=/usr/local/cuda
ENV LD_LIBRARY_PATH=/usr/local/cuda/lib64
ENV OPENCV_IO_ENABLE_OPENEXR=1
# Fix multiprocessing issues
ENV PYTHONUNBUFFERED=1
ENV OMP_NUM_THREADS=1


#RUN imageio_download_bin freeimage
# 🔥 彻底修复 freeimage 下载失败(100% 能用)
ENV IMAGEIO_NO_FREEIMAGE=1
ENV IMAGEIO_FREEIMAGE_LIB=/usr/lib/x86_64-linux-gnu/libfreeimage.so
RUN apt-get update && apt-get install -y --no-install-recommends libfreeimage-dev && rm -rf /var/lib/apt/lists/*

#### Kaolin will change numpy version
RUN pip3 install --break-system-packages --force-reinstall blinker
RUN pip3 install numpy==1.26.4 transformations einops scikit-image awscli-plugin-endpoint gputil xatlas pymeshlab rtree dearpygui pytinyrenderer PyQt5 cython-npm chardet openpyxl

RUN apt-get update --fix-missing
RUN apt install -y rsync lbzip2 pigz zip p7zip-full p7zip-rar

修改完后,编译Docker,命令含有ros2与docker共享网络的的参数(获得ros2话题的相机参数)

sudo docker run -it --rm \
  --gpus all \
  --net=host \
  --privileged \
  -e ROS_DOMAIN_ID=0 \
  -v /dev:/dev \
  -v /home/wyq/BundleSDF-master:/workspace \
  bundlesdf:latest \
  bash



#进入根目录
cd /workspace
ls

2, 下载loftr匹配用的权重文件及下载loftr源码

连接https://pan.baidu.com/s/1dwUDx6A9IRMBkCSowLIz5Q

提取马:jlcl

下载好LoFTR源码,权重文件放在BundleTrack/LoFTR/weights下面。

3.ros2获取相机参数

在bundlesdf的根目录下建立camera_datas文件夹,在里面写个ros2_save_to_docker.py文件,代码如下:

#!/usr/bin/env python3
import rclpy
from rclpy.node import Node
from sensor_msgs.msg import Image, CameraInfo
from cv_bridge import CvBridge
import cv2
import numpy as np
import os
import sys
import time
import matplotlib
import matplotlib.pyplot as plt

matplotlib.use('TkAgg')
plt.rcParams['font.sans-serif'] = ['DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False

os.environ["QT_LOGGING_RULES"] = "false"
os.environ["OPENCV_LOG_LEVEL"] = "FATAL"

SAM_WORK_DIR = "/home/wyq/BundleSDF-master/sam2-main"
os.chdir(SAM_WORK_DIR)

from sam2.build_sam import build_sam2
from sam2.sam2_image_predictor import SAM2ImagePredictor

cfg_path = "configs/sam2.1/sam2.1_hiera_b+.yaml"
ckpt_path = "weights/sam2.1_hiera_base_plus.pt"
sam2_model = build_sam2(cfg_path, ckpt_path, device="cuda")
predictor = SAM2ImagePredictor(sam2_model)
print("✅ SAM2 加载成功")

points = []
labels = []

def onclick(event):
    global points, labels
    if event.xdata is None or event.ydata is None:
        return
    x = int(event.xdata)
    y = int(event.ydata)
    if event.button == 1:
        print(f"🟢 目标点:({x}, {y})")
        points.append([x, y])
        labels.append(1)
    elif event.button == 3:
        print(f"🔴 背景点:({x}, {y})")
        points.append([x, y])
        labels.append(0)

def smooth_mask_with_gradient(img_bgr, clean_mask, grad_thresh=5, epsilon=0.7):
    gray = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)
    cnt_smooth, _ = cv2.findContours(clean_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
    smooth_mask = np.zeros_like(clean_mask)
    if not cnt_smooth:
        return clean_mask

    main_cnt = max(cnt_smooth, key=cv2.contourArea)
    h_img, w_img = gray.shape[:2]
    fix_cont = []
    for pt in main_cnt:
        x, y = int(pt[0][0]), int(pt[0][1])
        y_up = max(0, y - 1)
        y_down = min(h_img - 1, y + 1)
        x_left = max(0, x - 1)
        x_right = min(w_img - 1, x + 1)

        g_up = gray[y_up, x]
        g_down = gray[y_down, x]
        g_left = gray[y, x_left]
        g_right = gray[y, x_right]
        g_curr = gray[y, x]

        grad_up = abs(int(g_up) - int(g_curr))
        grad_down = abs(int(g_down) - int(g_curr))
        grad_left = abs(int(g_left) - int(g_curr))
        grad_right = abs(int(g_right) - int(g_curr))
        total_grad = grad_up + grad_down + grad_left + grad_right

        dirs = [[0,-1,clean_mask[y_up,x]], [0,1,clean_mask[y_down,x]],[-1,0,clean_mask[y,x_left]], [1,0,clean_mask[y,x_right]]]
        bg_dir = None
        for dx,dy,val in dirs:
            if val == 0:
                bg_dir = (dx, dy)
                break
        if bg_dir is not None:
            x += bg_dir[0]
            y += bg_dir[1]
        else:
            if total_grad < grad_thresh:
                max_grad = max(grad_up,grad_down,grad_left,grad_right)
                if max_grad == grad_up: y -=1
                elif max_grad == grad_down: y +=1
                elif max_grad == grad_left: x -=1
                else: x +=1
        fix_cont.append([[x,y]])

    fix_cont = np.array(fix_cont, dtype=np.int32)
    smooth_cont = cv2.approxPolyDP(fix_cont, epsilon, closed=True)
    cv2.drawContours(smooth_mask, [smooth_cont], -1,255,-1)
    smooth_mask = cv2.GaussianBlur(smooth_mask,(3,3),0.3)
    final_out = (smooth_mask>127).astype(np.uint8)*255
    return final_out

def do_segment(img_bgr, save_mask, save_overlay):
    global points, labels
    points.clear()
    labels.clear()
    if img_bgr is None:
        print("❌ 图像为空")
        return
    H, W = img_bgr.shape[:2]
    img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)
    img_rgb = cv2.medianBlur(img_rgb, 5)
    img_rgb = cv2.GaussianBlur(img_rgb, (3, 3), sigmaX=0.8)

    plt.figure(figsize=(10, 7))
    plt.imshow(img_rgb)
    plt.title("Left click: Foreground | Right click: Background")
    plt.axis("off")
    plt.gcf().canvas.mpl_connect('button_press_event', onclick)
    plt.show()

    if len(points) == 0 or len(points) != len(labels):
        print("❌ 点位数量不匹配,跳过分割")
        return

    pts_np = np.array(points, dtype=np.float32)
    lab_np = np.array(labels, dtype=np.int32)

    print("\n🔍 第一步:全图粗分割")
    predictor.set_image(img_rgb)
    masks1, scores1, _ = predictor.predict(point_coords=pts_np,point_labels=lab_np,multimask_output=True,return_logits=False)
    best1_idx = np.argmax(scores1)
    rough_mask = (masks1[best1_idx] * 255).astype(np.uint8)

    cnts, _ = cv2.findContours(rough_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
    if len(cnts) == 0:
        crop_rgb = img_rgb
        off_x, off_y = 0, 0
    else:
        max_cnt = max(cnts, key=cv2.contourArea)
        x, y, w, h = cv2.boundingRect(max_cnt)
        pad = 15
        x1 = max(0, x - pad)
        y1 = max(0, y - pad)
        x2 = min(W, x + w + pad)
        y2 = min(H, y + h + pad)
        off_x, off_y = x1, y1
        crop_rgb = img_rgb[y1:y2, x1:x2]
        pts_np[:, 0] -= off_x
        pts_np[:, 1] -= off_y

    print("🔍 第二步:裁剪区域精细分割")
    predictor.set_image(crop_rgb)
    masks2, scores2, _ = predictor.predict(point_coords=pts_np,point_labels=lab_np,multimask_output=True,return_logits=False)
    best2_idx = np.argmax(scores2)
    crop_mask = (masks2[best2_idx] * 255).astype(np.uint8)

    final_mask = np.zeros((H, W), dtype=np.uint8)
    final_mask[off_y:off_y+crop_mask.shape[0], off_x:off_x+crop_mask.shape[1]] = crop_mask

    kernel_close = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(3,3))
    kernel_erode = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(2,2))
    final_mask = cv2.morphologyEx(final_mask, cv2.MORPH_CLOSE, kernel_close, iterations=1)
    final_mask = cv2.morphologyEx(final_mask, cv2.MORPH_ERODE, kernel_erode, iterations=1)

    cnt_all,_ = cv2.findContours(final_mask,cv2.RETR_EXTERNAL,cv2.CHAIN_APPROX_SIMPLE)
    clean_mask = np.zeros_like(final_mask)
    if len(cnt_all)>0:
        maxc = max(cnt_all,key=cv2.contourArea)
        cv2.drawContours(clean_mask,[maxc],-1,255,-1)
    clean_mask = smooth_mask_with_gradient(img_bgr, clean_mask, grad_thresh=8, epsilon=1.0)

    overlay = img_bgr.copy()
    roi = overlay[clean_mask > 0]
    if roi.size > 0:
        color_fill = np.full_like(roi, (0, 0, 255), dtype=np.uint8)
        overlay[clean_mask > 0] = cv2.addWeighted(roi, 0.3, color_fill, 0.7, 0)

    cv2.imwrite(save_mask, clean_mask)
    cv2.imwrite(save_overlay, overlay)
    print(f"\n🎉 分割完成!mask:{save_mask} | overlay:{save_overlay}")

# 切回主目录
os.chdir("/home/wyq/BundleSDF-master")
SAVE_DIR = "./realtime_input"
os.makedirs(SAVE_DIR, exist_ok=True)
K_SAVE_PATH = os.path.join(SAVE_DIR, "K.txt")

# 时序保存配置
SAVE_INTERVAL = 1.0   # 改成0.5就是半秒一帧
FRAME_IDX = 0
LAST_SAVE_T = 0.0
FIRST_FRAME_DONE = False
K_SAVED = False

class CameraSaver(Node):
    def __init__(self):
        super().__init__("camera_to_docker")
        self.bridge = CvBridge()
        self.latest_rgb = None
        self.latest_depth = None
        self.create_subscription(Image, "/camera/color/image_raw", self.rgb_cb, 10)
        self.create_subscription(Image, "/camera/depth/image_raw", self.depth_cb, 10)
        self.create_subscription(CameraInfo, "/camera/color/camera_info", self.k_cb, 10)

    def rgb_cb(self, msg):
        self.latest_rgb = self.bridge.imgmsg_to_cv2(msg, "bgr8")
        self.check_save()

    def depth_cb(self, msg):
        self.latest_depth = self.bridge.imgmsg_to_cv2(msg, "16UC1")
        self.check_save()

    def k_cb(self, msg):
        global K_SAVED
        if not K_SAVED:
            K = np.array(msg.k).reshape(3, 3)
            np.savetxt(K_SAVE_PATH, K)
            K_SAVED = True
            self.get_logger().info("✅ 相机内参已保存至K.txt")

    def check_save(self):
        global FRAME_IDX, LAST_SAVE_T, FIRST_FRAME_DONE
        now = time.time()
        # 间隔节流
        if now - LAST_SAVE_T < SAVE_INTERVAL:
            return
        if self.latest_rgb is None or self.latest_depth is None:
            return

        idx = FRAME_IDX
        rgb_p = os.path.join(SAVE_DIR, f"rgb_{idx:06d}.png")
        depth_p = os.path.join(SAVE_DIR, f"depth_{idx:06d}.png")
        cv2.imwrite(rgb_p, self.latest_rgb)
        cv2.imwrite(depth_p, self.latest_depth)

        # 只有第0帧执行分割
        if not FIRST_FRAME_DONE:
            self.get_logger().info(f"📌 第{idx:06d}帧,开始交互式分割")
            mask_p = os.path.join(SAVE_DIR, f"mask_{idx:06d}.png")
            overlay_p = os.path.join(SAVE_DIR, f"overlay_{idx:06d}.png")
            do_segment(self.latest_rgb, mask_p, overlay_p)
            FIRST_FRAME_DONE = True
            self.get_logger().info(f"✅ 首帧{idx:06d}分割完成,后续仅存RGB/Depth")
        else:
            self.get_logger().info(f"✅ 保存帧 {idx:06d} RGB/Depth")

        LAST_SAVE_T = now
        FRAME_IDX += 1

def main():
    rclpy.init()
    node = CameraSaver()
    print(f"✅ ROS2+SAM2时序存储启动,保存间隔{SAVE_INTERVAL}s")
    print("📌 仅第0帧弹窗打点分割,后续序列帧只存图像")
    rclpy.spin(node)
    node.destroy_node()
    rclpy.shutdown()

if __name__ == "__main__":
    main()

用ros2编译后运行该代码。

4. run_custom.py文件修改

代码修改成读取单张图片重建。代码如下:

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto.  Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.
import multiprocessing
multiprocessing.set_start_method('spawn', force=True)
from bundlesdf import *
import argparse
import os
import sys
import cv2
import numpy as np
import trimesh
import copy
import yaml
import glob
import imageio
from segmentation_utils import Segmenter
import time
import shutil

# 扩展:多帧序列读取器(区分首帧/后续帧mask逻辑)
class SequenceImageReader:
    def __init__(self, seq_dir, shorter_side=480):
        """
        序列图片读取器
        :param seq_dir: 序列目录
        :param shorter_side: 图片短边缩放尺寸
        """
        self.seq_dir = seq_dir
        self.shorter_side = shorter_side

        # 匹配6位序号文件 rgb_000000.png / depth_000000.png
        self.color_files = sorted(glob.glob(f"{seq_dir}/rgb_??????.png"),
                                  key=lambda x: int(os.path.basename(x).split('_')[1].split('.')[0]))
        self.depth_files = sorted(glob.glob(f"{seq_dir}/depth_??????.png"),
                                  key=lambda x: int(os.path.basename(x).split('_')[1].split('.')[0]))
        self.id_strs = [f"frame_{i}" for i in range(len(self.color_files))]

        if not self.color_files or not self.depth_files:
            raise Exception(f"序列目录{seq_dir}中未找到rgb_??????/depth_??????文件")
        if len(self.color_files) != len(self.depth_files):
            raise Exception(f"RGB文件数({len(self.color_files)})和Depth文件数({len(self.depth_files)})不匹配")

        # 加载相机内参
        self.K_path = os.path.join(seq_dir, "K.txt")
        self._load_K()
        self.frames = {}

    def _load_K(self):
        """加载并缩放相机内参"""
        if not os.path.exists(self.K_path):
            raise Exception(f"未找到内参文件:{self.K_path}")
        # 先赋值实例属性 self.K
        raw_K = np.loadtxt(self.K_path).reshape(3, 3)
        first_color = self._safe_read_image(self.color_files[0])
        scale = self.shorter_side / min(first_color.shape[:2])
        raw_K[:2, :] *= scale
        self.K = raw_K

    def _safe_read_image(self, path, max_retry=20):
        """安全读取图片,防止文件未写完"""
        for _ in range(max_retry):
            img = cv2.imread(path) if path.endswith('.png') else np.load(path)
            if img is not None:
                return img
            time.sleep(0.05)
        raise Exception(f"无法读取文件:{path}")

    def _load_frame(self, idx):
        if idx in self.frames:
            return self.frames[idx]
        color = self._safe_read_image(self.color_files[idx])
        H0, W0 = color.shape[:2]
        scale = self.shorter_side / min(H0, W0)
        H = int(H0 * scale)
        W = int(W0 * scale)
        color = cv2.resize(color, (W, H), interpolation=cv2.INTER_NEAREST)

        # 强制以单通道模式读取深度图
        depth = cv2.imread(self.depth_files[idx], cv2.IMREAD_UNCHANGED)
        depth = depth.astype(np.float32) / 1000.0
        depth = cv2.resize(depth, (W, H), interpolation=cv2.INTER_NEAREST)

        self.frames[idx] = {
            'color': color,
            'depth': depth,
            'H': H,
            'W': W
        }
        return self.frames[idx]

    def get_color(self, idx):
        return self._load_frame(idx)['color']

    def get_depth(self, idx):
        return self._load_frame(idx)['depth']

    def get_mask(self, idx):
        """
        首帧读取 6位编号掩码 mask_000000.png
        后续帧使用BundleSDF自带追踪
        """
        frame_data = self._load_frame(idx)
        H, W = frame_data['H'], frame_data['W']
        if idx == 0:
            mask_path = os.path.join(self.seq_dir, "mask_000000.png")
            mask = cv2.imread(mask_path, 0)
            if mask is None:
                raise Exception(f"首帧掩码不存在:{mask_path}")
            mask = cv2.resize(mask, (W, H), interpolation=cv2.INTER_NEAREST)
            return mask
        # 后续帧返回全1,交由BundleSDF追踪
        return np.ones((H, W), dtype=np.uint8)

    def get_K(self):
        return self.K.copy()

    def __len__(self):
        return len(self.color_files)


def wait_for_sequence_ready(seq_dir, min_frames=1, timeout=300):
    """
    等待序列就绪,适配6位文件名 + K.txt + mask_000000.png
    增加实时日志,方便排查卡住问题
    """
    start_time = time.time()
    print(f"[INFO] 等待序列目录: {seq_dir} (超时{timeout}s)")
    while True:
        # 核心检测:内参 + 6位掩码 + 序列图片
        has_K = os.path.exists(os.path.join(seq_dir, "K.txt"))
        has_first_mask = os.path.exists(os.path.join(seq_dir, "mask_000000.png"))
        color_files = glob.glob(f"{seq_dir}/rgb_*.png")
        depth_files = glob.glob(f"{seq_dir}/depth_*.png")

        # 实时打印状态(排错用)
        print(f"[DEBUG] K文件:{has_K} | 首帧掩码:{has_first_mask} | 帧数:{len(color_files)}")

        if has_K and has_first_mask and len(color_files) >= min_frames and len(depth_files) >= min_frames:
            print(f"[INFO] ✅ 序列就绪!总帧数: {len(color_files)}")
            break

        if time.time() - start_time > timeout:
            raise TimeoutError(f"等待超时!请检查 2.py 是否正常生成文件")
        time.sleep(1)


def run_sequence_recon(
        seq_dir='/workspace/realtime_input',
        out_folder='/workspace/recon_output',
        use_segmenter=False,
        use_gui=False,
        wait_for_seq=True
):
    import os
    import shutil  # 加上这一行
    set_seed(0)
    # 安全创建/清空目录,替换 os.system
    if os.path.exists(out_folder):
        shutil.rmtree(out_folder)
    os.makedirs(out_folder, exist_ok=True)

    if wait_for_seq:
        wait_for_sequence_ready(seq_dir)

    code_dir = os.path.dirname(os.path.realpath(__file__))
    cfg_bundletrack = yaml.load(open(os.path.join(code_dir, "BundleTrack", "config_ho3d.yml"), 'r'), Loader=yaml.FullLoader)
    cfg_bundletrack['SPDLOG'] = int(args.debug_level)
    cfg_bundletrack['depth_processing']["percentile"] = 95
    cfg_bundletrack['erode_mask'] = 1
    cfg_bundletrack['debug_dir'] = out_folder
    cfg_bundletrack['bundle']['max_BA_frames'] = 20
    cfg_bundletrack['bundle']['max_optimized_feature_loss'] = 0.05
    cfg_bundletrack['feature_corres']['max_dist_neighbor'] = 0.06
    cfg_bundletrack['feature_corres']['max_normal_neighbor'] = 45
    cfg_bundletrack['feature_corres']['max_dist_no_neighbor'] = 0.04
    cfg_bundletrack['feature_corres']['max_normal_no_neighbor'] = 35
    cfg_bundletrack['feature_corres']['map_points'] = True
    cfg_bundletrack['feature_corres']['resize'] = 400
    cfg_bundletrack['feature_corres']['rematch_after_nerf'] = True
    cfg_bundletrack['keyframe']['min_rot'] = 8
    cfg_bundletrack['ransac']['inlier_dist'] = 0.03
    cfg_bundletrack['ransac']['inlier_normal_angle'] = 35
    cfg_bundletrack['ransac']['max_trans_neighbor'] = 0.02
    cfg_bundletrack['ransac']['max_rot_deg_neighbor'] = 30
    cfg_bundletrack['ransac']['max_trans_no_neighbor'] = 0.01
    cfg_bundletrack['ransac']['max_rot_no_neighbor'] = 10
    cfg_bundletrack['p2p']['max_dist'] = 0.04
    cfg_bundletrack['p2p']['max_normal_angle'] = 55

    cfg_track_dir = os.path.join(out_folder, "config_bundletrack.yml")
    yaml.dump(cfg_bundletrack, open(cfg_track_dir, 'w'))

    # ========= 关键修复:此处不提前读取nerf配置 =========
    cfg_track_dir = os.path.join(out_folder, "config_bundletrack.yml")
    yaml.dump(cfg_bundletrack, open(cfg_track_dir, 'w'))

    # 修复:加载并补齐nerf基础配置,防止KeyError
    code_dir = os.path.dirname(os.path.realpath(__file__))
    base_nerf_cfg_file = os.path.join(code_dir, "config.yml")
    with open(base_nerf_cfg_file, 'r') as f:
        cfg_nerf_base = yaml.load(f, Loader=yaml.FullLoader)

    nerf_online_dir = os.path.join(out_folder, "nerf_with_bundletrack_online")
    os.makedirs(nerf_online_dir, exist_ok=True)
    # 填充缺失关键key
    cfg_nerf_base['datadir'] = nerf_online_dir
    cfg_nerf_base['save_dir'] = nerf_online_dir
    cfg_nerf_base['dbscan_eps'] = 0.02
    cfg_nerf_base['dbscan_eps_min_samples'] = 15

    temp_nerf_cfg = os.path.join(out_folder, "init_nerf_config.yml")
    yaml.dump(cfg_nerf_base, open(temp_nerf_cfg, 'w'))

    tracker = BundleSdf(
        cfg_track_dir=cfg_track_dir,
        cfg_nerf_dir=temp_nerf_cfg,
        start_nerf_keyframes=1,
        use_gui=use_gui
    )

    # 初始化序列读取器,开始逐帧处理
    reader = SequenceImageReader(seq_dir)
    total_frames = len(reader)
    print(f"[INFO] 开始处理 {total_frames} 帧序列")

    for i in range(total_frames):
        print(f"\n[INFO] 处理第 {i+1}/{total_frames} 帧")
        color = reader.get_color(i)
        depth = reader.get_depth(i)
        H, W = depth.shape[:2]
        K = reader.get_K()
        id_str = reader.id_strs[i]
        pose_in_model = np.eye(4)

        if use_segmenter:
            mask = reader.get_mask(i)
        else:
            mask = reader.get_mask(i)

        # 仅首帧做掩码腐蚀
        if cfg_bundletrack['erode_mask'] > 0 and i == 0:
            k_size = cfg_bundletrack['erode_mask']
            kernel = np.ones((k_size, k_size), np.uint8)
            mask = cv2.erode(mask.astype(np.uint8), kernel)

        tracker.run(
            color=color,
            depth=depth,
            K=K,
            id_str=id_str,
            mask=mask,
            occ_mask=None,
            pose_in_model=pose_in_model
        )

    # 所有帧处理完毕,再执行全局NeRF与后处理
    import shutil
    nerf_root_dir = os.path.join(out_folder, "nerf_with_bundletrack_online")
    # 拿frame_0里已经生成好的config.yml复制到顶层
    src_cfg = os.path.join(out_folder, "frame_0", "nerf", "config.yml")
    dst_cfg = os.path.join(nerf_root_dir, "config.yml")
    os.makedirs(nerf_root_dir, exist_ok=True)
    if os.path.exists(src_cfg):
        shutil.copy(src_cfg, dst_cfg)
        print(f"[INFO] 复制nerf配置 {src_cfg} → {dst_cfg}")
    else:
        # 兜底用初始配置
        base_cfg = os.path.join(out_folder, "init_nerf_config.yml")
        shutil.copy(base_cfg, dst_cfg)

    tracker.on_finish()
    run_sequence_global_nerf(out_folder, total_frames)
    postprocess_mesh(out_folder)

    print(f"\n✅ 重建完成!结果目录: {os.path.join(out_folder, 'mesh')}")


def run_sequence_global_nerf(out_folder='/workspace/recon_output', total_frames=1):
    set_seed(0)
    cfg_track_dir = os.path.join(out_folder, "config_bundletrack.yml")
    cfg_bundletrack = yaml.load(open(cfg_track_dir, 'r'), Loader=yaml.FullLoader)
    cfg_bundletrack['debug_dir'] = out_folder
    yaml.dump(cfg_bundletrack, open(cfg_track_dir, 'w'))

    # 正确读取nerf子目录下的配置文件
    nerf_dir = os.path.join(out_folder, "nerf_with_bundletrack_online")
    nerf_cfg_path = os.path.join(nerf_dir, "config.yml")
    cfg_nerf = yaml.load(open(nerf_cfg_path, 'r'), Loader=yaml.FullLoader)

    cfg_nerf['n_step'] = 5000 if total_frames > 1 else 2000
    cfg_nerf['N_samples'] = 64
    cfg_nerf['N_samples_around_depth'] = 256
    cfg_nerf['first_frame_weight'] = 1
    cfg_nerf['down_scale_ratio'] = 1
    cfg_nerf['finest_res'] = 256
    cfg_nerf['num_levels'] = 16
    cfg_nerf['mesh_resolution'] = 0.002
    cfg_nerf['n_train_image'] = total_frames
    cfg_nerf['fs_sdf'] = 0.1
    cfg_nerf['frame_features'] = 2
    cfg_nerf['rgb_weight'] = 100
    cfg_nerf['i_img'] = np.inf
    cfg_nerf['i_mesh'] = np.inf
    cfg_nerf['i_nerf_normals'] = np.inf
    cfg_nerf['i_save_ray'] = np.inf

    cfg_nerf['datadir'] = nerf_dir
    cfg_nerf['save_dir'] = copy.deepcopy(cfg_nerf['datadir'])
    os.makedirs(cfg_nerf['datadir'], exist_ok=True)
    cfg_nerf_dir = os.path.join(cfg_nerf['datadir'], "config.yml")
    yaml.dump(cfg_nerf, open(cfg_nerf_dir, 'w'))

    tracker = BundleSdf(cfg_track_dir=cfg_track_dir, cfg_nerf_dir=cfg_nerf_dir, start_nerf_keyframes=1)
    tracker.cfg_nerf = cfg_nerf
    #tracker.run_global_nerf(reader=None, get_texture=True, tex_res=512)
    tracker.run_global_nerf(reader=None, get_texture=False, tex_res=512)
    tracker.on_finish()
    print("[INFO] 全局NeRF优化完成!")


def postprocess_mesh(out_folder):
    mesh_dir = os.path.join(out_folder, "mesh")
    os.makedirs(mesh_dir, exist_ok=True)  # 新建mesh文件夹,不存在就创建
    mesh_files = sorted(glob.glob(f'{out_folder}/**/nerf/*normalized_space.obj', recursive=True))
    if not mesh_files:
        raise FileNotFoundError("未生成mesh文件")
    print(f"[INFO] 加载mesh: {mesh_files[-1]}")
    mesh = trimesh.load(mesh_files[-1])
    with open(os.path.join(os.path.dirname(mesh_files[-1]), 'config.yml'), 'r') as ff:
        cfg = yaml.load(ff, Loader=yaml.FullLoader)
    tf = np.eye(4)
    tf[:3, 3] = cfg['translation']
    tf1 = np.eye(4)
    tf1[:3, :3] *= cfg['sc_factor']
    tf = tf1 @ tf
    mesh.apply_transform(np.linalg.inv(tf))
    mesh.export(os.path.join(out_folder, "mesh", "mesh_real_scale.obj"))

    components = trimesh_split(mesh, min_edge=10) # 1000
    best_component = None
    best_size = 0
    for component in components:
        if len(component.vertices) > best_size:
            best_size = len(component.vertices)
            best_component = component
    mesh = trimesh_clean(best_component)
    mesh.export(os.path.join(out_folder, "mesh", "mesh_biggest_component.obj"))

    mesh = trimesh.smoothing.filter_laplacian(mesh, lamb=0.5, iterations=3)
    mesh.export(os.path.join(out_folder, "mesh", "mesh_biggest_component_smoothed.obj"))


def draw_pose(args):
    K = np.loadtxt(os.path.join(args.out_folder, "cam_K.txt")).reshape(3, 3)
    color_files = sorted(glob.glob(f'{args.out_folder}/color/*'))
    mesh = trimesh.load(os.path.join(args.out_folder, "textured_mesh.obj"))
    to_origin, extents = trimesh.bounds.oriented_bounds(mesh)
    bbox = np.stack([-extents/2, extents/2], axis=0).reshape(2, 3)
    out_dir = os.path.join(args.out_folder, "pose_vis")
    os.makedirs(out_dir, exist_ok=True)
    for color_file in color_files:
        color = imageio.imread(color_file)
        pose = np.loadtxt(color.replace('.png','.txt').replace('color','ob_in_cam'))
        pose = pose @ np.linalg.inv(to_origin)
        vis = draw_posed_3d_box(K, color, ob_in_cam=pose, bbox=bbox)
        id_str = os.path.basename(color_file).replace('.png','')
        imageio.imwrite(os.path.join(out_dir, f"{id_str}.png"), vis)


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument('--mode', type=str, default="run_sequence")
    parser.add_argument('--seq_dir', type=str, default="/workspace/realtime_input")
    parser.add_argument('--out_folder', type=str, default="/workspace/recon_output")
    parser.add_argument('--use_segmenter', type=int, default=0)
    parser.add_argument('--use_gui', type=int, default=0)
    parser.add_argument('--debug_level', type=int, default=2)
    parser.add_argument('--no_wait', action='store_true')
    args = parser.parse_args()

    if args.mode == 'run_sequence':
        run_sequence_recon(
            seq_dir=args.seq_dir,
            out_folder=args.out_folder,
            use_segmenter=args.use_segmenter,
            use_gui=args.use_gui,
            wait_for_seq=not args.no_wait
        )
    elif args.mode == 'global_refine':
        color_files = sorted(glob.glob(f"{args.seq_dir}/rgb_??????.png"))
        total_frames = len(color_files)
        run_sequence_global_nerf(args.out_folder, total_frames)
    elif args.mode == 'draw_pose':
        draw_pose(args)
    else:
        raise RuntimeError("不支持的模式")

6.ros2获取rgbd数据

代码如下:

#!/usr/bin/env python3
import rclpy
from rclpy.node import Node
from sensor_msgs.msg import Image, CameraInfo
from cv_bridge import CvBridge
import cv2
import numpy as np
import os
import sys
import time
import matplotlib
import matplotlib.pyplot as plt

matplotlib.use('TkAgg')
plt.rcParams['font.sans-serif'] = ['DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False

os.environ["QT_LOGGING_RULES"] = "false"
os.environ["OPENCV_LOG_LEVEL"] = "FATAL"

SAM_WORK_DIR = "/home/wyq/BundleSDF-master/sam2-main"
os.chdir(SAM_WORK_DIR)

from sam2.build_sam import build_sam2
from sam2.sam2_image_predictor import SAM2ImagePredictor

cfg_path = "configs/sam2.1/sam2.1_hiera_b+.yaml"
ckpt_path = "weights/sam2.1_hiera_base_plus.pt"
sam2_model = build_sam2(cfg_path, ckpt_path, device="cuda")
predictor = SAM2ImagePredictor(sam2_model)
print("✅ SAM2 加载成功")

points = []
labels = []

def onclick(event):
    global points, labels
    if event.xdata is None or event.ydata is None:
        return
    x = int(event.xdata)
    y = int(event.ydata)
    if event.button == 1:
        print(f"🟢 目标点:({x}, {y})")
        points.append([x, y])
        labels.append(1)
    elif event.button == 3:
        print(f"🔴 背景点:({x}, {y})")
        points.append([x, y])
        labels.append(0)

def smooth_mask_with_gradient(img_bgr, clean_mask, grad_thresh=5, epsilon=0.7):
    gray = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)
    cnt_smooth, _ = cv2.findContours(clean_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
    smooth_mask = np.zeros_like(clean_mask)
    if not cnt_smooth:
        return clean_mask

    main_cnt = max(cnt_smooth, key=cv2.contourArea)
    h_img, w_img = gray.shape[:2]
    fix_cont = []
    for pt in main_cnt:
        x, y = int(pt[0][0]), int(pt[0][1])
        y_up = max(0, y - 1)
        y_down = min(h_img - 1, y + 1)
        x_left = max(0, x - 1)
        x_right = min(w_img - 1, x + 1)

        g_up = gray[y_up, x]
        g_down = gray[y_down, x]
        g_left = gray[y, x_left]
        g_right = gray[y, x_right]
        g_curr = gray[y, x]

        grad_up = abs(int(g_up) - int(g_curr))
        grad_down = abs(int(g_down) - int(g_curr))
        grad_left = abs(int(g_left) - int(g_curr))
        grad_right = abs(int(g_right) - int(g_curr))
        total_grad = grad_up + grad_down + grad_left + grad_right

        dirs = [[0,-1,clean_mask[y_up,x]], [0,1,clean_mask[y_down,x]],[-1,0,clean_mask[y,x_left]], [1,0,clean_mask[y,x_right]]]
        bg_dir = None
        for dx,dy,val in dirs:
            if val == 0:
                bg_dir = (dx, dy)
                break
        if bg_dir is not None:
            x += bg_dir[0]
            y += bg_dir[1]
        else:
            if total_grad < grad_thresh:
                max_grad = max(grad_up,grad_down,grad_left,grad_right)
                if max_grad == grad_up: y -=1
                elif max_grad == grad_down: y +=1
                elif max_grad == grad_left: x -=1
                else: x +=1
        fix_cont.append([[x,y]])

    fix_cont = np.array(fix_cont, dtype=np.int32)
    smooth_cont = cv2.approxPolyDP(fix_cont, epsilon, closed=True)
    cv2.drawContours(smooth_mask, [smooth_cont], -1,255,-1)
    smooth_mask = cv2.GaussianBlur(smooth_mask,(3,3),0.3)
    final_out = (smooth_mask>127).astype(np.uint8)*255
    return final_out

def do_segment(img_bgr, save_mask, save_overlay):
    global points, labels
    points.clear()
    labels.clear()
    if img_bgr is None:
        print("❌ 图像为空")
        return
    H, W = img_bgr.shape[:2]
    img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)
    img_rgb = cv2.medianBlur(img_rgb, 5)
    img_rgb = cv2.GaussianBlur(img_rgb, (3, 3), sigmaX=0.8)

    plt.figure(figsize=(10, 7))
    plt.imshow(img_rgb)
    plt.title("Left click: Foreground | Right click: Background")
    plt.axis("off")
    plt.gcf().canvas.mpl_connect('button_press_event', onclick)
    plt.show()

    if len(points) == 0 or len(points) != len(labels):
        print("❌ 点位数量不匹配,跳过分割")
        return

    pts_np = np.array(points, dtype=np.float32)
    lab_np = np.array(labels, dtype=np.int32)

    print("\n🔍 第一步:全图粗分割")
    predictor.set_image(img_rgb)
    masks1, scores1, _ = predictor.predict(point_coords=pts_np,point_labels=lab_np,multimask_output=True,return_logits=False)
    best1_idx = np.argmax(scores1)
    rough_mask = (masks1[best1_idx] * 255).astype(np.uint8)

    cnts, _ = cv2.findContours(rough_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
    if len(cnts) == 0:
        crop_rgb = img_rgb
        off_x, off_y = 0, 0
    else:
        max_cnt = max(cnts, key=cv2.contourArea)
        x, y, w, h = cv2.boundingRect(max_cnt)
        pad = 15
        x1 = max(0, x - pad)
        y1 = max(0, y - pad)
        x2 = min(W, x + w + pad)
        y2 = min(H, y + h + pad)
        off_x, off_y = x1, y1
        crop_rgb = img_rgb[y1:y2, x1:x2]
        pts_np[:, 0] -= off_x
        pts_np[:, 1] -= off_y

    print("🔍 第二步:裁剪区域精细分割")
    predictor.set_image(crop_rgb)
    masks2, scores2, _ = predictor.predict(point_coords=pts_np,point_labels=lab_np,multimask_output=True,return_logits=False)
    best2_idx = np.argmax(scores2)
    crop_mask = (masks2[best2_idx] * 255).astype(np.uint8)

    final_mask = np.zeros((H, W), dtype=np.uint8)
    final_mask[off_y:off_y+crop_mask.shape[0], off_x:off_x+crop_mask.shape[1]] = crop_mask

    kernel_close = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(3,3))
    kernel_erode = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(2,2))
    final_mask = cv2.morphologyEx(final_mask, cv2.MORPH_CLOSE, kernel_close, iterations=1)
    final_mask = cv2.morphologyEx(final_mask, cv2.MORPH_ERODE, kernel_erode, iterations=1)

    cnt_all,_ = cv2.findContours(final_mask,cv2.RETR_EXTERNAL,cv2.CHAIN_APPROX_SIMPLE)
    clean_mask = np.zeros_like(final_mask)
    if len(cnt_all)>0:
        maxc = max(cnt_all,key=cv2.contourArea)
        cv2.drawContours(clean_mask,[maxc],-1,255,-1)
    clean_mask = smooth_mask_with_gradient(img_bgr, clean_mask, grad_thresh=8, epsilon=1.0)

    overlay = img_bgr.copy()
    roi = overlay[clean_mask > 0]
    if roi.size > 0:
        color_fill = np.full_like(roi, (0, 0, 255), dtype=np.uint8)
        overlay[clean_mask > 0] = cv2.addWeighted(roi, 0.3, color_fill, 0.7, 0)

    cv2.imwrite(save_mask, clean_mask)
    cv2.imwrite(save_overlay, overlay)
    print(f"\n🎉 分割完成!mask:{save_mask} | overlay:{save_overlay}")

# 切回主目录
os.chdir("/home/wyq/BundleSDF-master")
SAVE_DIR = "./realtime_input"
os.makedirs(SAVE_DIR, exist_ok=True)
K_SAVE_PATH = os.path.join(SAVE_DIR, "K.txt")

# 时序保存配置
SAVE_INTERVAL = 1.0   # 改成0.5就是半秒一帧
FRAME_IDX = 0
LAST_SAVE_T = 0.0
FIRST_FRAME_DONE = False
K_SAVED = False

class CameraSaver(Node):
    def __init__(self):
        super().__init__("camera_to_docker")
        self.bridge = CvBridge()
        self.latest_rgb = None
        self.latest_depth = None
        self.create_subscription(Image, "/camera/color/image_raw", self.rgb_cb, 10)
        self.create_subscription(Image, "/camera/depth/image_raw", self.depth_cb, 10)
        self.create_subscription(CameraInfo, "/camera/color/camera_info", self.k_cb, 10)

    def rgb_cb(self, msg):
        self.latest_rgb = self.bridge.imgmsg_to_cv2(msg, "bgr8")
        self.check_save()

    def depth_cb(self, msg):
        self.latest_depth = self.bridge.imgmsg_to_cv2(msg, "16UC1")
        self.check_save()

    def k_cb(self, msg):
        global K_SAVED
        if not K_SAVED:
            K = np.array(msg.k).reshape(3, 3)
            np.savetxt(K_SAVE_PATH, K)
            K_SAVED = True
            self.get_logger().info("✅ 相机内参已保存至K.txt")

    def check_save(self):
        global FRAME_IDX, LAST_SAVE_T, FIRST_FRAME_DONE
        now = time.time()
        # 间隔节流
        if now - LAST_SAVE_T < SAVE_INTERVAL:
            return
        if self.latest_rgb is None or self.latest_depth is None:
            return

        idx = FRAME_IDX
        rgb_p = os.path.join(SAVE_DIR, f"rgb_{idx:06d}.png")
        depth_p = os.path.join(SAVE_DIR, f"depth_{idx:06d}.png")
        cv2.imwrite(rgb_p, self.latest_rgb)
        cv2.imwrite(depth_p, self.latest_depth)

        # 只有第0帧执行分割
        if not FIRST_FRAME_DONE:
            self.get_logger().info(f"📌 第{idx:06d}帧,开始交互式分割")
            mask_p = os.path.join(SAVE_DIR, f"mask_{idx:06d}.png")
            overlay_p = os.path.join(SAVE_DIR, f"overlay_{idx:06d}.png")
            do_segment(self.latest_rgb, mask_p, overlay_p)
            FIRST_FRAME_DONE = True
            self.get_logger().info(f"✅ 首帧{idx:06d}分割完成,后续仅存RGB/Depth")
        else:
            self.get_logger().info(f"✅ 保存帧 {idx:06d} RGB/Depth")

        LAST_SAVE_T = now
        FRAME_IDX += 1

def main():
    rclpy.init()
    node = CameraSaver()
    print(f"✅ ROS2+SAM2时序存储启动,保存间隔{SAVE_INTERVAL}s")
    print("📌 仅第0帧弹窗打点分割,后续序列帧只存图像")
    rclpy.spin(node)
    node.destroy_node()
    rclpy.shutdown()

if __name__ == "__main__":
    main()

打开获取图片的代码,鼠标左键点击目标物体,右键点击背景,会对第一张图片分割mask图片,保存图片后再运行run_custom.py文件。

5. 可视化查看obj文件

将obj文件转换为ply文件,代码如下:

import trimesh
src_path = "/home/wyq/BundleSDF-master/recon_output/mesh/mesh_biggest_component_smoothed.obj"
ply_path = "/home/wyq/model.ply"

mesh = trimesh.load(src_path)
mesh.merge_vertices()
mesh.remove_unreferenced_vertices()
# 导出ply格式
mesh.export(ply_path, file_type="ply")
print("PLY模型生成完成:", ply_path)
print(f"顶点:{len(mesh.vertices)} 面片:{len(mesh.faces)}")

运行以上代码后,终端再用 meshlab 打开 ply 文件:

meshlab /home/wyq/model.ply

4.编译

编译如下:

cd ~/ros2_ws
colcon build --packages-select camera_info_get
source install/setup.bash
ros2 run camera_info_get get_camera_info

或者直接运行

conda activate ros2_tensorrt
cd ~/ros2_ws
source install/setup.bash
python -m camera_info_get.get_camera_info

5.foundationpose在项目中运行

该代码集成了ros2环境下获取深度相机数据、yolo分割mask、foundationpose获取姿态、将姿态转换到基坐标系、姿态发布。

代码如下:

#!/usr/bin/env python3
import sys
sys.path.insert(0, "/home/wyq/ros2\_ws/FoundationPose2/mycpp/build")
# 添加FoundationPose路径
sys.path.insert(0, "/home/wyq/ros2_ws/FoundationPose2")
import mycpp
import rclpy
from rclpy.node import Node
from sensor_msgs.msg import Image, CameraInfo, PointCloud2, PointField
from geometry_msgs.msg import PoseStamped, Pose, Point, Quaternion
from cv_bridge import CvBridge
from std_msgs.msg import Header
import cv2
import numpy as np
import struct
import time
import sys
import os
import json
import base64
import socket
import threading
import math
from typing import Tuple, Optional


# 导入FoundationPose
try:
    from estimater import FoundationPose
    import trimesh
    import torch

    print("✅ FoundationPose 导入成功")
except ImportError as e:
    print(f"❌ FoundationPose 导入失败: {e}")
    sys.exit(1)

from ultralytics import YOLO
import open3d as o3d
from scipy.spatial.transform import Rotation
from com_interfaces.action import VisionDetection
from rclpy.action import ActionServer, GoalResponse, CancelResponse
from rclpy.action.server import ServerGoalHandle

# 屏蔽无关警告
os.environ["QT_LOGGING_RULES"] = "qt.fonts.warning=false"
os.environ["OPENCV_LOG_LEVEL"] = "FATAL"
os.environ["CV_LOG_LEVEL"] = "FATAL"

# ============================================================================
# 📋 可调参数配置区域
# ============================================================================

# -------------------- FoundationPose配置 --------------------
FOUNDATIONPOSE_CONFIG = {
    "mesh_path": "/home/wyq/ros2_ws/phones.obj",  # 3D模型路径
    "device": "cuda",  # 或 "cpu"
    "input_width": 1280,
    "input_height": 720,
    "iteration_register": 10,  # 迭代次数
    "scale_mesh": 1.0,  # 模型缩放
    "debug": 0,  # 调试模式
}

# -------------------- YOLO模型配置 --------------------
YOLO_MODEL_PATH = "/home/wyq/ros2_ws/weights/best5.pt"
YOLO_CONFIDENCE_THRESHOLD = 0.5

# -------------------- 相机话题配置 --------------------
COLOR_TOPIC = "/camera/color/image_raw"
DEPTH_TOPIC = "/camera/depth/image_raw"
CAMERA_INFO_TOPIC = "/camera/color/camera_info"
END_EFFECTOR_POSE_TOPIC = "/end_effector_pose"

# -------------------- 发布话题配置 --------------------
POINTCLOUD_TOPIC = "/detection/pointcloud"
RESULT_IMAGE_TOPIC = "/detection/result_image"
FINAL_POSE_TOPIC = "/detection/final_pose"
ACTION_NAME = "/vision/detection_pose_cloud"

# -------------------- 处理间隔配置 --------------------
PROCESS_INTERVAL = 0.3

# -------------------- 深度图配置 --------------------
DEPTH_MIN = 0.1  # 0.1米
DEPTH_MAX = 2.0  # 2.0米
DEPTH_SCALE = 1000.0  # 毫米转米

# -------------------- 点云配置 --------------------
MIN_MASK_PIXELS = 500
SAMPLE_MAX_POINTS = 20000

# -------------------- 形态学操作配置 --------------------
MORPH_KERNEL_SIZE = 5
MORPH_DILATE_ITERATIONS = 2

# -------------------- 坐标系可视化配置 --------------------
AXIS_LENGTH = 0.15
SHOW_BASE_FRAME = True
SHOW_WORLD_ORIGIN = True
ENABLE_3D_VISUALIZATION = False
WINDOW_WIDTH = 800
WINDOW_HEIGHT = 600

# -------------------- 手眼标定外参配置 --------------------
HAND_EYE_ROTATION = np.array([
    [0.00271188, -0.01911421, 0.99981363],
    [-0.99982551, -0.01853091, 0.00235764],
    [0.0184824, -0.99964556, -0.01916113]
])
HAND_EYE_TRANSLATION = np.array([0.00730634, 0.00320815, 0.00048779]).reshape(3, 1)
HAND_EYE_MATRIX = np.eye(4)
HAND_EYE_MATRIX[:3, :3] = HAND_EYE_ROTATION
HAND_EYE_MATRIX[:3, 3:4] = HAND_EYE_TRANSLATION
ENABLE_HAND_EYE_TRANSFORM = True
M_TO_MM = 1000.0

# -------------------- 姿态偏移参数 --------------------
x_pianyi = -30.0
y_pianyi = 20.0
z_pianyi = 0.0


# ============================================================================
# 辅助函数
# ============================================================================

def euler_zyx_to_quaternion(roll: float, pitch: float, yaw: float, degrees: bool = False) -> np.ndarray:
    """将欧拉角 (ZYX 顺序) 转换为四元数"""
    if degrees:
        roll = math.radians(roll)
        pitch = math.radians(pitch)
        yaw = math.radians(yaw)

    # 根据你的需求调整
    yaw = math.pi / 2 - yaw if yaw > -math.pi / 2 else math.pi / 2 - yaw - math.pi

    cr = math.cos(roll * 0.5)
    sr = math.sin(roll * 0.5)
    cp = math.cos(pitch * 0.5)
    sp = math.sin(pitch * 0.5)
    cy = math.cos(yaw * 0.5)
    sy = math.sin(yaw * 0.5)

    w = cr * cp * cy + sr * sp * sy
    x = sr * cp * cy - cr * sp * sy
    y = cr * sp * cy + sr * cp * sy
    z = cr * cp * sy - sr * sp * cy

    return np.array([x, y, z, w])


def create_point_cloud(points_3d, colors, frame_id, clock):
    """创建ROS点云消息"""
    if len(points_3d) == 0:
        return PointCloud2()

    cloud_msg = PointCloud2()
    cloud_msg.header = Header()
    cloud_msg.header.stamp = clock.now().to_msg()
    cloud_msg.header.frame_id = frame_id
    cloud_msg.height = 1
    cloud_msg.width = len(points_3d)
    cloud_msg.is_bigendian = False
    cloud_msg.is_dense = True
    cloud_msg.fields = [
        PointField(name='x', offset=0, datatype=PointField.FLOAT32, count=1),
        PointField(name='y', offset=4, datatype=PointField.FLOAT32, count=1),
        PointField(name='z', offset=8, datatype=PointField.FLOAT32, count=1),
        PointField(name='rgb', offset=12, datatype=PointField.UINT32, count=1),
    ]
    cloud_msg.point_step = 16
    cloud_msg.row_step = cloud_msg.point_step * len(points_3d)

    data = []
    for pt, col in zip(points_3d, colors):
        rgb = (int(col[2]) << 16) | (int(col[1]) << 8) | int(col[0])
        data.append(struct.pack('ffff', pt[0], pt[1], pt[2], float(rgb)))
    cloud_msg.data = b''.join(data)
    return cloud_msg


# ============================================================================
# FoundationPose 6D位姿估计器
# ============================================================================
class FoundationPoseEstimator:
    """使用FoundationPose进行6D位姿估计"""

    def __init__(self, config):
        self.config = config
        self.fp = None
        self.device = None
        self.initialized = False

        # 初始化FoundationPose
        self._initialize()

    def _initialize(self):
        """初始化FoundationPose"""
        try:
            print("📦 加载3D模型...")
            mesh_path = self.config["mesh_path"]

            if not os.path.exists(mesh_path):
                print(f"❌ 模型文件不存在: {mesh_path}")
                return

            mesh = trimesh.load(mesh_path)
            mesh.vertices /= self.config["scale_mesh"]
            model_pts = mesh.vertices.astype(np.float32)
            model_normals = mesh.vertex_normals.astype(np.float32)

            print(f"✅ 模型加载成功,顶点数: {len(model_pts)}")
            print(f"   模型范围: x [{model_pts[:, 0].min():.3f}, {model_pts[:, 0].max():.3f}]")
            print(f"   模型范围: y [{model_pts[:, 1].min():.3f}, {model_pts[:, 1].max():.3f}]")
            print(f"   模型范围: z [{model_pts[:, 2].min():.3f}, {model_pts[:, 2].max():.3f}]")

            print("🔧 初始化FoundationPose...")
            self.fp = FoundationPose(
                model_pts=model_pts,
                model_normals=model_normals,
                mesh=mesh,
                symmetry_tfs=None,
                debug=self.config["debug"]
            )

            self.device = torch.device(self.config["device"])
            torch.set_float32_matmul_precision('medium')
            print(f"✅ FoundationPose 初始化成功 (设备: {self.device})")
            self.initialized = True

        except Exception as e:
            print(f"❌ FoundationPose 初始化失败: {e}")
            import traceback
            traceback.print_exc()

    def estimate_pose(self, rgb: np.ndarray, depth: np.ndarray, mask: np.ndarray, K: np.ndarray) -> Optional[
        np.ndarray]:
        """
        使用FoundationPose估计6D位姿
        
        Args:
            rgb: RGB图像 (H, W, 3)
            depth: 深度图 (H, W),单位:米
            mask: 掩码 (H, W)
            K: 相机内参矩阵 (3, 3)

        Returns:
            4x4 位姿矩阵,或 None(如果失败)
        """
        if not self.initialized or self.fp is None:
            print("❌ FoundationPose未初始化")
            return None

        try:
            # 检查mask是否有效
            if np.sum(mask) < self.config.get("min_mask_area", 500):
                print(f"⚠️ mask面积太小: {np.sum(mask)}")
                return None

            # 确保图像尺寸正确
            target_h = self.config["input_height"]
            target_w = self.config["input_width"]

            if rgb.shape[0] != target_h or rgb.shape[1] != target_w:
                rgb = cv2.resize(rgb, (target_w, target_h))
                depth = cv2.resize(depth, (target_w, target_h), interpolation=cv2.INTER_NEAREST)
                mask = cv2.resize(mask, (target_w, target_h), interpolation=cv2.INTER_NEAREST)

            # 调用FoundationPose register
            print("⚡ FoundationPose register...")
            pose_1d = self.fp.register(
                K=K,
                rgb=rgb,
                depth=depth,
                ob_mask=mask,
                iteration=self.config["iteration_register"]
            )

            # 转换为4x4矩阵
            pose_mat = torch.tensor(pose_1d, device=self.device).float()
            pose_np = pose_mat.cpu().numpy()

            # 检查是否为有效位姿
            if np.allclose(pose_np, np.eye(4)):
                print("⚠️ 位姿为单位矩阵,计算失败!")
                return None

            print("✅ FoundationPose 位姿估计成功")
            print(f"   位姿矩阵:\n{pose_np}")

            return pose_np

        except Exception as e:
            print(f"❌ FoundationPose 位姿估计失败: {e}")
            import traceback
            traceback.print_exc()
            return None


# ============================================================================
# 主节点
# ============================================================================
class CameraDetectionNode(Node):
    def __init__(self):
        super().__init__('ros2_foundationpose_detection')
        self.bridge = CvBridge()

        # 图像缓存
        self.rgb_img = None
        self.depth_img = None
        self.K = None
        self.last_process_time = 0
        self.process_interval = PROCESS_INTERVAL
        self.end_effector_pose_matrix = None
        self.end_effector_received = False

        # 结果缓存
        self.latest_points_3d = None
        self.latest_colors = None
        self.latest_pose_matrix = None
        self.latest_pose_6d = None
        self.latest_pose_base = None
        self.latest_bbox = None
        self.latest_detection_img = None

        # 线程锁
        self.pipeline_lock = threading.Lock()

        # 加载手眼标定
        self.load_hand_eye_calibration()

        # 初始化YOLO
        self.yolo = YOLO(YOLO_MODEL_PATH)
        self.yolo.to("cpu")
        self.get_logger().info("✅ YOLO模型加载成功")

        # ============================================================
        # 🔑 初始化FoundationPose
        # ============================================================
        self.get_logger().info("=" * 60)
        self.get_logger().info("🔧 初始化FoundationPose...")
        self.foundationpose = FoundationPoseEstimator(FOUNDATIONPOSE_CONFIG)

        if not self.foundationpose.initialized:
            self.get_logger().error("❌ FoundationPose初始化失败!")
        else:
            self.get_logger().info("✅ FoundationPose初始化成功")
        self.get_logger().info("=" * 60)

        # 创建订阅
        self.create_subscription(Image, COLOR_TOPIC, self.rgb_cb, 10)
        self.create_subscription(Image, DEPTH_TOPIC, self.depth_cb, 10)
        self.create_subscription(CameraInfo, CAMERA_INFO_TOPIC, self.caminfo_cb, 10)
        self.create_subscription(PoseStamped, END_EFFECTOR_POSE_TOPIC, self.end_effector_pose_cb, 10)

        # 创建发布器
        self.pointcloud_pub = self.create_publisher(PointCloud2, POINTCLOUD_TOPIC, 10)
        self.result_image_pub = self.create_publisher(Image, RESULT_IMAGE_TOPIC, 10)
        self.final_pose_pub = self.create_publisher(PoseStamped, FINAL_POSE_TOPIC, 10)

        # 创建Action服务器
        self.action_server = ActionServer(
            node=self,
            action_name=ACTION_NAME,
            action_type=VisionDetection,
            execute_callback=self.execute_callback,
            goal_callback=self.goal_callback,
            cancel_callback=self.cancel_callback,
        )

        self.get_logger().info("=" * 60)
        self.get_logger().info("✅ 节点初始化完成 (使用FoundationPose)")
        self.get_logger().info("=" * 60)

    def load_hand_eye_calibration(self):
        """加载手眼标定参数"""
        self.T_cam2gripper = HAND_EYE_MATRIX.copy()
        self.enable_transform = ENABLE_HAND_EYE_TRANSFORM
        self.get_logger().info("✅ 手眼标定参数加载完成")

    def end_effector_pose_cb(self, msg: PoseStamped):
        """末端位姿回调"""
        try:
            position = np.array([msg.pose.position.x, msg.pose.position.y, msg.pose.position.z])
            quat = np.array([msg.pose.orientation.x, msg.pose.orientation.y,
                             msg.pose.orientation.z, msg.pose.orientation.w])

            rotation = Rotation.from_quat(quat)
            self.end_effector_pose_matrix = np.eye(4)
            self.end_effector_pose_matrix[:3, :3] = rotation.as_matrix()
            self.end_effector_pose_matrix[:3, 3] = position
            self.end_effector_received = True

            self.get_logger().debug(f"📥 末端位姿: ({position[0]:.3f}, {position[1]:.3f}, {position[2]:.3f})")

        except Exception as e:
            self.get_logger().error(f"末端位姿处理错误: {e}")

    def caminfo_cb(self, msg):
        """相机内参回调"""
        self.K = np.array(msg.k).reshape(3, 3)

    def rgb_cb(self, msg):
        """RGB图像回调"""
        self.rgb_img = self.bridge.imgmsg_to_cv2(msg, 'bgr8')

    def depth_cb(self, msg):
        """深度图像回调"""
        try:
            if msg.encoding == "16UC1":
                depth_data = np.ndarray(shape=(msg.height, msg.width), dtype=np.uint16, buffer=msg.data)
            else:
                self.get_logger().warning(f"未知深度编码: {msg.encoding}")
                return

            self.depth_img = depth_data
            self.process()

        except Exception as e:
            self.get_logger().error(f"深度图像处理错误: {e}")

    def process(self):
        """主处理流程"""
        if self.rgb_img is None or self.depth_img is None or self.K is None:
            return

        # 处理间隔控制
        current_time = time.time()
        if current_time - self.last_process_time < self.process_interval:
            return
        self.last_process_time = current_time

        with self.pipeline_lock:
            start_time = time.time()

            rgb = self.rgb_img.copy()
            h, w = rgb.shape[:2]

            # 深度图处理
            depth = self.depth_img.astype(np.float32) / DEPTH_SCALE
            depth = np.clip(depth, DEPTH_MIN, DEPTH_MAX)
            depth[np.isnan(depth)] = DEPTH_MIN
            depth[np.isinf(depth)] = DEPTH_MIN

            # YOLO检测
            res = self.yolo.predict(rgb, conf=YOLO_CONFIDENCE_THRESHOLD, verbose=False)

            if len(res) == 0 or res[0].boxes is None or len(res[0].boxes) == 0:
                return

            # 获取边界框
            box = res[0].boxes[0].xyxy.cpu().numpy()[0]
            x1, y1, x2, y2 = map(int, box)
            bbox = (x1, y1, x2, y2)

            # 获取Mask
            if res[0].masks is None or len(res[0].masks) <= 0:
                self.get_logger().warning("⚠️ 没有分割掩码")
                return

            mask = res[0].masks.data[0].cpu().numpy()
            mask = cv2.resize(mask, (w, h))
            mask = (mask > 0.5).astype(np.uint8)

            # 形态学操作
            kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (MORPH_KERNEL_SIZE, MORPH_KERNEL_SIZE))
            mask = cv2.dilate(mask, kernel, iterations=MORPH_DILATE_ITERATIONS)

            mask_pixels = np.sum(mask)
            if mask_pixels < MIN_MASK_PIXELS:
                self.get_logger().warning(f"⚠️ mask面积太小: {mask_pixels}")
                return

            # ============================================================
            # 🔑 使用FoundationPose估计6D位姿
            # ============================================================
            pose_matrix = self.foundationpose.estimate_pose(rgb, depth, mask, self.K)

            if pose_matrix is None:
                self.get_logger().warning("❌ FoundationPose位姿估计失败")
                return

            # 保存结果
            self.latest_pose_matrix = pose_matrix
            self.latest_bbox = bbox

            # ============================================================
            # 坐标转换:相机 -> 末端 -> 基座
            # ============================================================

            # 1. 相机坐标系 -> 末端坐标系
            if self.enable_transform:
                pose_gripper = self.T_cam2gripper @ pose_matrix
            else:
                pose_gripper = pose_matrix

            # 2. 末端坐标系 -> 基座坐标系
            if self.end_effector_received and self.end_effector_pose_matrix is not None:
                pose_base = self.end_effector_pose_matrix @ pose_gripper
            else:
                pose_base = pose_gripper
                self.get_logger().warning("⚠️ 未收到末端位姿,使用相机坐标系")

            self.latest_pose_base = pose_base

            # 提取位置和姿态
            position = pose_base[:3, 3]
            rotation = pose_base[:3, :3]
            r = Rotation.from_matrix(rotation)
            euler = r.as_euler('zyx', degrees=True)

            self.get_logger().info("=" * 60)
            self.get_logger().info("🎯 FoundationPose 6D位姿结果:")
            self.get_logger().info(
                f"  位置 (mm): ({position[0] * 1000:.1f}, {position[1] * 1000:.1f}, {position[2] * 1000:.1f})")
            self.get_logger().info(f"  姿态 (deg): ({euler[0]:.2f}, {euler[1]:.2f}, {euler[2]:.2f})")
            self.get_logger().info("=" * 60)

            # 生成点云(用于发布和可视化)
            points_3d, colors = self.generate_pointcloud(rgb, depth, mask)
            self.latest_points_3d = points_3d
            self.latest_colors = colors

            # 发布点云
            if points_3d is not None and len(points_3d) > 0:
                self.publish_pointcloud(points_3d, colors)

            # 发布最终位姿
            self.publish_final_pose(pose_base)

            # 显示结果
            self.segment_display_and_publish(rgb, mask, bbox)

            # 3D可视化(可选)
            if ENABLE_3D_VISUALIZATION:
                self.visualize_3d(points_3d, colors, pose_base)

            elapsed = time.time() - start_time
            self.get_logger().info(f"⏱️ 处理耗时: {elapsed:.3f}s")

    def generate_pointcloud(self, rgb, depth, mask):
        """从深度图生成点云"""
        h, w = rgb.shape[:2]
        fx = self.K[0, 0]
        fy = self.K[1, 1]
        cx = self.K[0, 2]
        cy = self.K[1, 2]

        ys, xs = np.where(mask > 0)
        if len(ys) < MIN_MASK_PIXELS:
            return None, None

        points_3d = []
        colors = []

        for v, u in zip(ys, xs):
            z = depth[v, u]
            if z <= 0 or z > DEPTH_MAX:
                continue
            x = (u - cx) * z / fx
            y = (v - cy) * z / fy
            b, g, r = rgb[v, u]
            points_3d.append([x, y, z])
            colors.append([r, g, b])

        if len(points_3d) == 0:
            return None, None

        points_3d = np.array(points_3d, dtype=np.float32)
        colors = np.array(colors, dtype=np.uint8)

        # 降采样
        if len(points_3d) > SAMPLE_MAX_POINTS:
            indices = np.random.choice(len(points_3d), SAMPLE_MAX_POINTS, replace=False)
            points_3d = points_3d[indices]
            colors = colors[indices]

        return points_3d, colors

    def publish_pointcloud(self, points_3d, colors):
        """发布点云"""
        if len(points_3d) == 0:
            return

        try:
            # 转换到基座坐标系
            if self.end_effector_received and self.end_effector_pose_matrix is not None:
                ones = np.ones((len(points_3d), 1))
                points_homogeneous = np.hstack([points_3d, ones])
                points_gripper = (self.T_cam2gripper @ points_homogeneous.T).T
                points_base = (self.end_effector_pose_matrix @ points_gripper.T).T
                points = points_base[:, :3]
                frame_id = "base_link"
            else:
                points = points_3d
                frame_id = "camera_link"

            cloud_msg = create_point_cloud(points, colors, frame_id, self.get_clock())
            self.pointcloud_pub.publish(cloud_msg)

        except Exception as e:
            self.get_logger().error(f"发布点云失败: {e}")

    def publish_final_pose(self, pose_matrix):
        """发布最终位姿"""
        position = pose_matrix[:3, 3]
        rotation = pose_matrix[:3, :3]

        # 转换为四元数
        r = Rotation.from_matrix(rotation)
        quat = r.as_quat()  # [x, y, z, w]

        pose_msg = PoseStamped()
        pose_msg.header.stamp = self.get_clock().now().to_msg()
        pose_msg.header.frame_id = "base_link"

        # 应用偏移
        pose_msg.pose.position.x = position[0] * M_TO_MM + x_pianyi
        pose_msg.pose.position.y = position[1] * M_TO_MM + y_pianyi
        pose_msg.pose.position.z = position[2] * M_TO_MM + z_pianyi

        pose_msg.pose.orientation.x = quat[0]
        pose_msg.pose.orientation.y = quat[1]
        pose_msg.pose.orientation.z = quat[2]
        pose_msg.pose.orientation.w = quat[3]

        self.final_pose_pub.publish(pose_msg)

    def segment_display_and_publish(self, rgb, mask, bbox):
        """显示并发布检测结果"""
        try:
            overlay = rgb.copy()

            # 叠加Mask
            if mask is not None and np.sum(mask) > 0:
                mask_colored = np.zeros_like(rgb)
                mask_colored[:, :, 1] = mask * 255
                overlay = cv2.addWeighted(rgb, 0.7, mask_colored, 0.3, 0)

            # 绘制边界框
            if bbox is not None:
                x1, y1, x2, y2 = bbox
                cv2.rectangle(overlay, (x1, y1), (x2, y2), (0, 255, 0), 2)
                cv2.putText(overlay, "Target", (x1, y1 - 10),
                            cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2)

            # 添加FoundationPose信息
            if self.latest_pose_base is not None:
                pos = self.latest_pose_base[:3, 3]
                text = f"FPose: ({pos[0] * 1000:.0f}, {pos[1] * 1000:.0f}, {pos[2] * 1000:.0f})mm"
                cv2.putText(overlay, text, (10, 60),
                            cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 255), 2)

            cv2.putText(overlay, "FoundationPose 6D Pose", (10, 30),
                        cv2.FONT_HERSHEY_SIMPLEX, 0.7, (0, 255, 255), 2)

            # 发布
            if self.result_image_pub.get_subscription_count() > 0:
                result_msg = self.bridge.cv2_to_imgmsg(overlay, "bgr8")
                result_msg.header.stamp = self.get_clock().now().to_msg()
                result_msg.header.frame_id = "camera_link"
                self.result_image_pub.publish(result_msg)

        except Exception as e:
            self.get_logger().error(f"显示错误: {e}")

    def visualize_3d(self, points_3d, colors, pose_matrix):
        """3D可视化"""
        if points_3d is None or len(points_3d) == 0:
            return

        try:
            # 创建点云
            pcd = o3d.geometry.PointCloud()
            pcd.points = o3d.utility.Vector3dVector(points_3d)
            if colors is not None:
                pcd.colors = o3d.utility.Vector3dVector(colors / 255.0)

            # 创建坐标系
            coord_frame = o3d.geometry.TriangleMesh.create_coordinate_frame(size=0.1)

            # 创建位姿坐标系
            pose_frame = o3d.geometry.TriangleMesh.create_coordinate_frame(size=0.15)
            pose_frame.transform(pose_matrix)

            # 显示
            o3d.visualization.draw_geometries(
                [pcd, coord_frame, pose_frame],
                window_name="FoundationPose 6D Pose Estimation",
                width=WINDOW_WIDTH,
                height=WINDOW_HEIGHT
            )

        except Exception as e:
            self.get_logger().error(f"3D可视化错误: {e}")

    # ========================================================================
    # Action Server 回调
    # ========================================================================

    def goal_callback(self, goal_request):
        self.get_logger().info(f"📥 收到目标请求: {goal_request.target_obj_name}")
        return GoalResponse.ACCEPT

    def cancel_callback(self, goal_handle):
        self.get_logger().info("❌ 任务取消")
        return CancelResponse.ACCEPT

    def execute_callback(self, goal_handle: ServerGoalHandle):
        self.get_logger().info("=" * 60)
        self.get_logger().info("🔴 EXECUTE_CALLBACK 被调用")
        self.get_logger().info("=" * 60)

        result = VisionDetection.Result()

        try:
            feedback = VisionDetection.Feedback()
            feedback.task_status = "Processing..."
            feedback.progress_rate = 0.0
            feedback.current_step_info = "Initializing..."

            if self.latest_pose_base is None:
                self.get_logger().error("❌ 没有可用数据")
                result.success = False
                result.status_code = 1
                result.error_message = "没有可用的检测结果"
                goal_handle.abort()
                return result

            feedback.progress_rate = 0.5
            feedback.current_step_info = "Creating pose message..."
            goal_handle.publish_feedback(feedback)

            # 构建结果
            result.success = True
            result.status_code = 0
            result.error_message = ""
            result.execution_time = time.time() - self.last_process_time
            result.header = Header()
            result.header.stamp = self.get_clock().now().to_msg()
            result.header.frame_id = "base_link"

            if goal_handle.request.need_6d_pose:
                pose_base = self.latest_pose_base
                position = pose_base[:3, 3]
                rotation = pose_base[:3, :3]
                r = Rotation.from_matrix(rotation)
                quat = r.as_quat()

                result.target_6d_pose = PoseStamped()
                result.target_6d_pose.header = Header()
                result.target_6d_pose.header.stamp = self.get_clock().now().to_msg()
                result.target_6d_pose.header.frame_id = "base_link"

                result.target_6d_pose.pose.position.x = position[0] * M_TO_MM + x_pianyi
                result.target_6d_pose.pose.position.y = position[1] * M_TO_MM + y_pianyi
                result.target_6d_pose.pose.position.z = position[2] * M_TO_MM + z_pianyi

                result.target_6d_pose.pose.orientation.x = quat[0]
                result.target_6d_pose.pose.orientation.y = quat[1]
                result.target_6d_pose.pose.orientation.z = quat[2]
                result.target_6d_pose.pose.orientation.w = quat[3]

                result.pose_confidence = 0.95

            if goal_handle.request.need_env_point_cloud and self.latest_points_3d is not None:
                result.env_point_cloud = create_point_cloud(
                    self.latest_points_3d,
                    self.latest_colors if self.latest_colors is not None else
                    np.ones((len(self.latest_points_3d), 3), dtype=np.uint8) * 255,
                    "base_link",
                    self.get_clock()
                )
                result.point_cloud_frame_id = "base_link"

            feedback.progress_rate = 1.0
            feedback.task_status = "Completed"
            feedback.current_step_info = "Done"
            goal_handle.publish_feedback(feedback)

            self.get_logger().info("✅ 返回结果")
            goal_handle.succeed()
            return result

        except Exception as e:
            self.get_logger().error(f"❌ execute_callback 异常: {e}")
            import traceback
            traceback.print_exc()
            result.success = False
            result.status_code = 3
            result.error_message = str(e)
            goal_handle.abort()
            return result


def main(args=None):
    rclpy.init(args=args)
    node = CameraDetectionNode()

    try:
        rclpy.spin(node)
    except KeyboardInterrupt:
        node.get_logger().info("接收到退出信号")
    finally:
        node.destroy_node()
        rclpy.shutdown()


if __name__ == '__main__':
    main()

更多推荐