108 lines
3.8 KiB
Python
108 lines
3.8 KiB
Python
#!/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()
|