This commit is contained in:
2026-04-05 07:37:08 +08:00
parent f3300ff1bf
commit d988ed941b
3 changed files with 1319 additions and 0 deletions

470
go1_unitree_rl_inference.py Normal file
View File

@@ -0,0 +1,470 @@
#!/usr/bin/env python3
"""
Go1 RL Policy 推理 - 移植自 unitree_rl 项目
基于 https://github.com/dstx123/unitree_rl
模型: body.jit + adapt.jit (GRU-based policy with history)
训练环境: IsaacGym
使用方法:
python3 /home/8x54zj-m/unitree_mujoco/go1_unitree_rl_inference.py
键盘控制:
W/S: 前进/后退
A/D: 左转/右转
Q/E: 侧向左移/右移
空格: 停止
ESC: 退出
"""
import numpy as np
import mujoco
from mujoco import viewer
import os
import sys
import tty
import termios
import signal
import threading
import time
import torch
# ============================================================
# 全局退出标志
# ============================================================
g_exit_requested = False
def signal_handler(signum, frame):
global g_exit_requested
g_exit_requested = True
signal.signal(signal.SIGINT, signal_handler)
# ============================================================
# 配置
# ============================================================
BODY_MODEL_PATH = "/tmp/unitree_rl/src/unitree_guide/unitree_guide/model/body.jit"
ADAPT_MODEL_PATH = "/tmp/unitree_rl/src/unitree_guide/unitree_guide/model/adapt.jit"
XML_PATH = "/home/8x54zj-m/unitree_mujoco/data/go1/xml/go1.xml"
MESH_PATH = "/home/8x54zj-m/unitree_mujoco/data/go1/meshes"
# ============================================================
# 模型参数 (来自 unitree_rl/State_RL.h)
# ============================================================
NUM_OBS = 45 # 观测维度
NUM_OBS_HISTORY = 10 # 历史观测步数
OBS_BUFFER_SIZE = NUM_OBS * NUM_OBS_HISTORY # 450
# 观测缩放
SCALE_LIN_VEL = 2.0
SCALE_ANG_VEL = 0.25
SCALE_COMMANDS = np.array([SCALE_LIN_VEL, SCALE_LIN_VEL, SCALE_ANG_VEL])
SCALE_DOF_POS = 1.0
SCALE_DOF_VEL = 0.05
# 动作缩放
ACTION_SCALE = 0.25
HIP_SCALE_REDUCTION = 0.5 # hip关节额外缩放
# 限幅
CLIP_OBSERVATIONS = 100.0
CLIP_ACTIONS = 100.0
# 默认关节角度 (与 unitree_rl 一致)
DEFAULT_DOF_POS = np.array([
-0.1, 0.8, -1.5, # FR: hip, thigh, calf
0.1, 0.8, -1.5, # FL: hip, thigh, calf
-0.1, 1.0, -1.5, # RR: hip, thigh, calf
0.1, 1.0, -1.5 # RL: hip, thigh, calf
])
# Motor order in code: FR(3,4,5), FL(0,1,2), RR(9,10,11), RL(6,7,8)
# Model order: FL(0,1,2), FR(3,4,5), RL(6,7,8), RR(9,10,11)
# dof_mapping = {3, 4, 5, 0, 1, 2, 9, 10, 11, 6, 7, 8}
# Inverse mapping (model idx -> motor idx): 3->0, 4->1, 5->2, 0->3, 1->4, 2->5, 9->6, 10->7, 11->8, 6->9, 7->10, 8->11
# This maps: model_output[0..11] -> motor_q[0..11]
DOF_MAPPING = np.array([3, 4, 5, 0, 1, 2, 9, 10, 11, 6, 7, 8])
# 速度命令范围
MAX_LIN_VEL = 1.0
MAX_ANG_VEL = 1.0
# PD控制参数 (来自 unitree_rl)
KP = 7.0
KD = 1.0
# MuJoCo 摩擦系数配置
# friction = [滑动摩擦(X), 滑动摩擦(Y), 扭转摩擦]
# - friction[0]: X方向滑动摩擦系数
# - friction[1]: Y方向滑动摩擦系数 (各向异性材料用)
# - friction[2]: 扭转摩擦系数 (绕法线的旋转阻力)
#
# 常见地面摩擦参数参考:
# ┌──────────┬───────────────────┐
# │ 地面类型 │ FLOOR_FRICTION │
# ├──────────┼───────────────────┤
# │ 冰面 │ [0.05, 0.02, 0.01]│
# │ 木地板 │ [0.4, 0.3, 0.2 ]│
# │ 瓷砖 │ [0.6, 0.4, 0.3 ]│
# │ 橡胶垫 │ [1.5, 1.0, 0.5 ]│
# │ 粗糙地面 │ [2.0, 1.5, 1.0 ]│
# └──────────┴───────────────────┘
# 机器人 geom 摩擦系数
FRICTION = [0.6, 0.3, 0.3]
# 地面摩擦系数 (单独设置)
FLOOR_FRICTION = [2.0, 1.5, 1.0]
# ============================================================
# 键盘输入读取 (termios 非阻塞方式,无 root 依赖)
# ============================================================
import select
class KeyboardReader:
def __init__(self):
self.keys_pressed = set()
self.last_key_time = 0
self.timeout = 0.1 # 无按键超时时间(秒),超时后清除所有按键
def init(self):
import termios
self.old_settings = termios.tcgetattr(sys.stdin)
# 非canonical模式关闭回显
new_settings = termios.tcgetattr(sys.stdin)
new_settings[3] = new_settings[3] & ~termios.ICANON & ~termios.ECHO
new_settings[6][termios.VMIN] = 0 # 非阻塞
new_settings[6][termios.VTIME] = 0
termios.tcsetattr(sys.stdin, termios.TCSADRAIN, new_settings)
print("[INFO] Keyboard reader initialized (termios mode, no root required)")
def _read_key(self):
"""尝试读取一个按键返回按键字符或None"""
try:
if select.select([sys.stdin], [], [], 0)[0]:
ch = sys.stdin.read(1)
return ch
except:
pass
return None
def update(self):
"""在主循环中调用,更新按键状态"""
import time
current_time = time.time()
# 尝试读取按键
ch = self._read_key()
if ch:
if ch == '\x1b': # ESC
self.keys_pressed.add('escape')
elif ch == ' ':
self.keys_pressed.add('space')
elif ch == '\n' or ch == '\r':
pass
else:
self.keys_pressed.add(ch.lower())
self.last_key_time = current_time
elif current_time - self.last_key_time > self.timeout:
# 超时后清除所有按键(模拟按键释放)
self.keys_pressed.clear()
def restore(self):
import termios
if hasattr(self, 'old_settings'):
termios.tcsetattr(sys.stdin, termios.TCSADRAIN, self.old_settings)
def is_key_pressed(self, key):
if key == '\x1b' or key == 'escape':
return 'escape' in self.keys_pressed
elif key == ' ':
return 'space' in self.keys_pressed
return key.lower() in self.keys_pressed
def clear(self):
self.keys_pressed.clear()
# ============================================================
# 辅助函数
# ============================================================
def quaternion_to_rotation_matrix(q):
"""四元数转旋转矩阵 (MuJoCo 格式: qx, qy, qz, qw)"""
qx, qy, qz, qw = q
norm = np.sqrt(qx**2 + qy**2 + qz**2 + qw**2)
qx, qy, qz, qw = qx/norm, qy/norm, qz/norm, qw/norm
return np.array([
[1-2*(qy**2+qz**2), 2*(qx*qy-qz*qw), 2*(qx*qz+qy*qw)],
[2*(qx*qy+qz*qw), 1-2*(qx**2+qz**2), 2*(qy*qz-qx*qw)],
[2*(qx*qz-qy*qw), 2*(qy*qz+qx*qw), 1-2*(qx**2+qy**2)]
])
def quat_rotate_inverse(q, v):
"""四元数逆旋转 (世界坐标系 -> 躯干坐标系)
对应 IsaacLab 的 quat_apply_inverse
q: (w, x, y, z) 格式四元数
v: (x, y, z) 向量
"""
q_w, q_x, q_y, q_z = q[3], q[0], q[1], q[2]
return np.array([
v[0] * (2*q_w**2 - 1) + 2*q_x*(q_y*v[2] - q_z*v[1]) + 2*q_w*(q_z*v[0] - q_x*v[2]),
v[1] * (2*q_w**2 - 1) + 2*q_y*(q_z*v[0] - q_x*v[2]) + 2*q_w*(q_x*v[1] - q_y*v[0]),
v[2] * (2*q_w**2 - 1) + 2*q_z*(q_x*v[1] - q_y*v[0]) + 2*q_w*(q_y*v[0] - q_x*v[1])
])
def compute_observations(data, prev_action, commands, default_dof_pos):
"""
计算 RL policy 观测向量 (45维)
对应 unitree_rl 的观测计算:
[0:3] body_ang_vel - 躯干角速度 (body frame)
[3:6] projected_gravity - 重力投影到躯干坐标系
[6:9] commands - 速度命令 (已缩放)
[9:21] dof_pos_rel - 相对关节位置
[21:33] dof_vel - 关节速度
[33:45] prev_action - 上一步动作
"""
obs = np.zeros(NUM_OBS, dtype=np.float32)
# 获取四元数 (MuJoCo 格式: x, y, z, w)
quat_mujoco = data.qpos[3:7] # (x, y, z, w)
# 转换为 unitree_rl 格式 (x, y, z, w) for quat_rotate_inverse
# 1. 躯干角速度 (body frame)
# unitree_rl 直接使用 imu.gyroscope (已是在 body frame)
# MuJoCo 的 qvel[3:6] 是世界坐标系,需要转换
ang_vel_world = data.qvel[3:6]
ang_vel_body = quat_rotate_inverse(np.concatenate([quat_mujoco[3:4], quat_mujoco[0:3]]), ang_vel_world)
obs[0:3] = ang_vel_body * SCALE_ANG_VEL
# 2. 重力投影 (body frame)
# gravity_vec = (0, 0, -1) in unitree_rl
gravity_world = np.array([0.0, 0.0, -1.0])
projected_gravity = quat_rotate_inverse(np.concatenate([quat_mujoco[3:4], quat_mujoco[0:3]]), gravity_world)
obs[3:6] = projected_gravity
# 3. 速度命令 (已缩放)
obs[6:9] = commands * SCALE_COMMANDS
# 4. 相对关节位置
# Motor order: FR(3,4,5), FL(0,1,2), RR(9,10,11), RL(6,7,8)
# Model order: FL(0,1,2), FR(3,4,5), RL(6,7,8), RR(9,10,11)
motor_pos = data.qpos[7:19] # MuJoCo motor order
# Map to model order
model_order_pos = motor_pos[DOF_MAPPING] # (12,)
dof_pos_rel = model_order_pos - default_dof_pos
obs[9:21] = dof_pos_rel * SCALE_DOF_POS
# 5. 关节速度
motor_vel = data.qvel[6:18]
model_order_vel = motor_vel[DOF_MAPPING]
obs[21:33] = model_order_vel * SCALE_DOF_VEL
# 6. 上一步动作
obs[33:45] = prev_action
# 限幅
obs = np.clip(obs, -CLIP_OBSERVATIONS, CLIP_OBSERVATIONS)
return obs
# ============================================================
# 主程序
# ============================================================
def main():
# 1. 加载 MuJoCo 模型
os.chdir(MESH_PATH)
with open(XML_PATH, 'r') as f:
xml_content = f.read()
xml_content = xml_content.replace('meshdir="../meshes/"', f'meshdir="{MESH_PATH}"')
xml_content = xml_content.replace('<default>\n=', '<default>\n')
model = mujoco.MjModel.from_xml_string(xml_content)
data = mujoco.MjData(model)
print(f"[INFO] MuJoCo Model: {model.nbody} bodies, {model.nq} DoF, {model.nu} actuators")
print(f"[INFO] 关节顺序: {[mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_JOINT, i) for i in range(1, 13)]}")
# 2. 设置摩擦系数
# friction = [滑动摩擦, 扭转摩擦, 滚动摩擦]
# 设置地面摩擦系数 (通过 geom 名称查找)
floor_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_GEOM, "floor")
if floor_id >= 0:
model.geom_friction[floor_id] = FLOOR_FRICTION
print(f"[INFO] 地面摩擦系数设置为: {FLOOR_FRICTION}")
else:
print(f"[WARNING] 未找到 floor geom使用默认摩擦系数")
# 设置机器人 geom 的摩擦系数 (排除地面)
for i in range(model.ngeom):
if i != floor_id:
model.geom_friction[i] = FRICTION
print(f"[INFO] 机器人摩擦系数设置为: {FRICTION}")
# 3. 加载 RL 模型
body_module = torch.jit.load(BODY_MODEL_PATH, map_location='cpu')
adapt_module = torch.jit.load(ADAPT_MODEL_PATH, map_location='cpu')
body_module.eval()
adapt_module.eval()
print(f"[INFO] Body model loaded from: {BODY_MODEL_PATH}")
print(f"[INFO] Adapt model loaded from: {ADAPT_MODEL_PATH}")
print(f"[INFO] 观测维度: {NUM_OBS}, 历史步数: {NUM_OBS_HISTORY}")
# 3. 初始化
obs = np.zeros(NUM_OBS, dtype=np.float32)
prev_action = np.zeros(12, dtype=np.float32)
obs_buffer = np.zeros(OBS_BUFFER_SIZE, dtype=np.float32) # 10 * 45 = 450
# 初始化机器人姿态
data.qpos[7:19] = DEFAULT_DOF_POS
data.qvel[:] = 0.0
data.qpos[2] = 0.35 # 抬高躯干
mujoco.mj_step(model, data)
print("[INFO] 机器人初始化完成,开始 RL 策略推理...")
# 4. 初始化键盘控制
keyboard_reader = KeyboardReader()
keyboard_reader.init()
print("[INFO] 键盘控制已启用!")
print(" W/S: 前进/后退")
print(" A/D: 左转/右转")
print(" Q/E: 侧向左移/右移")
print(" 空格: 停止")
print(" ESC: 退出")
# 5. 启动交互式查看器
import mujoco.viewer as mv
view = mv.launch_passive(model, data)
print("[INFO] 交互式查看器已启动!")
# 初始化观测 buffer (填充历史)
for _ in range(NUM_OBS_HISTORY):
obs = compute_observations(data, prev_action, np.array([0.0, 0.0, 0.0]), DEFAULT_DOF_POS)
obs_buffer = np.concatenate([obs_buffer[NUM_OBS:], obs])
step_count = 0
last_inference_time = time.time()
inference_interval = 0.02 # 50Hz
command = np.array([0.0, 0.0, 0.0], dtype=np.float32)
joint_targets = DEFAULT_DOF_POS.copy() # 初始化关节目标
while view.is_running() and not g_exit_requested:
current_time = time.time()
# 更新键盘状态
keyboard_reader.update()
# 读取键盘输入 (切换模式: 按下切换状态)
# 空格: 停止所有
if keyboard_reader.is_key_pressed(' '):
command[:] = 0.0
keyboard_reader.clear()
# W/S: 前进/后退 (vx)
if keyboard_reader.is_key_pressed('w'):
command[0] = MAX_LIN_VEL
command[1] = 0.0 # 清除侧向
command[2] = 0.0 # 清除转向
keyboard_reader.clear()
if keyboard_reader.is_key_pressed('s'):
command[0] = -MAX_LIN_VEL
command[1] = 0.0
command[2] = 0.0
keyboard_reader.clear()
# Q/E: 侧向移动 (vy)
if keyboard_reader.is_key_pressed('q'):
command[1] = MAX_LIN_VEL
command[0] = 0.0
command[2] = 0.0
keyboard_reader.clear()
if keyboard_reader.is_key_pressed('e'):
command[1] = -MAX_LIN_VEL
command[0] = 0.0
command[2] = 0.0
keyboard_reader.clear()
# A/D: 左转/右转 (yaw_rate)
if keyboard_reader.is_key_pressed('a'):
command[2] = MAX_ANG_VEL
command[0] = 0.0
command[1] = 0.0
keyboard_reader.clear()
if keyboard_reader.is_key_pressed('d'):
command[2] = -MAX_LIN_VEL
command[0] = 0.0
command[1] = 0.0
keyboard_reader.clear()
if keyboard_reader.is_key_pressed('\x1b'):
print("[INFO] ESC pressed, exiting...")
break
# 50Hz 推理
if current_time - last_inference_time >= inference_interval:
# 计算观测
obs = compute_observations(data, prev_action, command, DEFAULT_DOF_POS)
# 更新观测历史 buffer
obs_buffer = np.concatenate([obs_buffer[NUM_OBS:], obs])
# 模型推理
with torch.inference_mode():
obs_buffer_tensor = torch.from_numpy(obs_buffer).float().unsqueeze(0) # (1, 450)
obs_tensor = torch.from_numpy(obs).float().unsqueeze(0) # (1, 45)
# Adapt 模块: 编码历史观测为 latent
latent = adapt_module(obs_buffer_tensor) # (1, 21)
# Body 模块: 结合当前观测和 latent 生成动作
combined_input = torch.cat([obs_tensor, latent], dim=1) # (1, 66)
action = body_module(combined_input).numpy().flatten() # (12,)
# 动作后处理
action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
actions_scaled = action * ACTION_SCALE
# Hip关节额外缩放
hip_indices = [0, 3, 6, 9]
for i in hip_indices:
actions_scaled[i] *= HIP_SCALE_REDUCTION
# 转换为电机角度目标
joint_targets = actions_scaled[DOF_MAPPING] + DEFAULT_DOF_POS
# 保存上一步动作
prev_action = action.copy()
last_inference_time = current_time
# PD 控制 (每步执行)
current_pos = data.qpos[7:19]
current_vel = data.qvel[6:18]
torque = KP * (joint_targets - current_pos) - KD * current_vel
# 应用控制
data.ctrl[:] = torque
mujoco.mj_step(model, data)
view.sync()
step_count += 1
if step_count % 100 == 0:
trunk_z = data.qpos[2]
lin_vel = np.linalg.norm(data.qvel[0:3])
print(f" Step {step_count}: cmd=[{command[0]:.2f}, {command[1]:.2f}, {command[2]:.2f}] | "
f"trunk_z={trunk_z:.3f}m, vel={lin_vel:.3f}m/s")
keyboard_reader.restore()
view.close()
if __name__ == "__main__":
main()