ROS2+强化学习:机械臂抓取与Sim-to-Real迁移实战

ROS2+强化学习:机械臂抓取与Sim-to-Real迁移实战

一、引言

机器人抓取是"感知→规划→控制"的经典闭环。传统方法依赖精确的物体位姿估计和运动规划,而 RL 提供了端到端的学习范式。本文将 ROS2 + Gymnasium + Isaac Sim 构建完整的机械臂 RL 训练流水线。

二、ROS2基础

python 复制代码
# ROS2节点:发布关节指令 + 订阅传感器
import rclpy
from rclpy.node import Node
from sensor_msgs.msg import JointState, Image
from trajectory_msgs.msg import JointTrajectory, JointTrajectoryPoint

class RobotRLNode(Node):
    def __init__(self):
        super().__init__('robot_rl')
        
        # 发布关节控制
        self.cmd_pub = self.create_publisher(
            JointTrajectory, '/arm_controller/joint_trajectory', 10
        )
        
        # 订阅关节状态 + 相机
        self.create_subscription(JointState, '/joint_states', self.joint_cb, 10)
        self.create_subscription(Image, '/camera/rgb', self.camera_cb, 10)
        
        # 定时器:RL控制循环(50Hz)
        self.timer = self.create_timer(0.02, self.rl_control_loop)
        
        self.agent = PPOAgent()  # 训练好的RL策略
    
    def rl_control_loop(self):
        # 1. 构建观测
        obs = np.concatenate([
            self.joint_positions,
            self.joint_velocities,
            self.target_pose,  # 目标物体位姿
            self.gripper_state
        ])
        
        # 2. RL推理
        action = self.agent.predict(obs)  # [Δx,Δy,Δz,gripper]
        
        # 3. 逆运动学 → 关节角度
        joint_angles = self.ik_solver(action[:3], action[3:6])
        
        # 4. 发布控制指令
        msg = JointTrajectory()
        msg.joint_names = ['joint_1', ..., 'joint_6', 'gripper']
        msg.points = [JointTrajectoryPoint(positions=joint_angles)]
        self.cmd_pub.publish(msg)

三、RL环境定义

python 复制代码
import gymnasium as gym
import numpy as np

class RobotGraspEnv(gym.Env):
    def __init__(self):
        self.observation_space = gym.spaces.Dict({
            "joint_pos": gym.spaces.Box(-np.pi, np.pi, (7,)),  # 6DOF+gripper
            "joint_vel": gym.spaces.Box(-10, 10, (7,)),
            "rgb": gym.spaces.Box(0, 255, (128, 128, 3), dtype=np.uint8),
            "depth": gym.spaces.Box(0, 5, (128, 128, 1)),
            "target_pose": gym.spaces.Box(-1, 1, (3,)),  # 物体相对位置
        })
        
        self.action_space = gym.spaces.Box(-1, 1, (4,))  # [vx,vy,vz,grip]
        
        self.robot = PyBulletRobot()  # 或 Isaac Sim
        self.objects = []
        self._spawn_scene()
    
    def step(self, action):
        # 1. 执行动作
        target_vel = action[:3] * 0.05  # 速度缩放
        grip_cmd = 1 if action[3] > 0 else 0
        
        self.robot.move_ee(target_vel)
        self.robot.set_gripper(grip_cmd)
        
        # 2. 物理仿真步进
        self.sim.step()
        
        # 3. 观测
        obs = self._get_obs()
        
        # 4. 奖励设计
        reward = 0
        
        # 接近奖励
        dist = np.linalg.norm(self.robot.ee_pos - self.target_pos)
        reward -= dist * 0.1  # 惩罚距离
        
        # 抓取奖励
        if self.robot.is_grasping():
            if self._object_lifted():
                reward += 5.0  # 成功举起
            reward += 1.0  # 抓住
        
        # 掉落惩罚
        if self._object_dropped():
            reward -= 3.0
        
        # 5. 终止条件
        done = self._object_lifted() or self._steps > 200
        
        return obs, reward, done, False, {"dist": dist}
    
    def reset(self):
        self._spawn_scene()
        self.robot.reset()
        return self._get_obs(), {}
    
    def _get_obs(self):
        return {
            "joint_pos": self.robot.joint_positions,
            "joint_vel": self.robot.joint_velocities,
            "rgb": self.camera.rgb(),
            "depth": self.camera.depth(),
            "target_pose": self.target_pos - self.robot.ee_pos,
        }

四、PPO训练

python 复制代码
from stable_baselines3 import PPO
from stable_baselines3.common.vec_env import SubprocVecEnv
from stable_baselines3.common.callbacks import CheckpointCallback

def make_env():
    return RobotGraspEnv()

# 并行环境(8进程)
env = SubprocVecEnv([make_env for _ in range(8)])

# 策略网络
policy_kwargs = dict(
    features_extractor_class=CustomCNN,  # 处理RGB+D
    net_arch=dict(pi=[512, 256], vf=[512, 256])
)

model = PPO(
    "MultiInputPolicy",
    env,
    policy_kwargs=policy_kwargs,
    n_steps=2048,
    batch_size=256,
    n_epochs=10,
    learning_rate=3e-4,
    gamma=0.99,
    gae_lambda=0.95,
    clip_range=0.2,
    tensorboard_log="./logs/",
    verbose=1
)

model.learn(
    total_timesteps=5_000_000,
    callback=CheckpointCallback(save_freq=100_000, save_path="./models/")
)
model.save("grasp_policy_final")

五、Sim-to-Real迁移

python 复制代码
class DomainRandomizer:
    """域随机化:缩小仿真与现实差距"""
    
    def randomize(self, env):
        # 物理参数随机化
        env.gravity = [0, 0, random.uniform(-9.6, -9.9)]
        env.friction = random.uniform(0.5, 1.5)
        env.object_mass = random.uniform(0.05, 0.3)
        
        # 视觉随机化
        env.lighting_intensity = random.uniform(0.6, 1.5)
        env.lighting_color = np.random.uniform(0.8, 1.2, 3)
        env.camera_noise = random.uniform(0, 0.02)
        env.background_texture = random.choice(textures)
        
        # 物体位姿随机化
        env.object_pose = [random.uniform(-0.2, 0.2),
                          random.uniform(-0.2, 0.2),
                          random.uniform(0.02, 0.1),
                          random.uniform(0, 2*np.pi)]
        
        return env

# 训练时每episode随机化
def train_with_randomization():
    env = DomainRandomizer().randomize(RobotGraspEnv())
    # ... PPO训练

六、总结

RL机械臂的三大要素:

  1. 环境设计 → 合理的观测/动作空间 + 分层奖励
  2. 域随机化 → 缩小Sim-to-Real gap
  3. ROS2集成 → 50Hz实时控制闭环
相关推荐
quanjui1 小时前
【医学尝试】基于Segment Anything Model的医学图像分割研究:眼底OCT与X线胸片微调实战
人工智能·笔记·学习
就是一顿骚操作1 小时前
Dropout:神经网络正则化的经典解读
人工智能·深度学习·神经网络·论文解读
2401_843253702 小时前
数据隐私AI:从合规检查到隐私计算的Skill化
人工智能
音视频工程实战2 小时前
PromptQL 新手入门与实战指南
数据库·sql·算法
清川渡水2 小时前
NL2SQL 的正确打开方式:从「AI 猜 SQL」到「置信度闭环」的工程化落地
数据库·人工智能·sql
Lyra_Infra2 小时前
OpenClaw 服务异常故障分析报告
linux·人工智能
半个落月2 小时前
用 Node.js 搭建 EPUB 问答助手:从文本切片、向量检索到 RAG
人工智能·node.js
MomentYY2 小时前
RAG 建库:资料是怎么存进去的?
人工智能·agent·ai编程
wabs6662 小时前
关于图论【卡码网104.建造最大岛屿的思考】
数据结构·算法·图论
Java编程爱好者2 小时前
AI Agent、架构决策记录与工程上下文治理:团队如何把隐性约束留在仓库里。
人工智能