#!/usr/bin/env python3 """无头 MuJoCo 测试:前进 → 突然停止 → 观察姿态变化。""" import numpy as np import mujoco import onnxruntime as ort import os, sys, time _PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) XML_DIR = os.path.join(_PROJECT, "motrix_envs", "src", "motrix_envs", "locomotion", "go1", "xmls") ONNX = os.path.join(_PROJECT, "runs/go1-dreamwaq-walk/rslrl/26-07-02_20-03-23-_36054_PPO/policy.onnx") NUM_OBS = 45 NUM_ACTIONS = 12 HISTORY_LEN = 5 ACTION_SCALE = 0.25 KP = 28.0 KD = 0.7 DECIMATION = 4 DEFAULT_ANGLES = np.array([0.0, 0.9, -1.8, 0.0, 0.9, -1.8, 0.0, 0.9, -1.8, 0.0, 0.9, -1.8], dtype=np.float32) def compute_obs(model, data, commands, last_action): obs = np.zeros(NUM_OBS, dtype=np.float32) g = None # no gyro sensor in headless obs[0:3] = (g if g is not None else data.qvel[3:6]) * 0.25 grav_world = model.opt.gravity.copy() grav_world = grav_world / np.linalg.norm(grav_world) R = data.xmat[1].reshape(3, 3) obs[3:6] = (R.T @ grav_world).astype(np.float32) obs[6:9] = commands * np.array([2.0, 2.0, 0.25], dtype=np.float32) obs[9:21] = (data.qpos[7:19] - DEFAULT_ANGLES) * 1.0 obs[21:33] = data.qvel[6:18] * 0.05 obs[33:45] = last_action return obs def get_base_pose(data): """返回 base 的 z 高度和 roll/pitch(度)。""" quat = data.xquat[1] # base body quaternion w, x, y, z = quat[0], quat[1], quat[2], quat[3] # roll, pitch from quaternion sinr_cosp = 2 * (w * x + y * z) cosr_cosp = 1 - 2 * (x * x + y * y) roll = np.arctan2(sinr_cosp, cosr_cosp) sinp = 2 * (w * y - z * x) pitch = np.arcsin(np.clip(sinp, -1, 1)) return float(data.xpos[1, 2]), np.degrees(roll), np.degrees(pitch) def main(): session = ort.InferenceSession(ONNX, providers=['CPUExecutionProvider']) xml_file = os.path.join(XML_DIR, "scene_dreamwaq_flat.xml") model = mujoco.MjModel.from_xml_path(xml_file) data = mujoco.MjData(model) # 初始化 data.qpos[7:19] = DEFAULT_ANGLES data.qpos[2] = 0.35 mujoco.mj_forward(model, data) history = np.zeros((1, HISTORY_LEN, NUM_OBS), dtype=np.float32) last_action = np.zeros(NUM_ACTIONS, dtype=np.float32) print(f"{'Step':>5s} {'Cmd_vx':>7s} {'base_z':>8s} {'roll':>8s} {'pitch':>8s} {'speed_xy':>8s}") print("-" * 60) for step in range(3000): # 命令:前 1500 步前进,之后突然停止 if step < 1500: vx, vy, wz = 0.5, 0.0, 0.0 # 前进 else: vx, vy, wz = 0.0, 0.0, 0.0 # 突然停止! for _ in range(DECIMATION): mujoco.mj_step(model, data) if step % DECIMATION == 0: cmd = np.array([vx, vy, wz], dtype=np.float32) obs = compute_obs(model, data, cmd, last_action) history = np.concatenate([history[:, 1:, :], obs.reshape(1, 1, -1)], axis=1) outputs = session.run(None, { 'obs': obs.reshape(1, -1).astype(np.float32), 'obs_history': history.reshape(1, -1).astype(np.float32), }) action = outputs[0][0] last_action = action target = DEFAULT_ANGLES + action * ACTION_SCALE data.ctrl[:] = KP * (target - data.qpos[7:19]) - KD * data.qvel[6:18] # 每 100 步打印 if step % 100 == 0: bz, roll, pitch = get_base_pose(data) speed_xy = np.linalg.norm(data.qvel[0:2]) print(f"{step:5d} {vx:7.1f} {bz:8.3f} {roll:+8.1f} {pitch:+8.1f} {speed_xy:8.3f}") # 最终状态 bz, roll, pitch = get_base_pose(data) print(f"\n最终: base_z={bz:.3f} roll={roll:.1f}° pitch={pitch:.1f}°") if abs(roll) > 30 or abs(pitch) > 30: print("⚠ 机器人倾覆!突然停止导致翻跟头") else: print("✅ 机器人在突然停止后保持稳定") if __name__ == "__main__": main()