diff --git a/src/Neuro_Mujoco/main.py b/src/Neuro_Mujoco/main.py index d729535af0..1c240566dd 100644 --- a/src/Neuro_Mujoco/main.py +++ b/src/Neuro_Mujoco/main.py @@ -1,4 +1,9 @@ #! /usr/bin/env python +# -*- coding: utf-8 -*- +""" +MuJoCo功能整合工具(支持强化学习策略控制 + ROS 1通信) +优化点:新增基于PyTorch的策略网络推理,自动生成控制指令 +""" import os import sys import time @@ -10,7 +15,16 @@ import mujoco from mujoco import viewer -# ===================== ROS 1 相关导入(新增)===================== +# ===================== 机器学习(PyTorch)相关导入 ===================== +try: + import torch + import torch.nn as nn + TORCH_AVAILABLE = True +except ImportError: + TORCH_AVAILABLE = False + logging.warning("未检测到PyTorch,策略控制功能已禁用(安装:pip install torch)") + +# ===================== ROS 1 相关导入 ===================== try: import rospy from sensor_msgs.msg import JointState @@ -19,9 +33,9 @@ ROS_AVAILABLE = True except ImportError: ROS_AVAILABLE = False - logging.warning("未检测到 ROS 环境,ROS 功能已禁用(如需启用,请安装 ROS 1 Noetic 并配置环境)") + logging.warning("未检测到 ROS 环境,ROS 功能已禁用(如需启用,请安装 ROS 1 Noetic)") -# 配置日志系统 +# ===================== 日志系统配置 ===================== logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s", @@ -30,16 +44,31 @@ ) logger = logging.getLogger("mujoco_utils") +# ===================== 强化学习策略网络 ===================== +class PolicyNetwork(nn.Module): + """轻量级策略网络(适用于MuJoCo机器人控制) + 输入:观测(关节位置+速度),输出:归一化控制指令([-1,1]) + """ + def __init__(self, obs_dim: int, action_dim: int, hidden_dim: int = 64): + super().__init__() + self.net = nn.Sequential( + nn.Linear(obs_dim, hidden_dim), + nn.Tanh(), + nn.Linear(hidden_dim, hidden_dim), + nn.Tanh(), + nn.Linear(hidden_dim, action_dim), + nn.Tanh() # 输出范围[-1,1],后续映射到实际控制范围 + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """前向推理(禁用梯度计算以提升效率)""" + with torch.no_grad(): + return self.net(x) +# ===================== 核心功能函数 ===================== def load_model(model_path: str) -> Tuple[Optional[mujoco.MjModel], Optional[mujoco.MjData]]: """ - 加载MuJoCo模型(支持XML和MJB格式) - - 参数: - model_path: 模型文件路径 - - 返回: - 加载成功返回(model, data)元组,失败返回(None, None) + 加载MuJoCo模型(支持XML/MJB格式) """ if not os.path.exists(model_path): logger.error(f"模型文件不存在: {model_path}") @@ -61,12 +90,7 @@ def load_model(model_path: str) -> Tuple[Optional[mujoco.MjModel], Optional[mujo def convert_model(input_path: str, output_path: str) -> bool: """ - 转换模型格式(XML↔MJB) - 参数: - input_path: 输入模型路径 - output_path: 输出模型路径(需指定扩展名.xml或.mjb) - 返回: - 转换成功返回True,失败返回False + 转换模型格式(XML ↔ MJB) """ model, data = load_model(input_path) if not model or not data: @@ -104,13 +128,7 @@ def test_speed( ctrlnoise: float = 0.01 ) -> None: """ - 测试模型模拟速度 - - 参数: - model_path: 模型文件路径 - nstep: 每线程模拟步数 - nthread: 测试线程数 - ctrlnoise: 控制噪声强度 + 测试模型模拟速度(多线程) """ model, _ = load_model(model_path) if not model: @@ -164,9 +182,13 @@ def simulate_thread(thread_id: int) -> float: logger.info(f"实时因子: {realtime_factor:.2f}x") logger.info(f"线程平均耗时: {np.mean(thread_durations):.2f}秒 (±{np.std(thread_durations):.2f})") -# ===================== 可视化函数(仅优化ROS关节状态发布逻辑)===================== -def visualize(model_path: str, use_ros: bool = False) -> None: + +def visualize(model_path: str, use_ros: bool = False, policy_path: Optional[str] = None) -> None: """ + 可视化模型并运行模拟(支持ROS/策略控制) + :param model_path: 模型文件/目录路径 + :param use_ros: 是否启用ROS模式 + :param policy_path: 预训练策略模型路径(.pth) 可视化模型并运行模拟(支持ROS 1模式) 参数: @@ -189,49 +211,79 @@ def visualize(model_path: str, use_ros: bool = False) -> None: model_path: 模型文件路径 use_ros: 是否启用ROS模式(默认False) """ + # 智能校验模型路径(支持目录自动找XML/MJB) + if os.path.isdir(model_path): + model_files = [] + for file in os.listdir(model_path): + if file.endswith(('.xml', '.mjb')): + model_files.append(os.path.join(model_path, file)) + if not model_files: + logger.error(f"目录 {model_path} 中未找到.xml/.mjb模型文件") + return + model_path = model_files[0] + logger.info(f"自动选择目录中的模型文件: {model_path}") + model, data = load_model(model_path) if not model: return - # ===================== ROS 1 初始化(新增)===================== + # ===================== 策略网络初始化 ===================== + policy = None + obs_dim = model.nq + model.nv # 观测维度:关节位置 + 关节速度 + action_dim = model.nu + ctrl_range = None # 控制指令实际范围 + + if policy_path and TORCH_AVAILABLE and action_dim > 0: + try: + # 加载预训练策略模型 + policy = PolicyNetwork(obs_dim, action_dim) + policy.load_state_dict(torch.load(policy_path, map_location=torch.device('cpu'))) + policy.eval() # 推理模式 + logger.info(f"成功加载策略模型: {policy_path}") + + # 获取控制指令范围(映射[-1,1]到实际范围) + ctrl_range = [] + for i in range(action_dim): + if model.actuator_ctrllimited[i]: + ctrl_range.append(model.actuator_ctrlrange[i]) + else: + ctrl_range.append((-1.0, 1.0)) + ctrl_range = np.array(ctrl_range) + except Exception as e: + logger.error(f"策略模型加载失败: {str(e)}", exc_info=True) + policy = None + elif policy_path: + logger.warning("策略功能需满足:PyTorch已安装 + 模型有控制维度(nu>0)") + + # ===================== ROS 初始化 ===================== ros_publishers = None ros_subscribers = None ros_rate = None ctrl_cmd = None joint_msg = None - # 新增:存储非自由关节的ID和对应的qpos/qvel索引(解决索引错位+支持多自由度) joint_ids = [] joint_qpos_idxs = [] joint_qvel_idxs = [] if use_ros: if not ROS_AVAILABLE: - logger.error("ROS 环境未就绪,无法启用 ROS 模式(请检查ROS安装和环境配置)") + logger.error("ROS环境未就绪,无法启用ROS模式") return - # 初始化ROS节点 rospy.init_node("mujoco_ros_node", anonymous=True) - ros_rate = rospy.Rate(100) # 100Hz发布频率(与MuJoCo默认步长0.01s匹配) + ros_rate = rospy.Rate(100) # 100Hz匹配MuJoCo默认步长 logger.info("="*60) logger.info("ROS 1 模式已启用!") - logger.info(f"发布话题:/mujoco/joint_states(关节状态)、/mujoco/pose(基座姿态)") - logger.info(f"订阅话题:/mujoco/ctrl_cmd(控制指令,长度={model.nu})") + logger.info(f"发布话题:/mujoco/joint_states、/mujoco/pose") + logger.info(f"订阅话题:/mujoco/ctrl_cmd(长度={model.nu})") logger.info("="*60) - # 1. 创建ROS发布者 - joint_state_pub = rospy.Publisher( - "/mujoco/joint_states", - JointState, - queue_size=10 # 消息队列大小 - ) - pose_pub = rospy.Publisher( - "/mujoco/pose", - PoseStamped, - queue_size=10 - ) + # 创建ROS发布者 + joint_state_pub = rospy.Publisher("/mujoco/joint_states", JointState, queue_size=10) + pose_pub = rospy.Publisher("/mujoco/pose", PoseStamped, queue_size=10) ros_publishers = (joint_state_pub, pose_pub) - # 2. 初始化关节状态消息(精准映射非自由关节的索引,支持多自由度) + # 初始化关节状态消息(精准映射非自由关节) joint_msg = JointState() joint_msg.name = [] for i in range(model.njnt): @@ -239,76 +291,79 @@ def visualize(model_path: str, use_ros: bool = False) -> None: if joint_type != mujoco.mjtJoint.mjJNT_FREE: joint_msg.name.append(model.joint(i).name) joint_ids.append(i) - # 获取该关节在qpos中的起始索引(mjJNT_FREE=7维, mjJNT_BALL=3维, mjJNT_HINGE/SLIDE=1维) joint_qpos_idxs.append(model.jnt_qposadr[i]) - # 获取该关节在qvel中的起始索引 joint_qvel_idxs.append(model.jnt_dofadr[i]) njnt = len(joint_msg.name) logger.info(f"ROS将发布 {njnt} 个非自由关节状态:{joint_msg.name}") - if njnt > 0: - logger.debug(f"关节qpos索引映射:{dict(zip(joint_msg.name, joint_qpos_idxs))}") - # 3. 创建ROS订阅者(接收控制指令) + # 创建ROS订阅者(接收控制指令) ctrl_cmd = np.zeros(model.nu) if model.nu > 0 else None def ctrl_callback(msg: Float32MultiArray): nonlocal ctrl_cmd if model.nu == len(msg.data): ctrl_cmd = np.array(msg.data) - logger.debug(f"收到ROS控制指令:{ctrl_cmd[:5]}...") # 只打印前5个值,避免日志冗余 + logger.debug(f"收到ROS控制指令:{ctrl_cmd[:5]}...") else: - logger.warning(f"控制指令长度不匹配!期望 {model.nu} 个值,实际收到 {len(msg.data)} 个") + logger.warning(f"控制指令长度不匹配!期望 {model.nu},实际 {len(msg.data)}") if model.nu > 0: ros_subscribers = rospy.Subscriber( - "/mujoco/ctrl_cmd", - Float32MultiArray, - ctrl_callback, - queue_size=5 + "/mujoco/ctrl_cmd", Float32MultiArray, ctrl_callback, queue_size=5 ) else: logger.warning("模型无控制输入(nu=0),不订阅控制指令话题") - # ===================== 可视化主循环(原有逻辑+优化ROS发布)===================== - logger.info("启动可视化窗口(按ESC键退出,鼠标可交互:拖拽旋转、滚轮缩放)") + # ===================== 可视化主循环 ===================== + logger.info("启动可视化窗口(按ESC键退出 | 鼠标交互:拖拽旋转、滚轮缩放)") try: with viewer.launch_passive(model, data) as v: while v.is_running() and (not use_ros or not rospy.is_shutdown()): - # ROS模式:应用控制指令(新增) + # 控制指令优先级:ROS指令 > 策略推理 > 无控制 if use_ros and ctrl_cmd is not None: data.ctrl[:] = ctrl_cmd + elif policy is not None: + # 提取观测:关节位置 + 关节速度 + obs = np.concatenate([data.qpos, data.qvel]) + obs_tensor = torch.tensor(obs, dtype=torch.float32).unsqueeze(0) + + # 策略推理 + action = policy(obs_tensor).squeeze().numpy() + + # 映射到实际控制范围 + if ctrl_range is not None: + action = ctrl_range[:, 0] + (ctrl_range[:, 1] - ctrl_range[:, 0]) * (action + 1) / 2 + + data.ctrl[:] = action - # 执行MuJoCo模拟步(原有逻辑) + # 执行模拟步 mujoco.mj_step(model, data) v.sync() - # ===================== ROS 消息发布(仅优化关节状态部分)===================== + # ROS消息发布 if use_ros and ros_publishers is not None: joint_state_pub, pose_pub = ros_publishers - # 1. 发布关节状态(位置、速度)- 优化后:精准映射+支持多自由度 + # 发布关节状态 joint_msg.header.stamp = rospy.Time.now() joint_msg.position = [] joint_msg.velocity = [] - for idx, (joint_id, qpos_idx, qvel_idx) in enumerate(zip(joint_ids, joint_qpos_idxs, joint_qvel_idxs)): + for joint_id, qpos_idx, qvel_idx in zip(joint_ids, joint_qpos_idxs, joint_qvel_idxs): joint_type = model.joint(joint_id).type - # 球关节(3自由度):补充3维位置/速度 if joint_type == mujoco.mjtJoint.mjJNT_BALL: joint_msg.position.extend(data.qpos[qpos_idx:qpos_idx+3]) joint_msg.velocity.extend(data.qvel[qvel_idx:qvel_idx+3]) - # 铰链/滑动关节(1自由度):仅取1维 elif joint_type in [mujoco.mjtJoint.mjJNT_HINGE, mujoco.mjtJoint.mjJNT_SLIDE]: joint_msg.position.append(data.qpos[qpos_idx]) joint_msg.velocity.append(data.qvel[qvel_idx]) joint_state_pub.publish(joint_msg) - # 2. 发布基座姿态(原有逻辑不变) + # 发布基座姿态 pose_msg = PoseStamped() pose_msg.header.stamp = rospy.Time.now() - pose_msg.header.frame_id = "world" # 坐标系名称(可自定义) + pose_msg.header.frame_id = "world" - # 位置信息(x,y,z) if model.nq >= 1: pose_msg.pose.position.x = data.qpos[0] if model.nq >= 2: @@ -316,7 +371,6 @@ def ctrl_callback(msg: Float32MultiArray): if model.nq >= 3: pose_msg.pose.position.z = data.qpos[2] - # 姿态信息(四元数 qx,qy,qz,qw) if model.nq >= 4: pose_msg.pose.orientation.x = data.qpos[3] if model.nq >= 5: @@ -327,23 +381,27 @@ def ctrl_callback(msg: Float32MultiArray): pose_msg.pose.orientation.w = data.qpos[6] pose_pub.publish(pose_msg) - - # 按ROS频率休眠,确保消息发布稳定 ros_rate.sleep() logger.info("可视化窗口已关闭") except Exception as e: logger.error(f"可视化过程出错: {str(e)}", exc_info=True) +# ===================== 主函数(命令行入口) ===================== # ===================== 主函数(仅优化model参数help+保持其他逻辑不变)===================== # ===================== 主函数(完全保持原有逻辑不变)===================== def main() -> None: parser = argparse.ArgumentParser( - description="MuJoCo功能整合工具(支持ROS 1消息封装)", + description="MuJoCo功能整合工具(支持强化学习策略控制 + ROS 1通信)", formatter_class=argparse.ArgumentDefaultsHelpFormatter ) subparsers = parser.add_subparsers(dest="command", required=True) + # 1. 可视化命令 + viz_parser = subparsers.add_parser("visualize", help="可视化模型并运行模拟") + viz_parser.add_argument("model", help="模型文件路径/目录(支持.xml/.mjb)") + viz_parser.add_argument("--ros", action="store_true", help="启用ROS 1模式") + viz_parser.add_argument("--policy", help="预训练策略模型路径(.pth文件)") # 1. 可视化命令(优化model参数help为通用提示,移除硬编码路径) viz_parser = subparsers.add_parser("visualize", help="可视化模型并运行模拟") viz_parser.add_argument("model", help="模型文件路径或包含模型的目录(支持.xml/.mjb格式)") @@ -354,24 +412,24 @@ def main() -> None: help="启用ROS模式(发布关节状态/基座姿态,订阅控制指令)" ) - # 2. 速度测试命令(原有功能不变) - speed_parser = subparsers.add_parser("testspeed", help="测试模型模拟速度") + # 2. 速度测试命令 + speed_parser = subparsers.add_parser("testspeed", help="测试模型模拟速度(多线程)") speed_parser.add_argument("model", help="模型文件路径") speed_parser.add_argument("--nstep", type=int, default=10000, help="每线程模拟步数") - speed_parser.add_argument("--nthread", type=int, default=1, help="测试线程数量") + speed_parser.add_argument("--nthread", type=int, default=1, help="测试线程数") speed_parser.add_argument("--ctrlnoise", type=float, default=0.01, help="控制噪声强度") - # 3. 模型转换命令(原有功能不变) - convert_parser = subparsers.add_parser("convert", help="转换模型格式(XML↔MJB)") + # 3. 模型转换命令 + convert_parser = subparsers.add_parser("convert", help="转换模型格式(XML ↔ MJB)") convert_parser.add_argument("input", help="输入模型路径") - convert_parser.add_argument("output", help="输出模型路径(需指定.xml或.mjb扩展名)") + convert_parser.add_argument("output", help="输出模型路径(指定.xml/.mjb)") + # 解析命令行参数 args, unknown = parser.parse_known_args() - - # 命令映射(更新visualize,支持use_ros参数) + # 命令映射 command_handlers: Dict[str, callable] = { - "visualize": lambda: visualize(args.model, use_ros=args.ros), + "visualize": lambda: visualize(args.model, use_ros=args.ros, policy_path=args.policy), "testspeed": lambda: test_speed(args.model, args.nstep, args.nthread, args.ctrlnoise), "convert": lambda: convert_model(args.input, args.output) } @@ -386,7 +444,6 @@ def main() -> None: logger.critical(f"程序执行失败: {str(e)}", exc_info=True) sys.exit(1) - - +# ===================== 程序入口 ===================== if __name__ == "__main__": main() \ No newline at end of file