332 lines
9.8 KiB
Markdown
332 lines
9.8 KiB
Markdown
# Go1 MoE 策略输入输出与部署说明
|
||
|
||
本文说明 `go1` 分支中 Go1 MoE locomotion 策略的输入、输出、状态管理和
|
||
RoboGauge/MuJoCo 接入方式。适用于由 `go2_rl_gym` 导出的 Go1 MoE CTS
|
||
TorchScript 模型,例如:
|
||
|
||
- `resources/models/go1/policy.pt`
|
||
- `resources/models/go1/policy_10k.pt`
|
||
- `resources/models/go1/policy_15k.pt`
|
||
- `resources/models/go1/policy_15k.onnx`
|
||
- `resources/models/go1/policy_25k.pt`
|
||
- `resources/models/go1/policy_25k.onnx`
|
||
|
||
## 1. 接口概览
|
||
|
||
| 项目 | 规格 |
|
||
| - | - |
|
||
| 控制频率 | 50 Hz,控制周期 `0.02 s` |
|
||
| TorchScript 输入 | `torch.float32`,形状 `[1, 45]` |
|
||
| 内部历史 | 5 帧,形状 `[1, 5, 45]` |
|
||
| 策略动作 | `torch.float32`,形状 `[1, 12]` |
|
||
| MoE 权重 | `torch.float32`,形状 `[1, 8]` |
|
||
| Student latent | `torch.float32`,形状 `[1, 32]` |
|
||
| 控制方式 | 关节位置 PD 控制 |
|
||
| 关节位置增量比例 | `action_scale = 0.25` |
|
||
| PD 参数 | `Kp = 28.0`,`Kd = 0.7` |
|
||
|
||
导出的 TorchScript 是有状态模型。调用方每个控制周期只输入当前 45 维
|
||
观测,模型会在内部更新 5 帧历史。不要在 TorchScript 输入端再次拼接成
|
||
225 维。
|
||
|
||
## 2. 坐标系和关节顺序
|
||
|
||
机身坐标系约定:
|
||
|
||
- `+x`:前方
|
||
- `+y`:左方
|
||
- `+z`:上方
|
||
- 正 yaw:绕 `+z` 轴逆时针旋转
|
||
- 四元数顺序:`[w, x, y, z]`
|
||
|
||
策略的 12 个关节严格按照以下顺序排列:
|
||
|
||
```text
|
||
0 FR_hip 1 FR_thigh 2 FR_calf
|
||
3 FL_hip 4 FL_thigh 5 FL_calf
|
||
6 RR_hip 7 RR_thigh 8 RR_calf
|
||
9 RL_hip 10 RL_thigh 11 RL_calf
|
||
```
|
||
|
||
即腿序为 `FR, FL, RR, RL`,每条腿内部为 `hip, thigh, calf`。
|
||
|
||
默认站立关节位置为:
|
||
|
||
```python
|
||
DEFAULT_DOF_POS = np.array([
|
||
-0.1, 0.8, -1.5, # FR
|
||
0.1, 0.8, -1.5, # FL
|
||
-0.1, 1.0, -1.5, # RR
|
||
0.1, 1.0, -1.5, # RL
|
||
], dtype=np.float32)
|
||
```
|
||
|
||
Isaac Gym 中 Go1 资产的原始关节顺序不同,但训练环境已经转换成上述策略
|
||
顺序。MuJoCo 和真机部署端必须直接使用上述顺序,不能再次交换左右腿。
|
||
|
||
## 3. 45 维单帧观测
|
||
|
||
每一帧观测按下表顺序拼接:
|
||
|
||
| 索引 | 维度 | 原始信号 | 输入值 |
|
||
| - | -: | - | - |
|
||
| `[0:3]` | 3 | 机身坐标系角速度 `[wx, wy, wz]` | `base_ang_vel * 0.25` |
|
||
| `[3:6]` | 3 | 世界重力方向在机身坐标系中的投影 | `projected_gravity` |
|
||
| `[6:9]` | 3 | 速度指令 `[vx, vy, yaw_rate]` | `command * [2.0, 2.0, 0.25]` |
|
||
| `[9:21]` | 12 | 关节位置 | `(q - q_default) * 1.0` |
|
||
| `[21:33]` | 12 | 关节速度 | `dq * 0.05` |
|
||
| `[33:45]` | 12 | 上一个控制周期的原始策略动作 | `last_action` |
|
||
|
||
所有关节相关数组都必须先转换成 `FR, FL, RR, RL` 策略顺序。
|
||
|
||
### 3.1 角速度
|
||
|
||
角速度必须是机身坐标系下的陀螺仪读数,单位为 `rad/s`。输入模型前乘以
|
||
`0.25`。不要传入世界坐标系角速度。
|
||
|
||
### 3.2 投影重力
|
||
|
||
策略不直接输入 roll、pitch 或四元数,而是输入单位重力向量在机身坐标系
|
||
中的投影。对于 `[w, x, y, z]` 四元数,可使用 RoboGauge 的实现:
|
||
|
||
```python
|
||
from robogauge.utils.math_utils import get_projected_gravity
|
||
|
||
projected_gravity = get_projected_gravity(base_quat_wxyz)
|
||
```
|
||
|
||
机器人水平站立且朝向不影响姿态时,该向量约为 `[0, 0, -1]`。
|
||
|
||
### 3.3 速度指令
|
||
|
||
输入的是期望机身速度,不是目标点坐标:
|
||
|
||
```text
|
||
[forward_velocity, lateral_velocity, yaw_rate]
|
||
```
|
||
|
||
RoboGauge 的目标点控制器会先把目标点误差转换为速度指令,再构建观测。
|
||
|
||
15k 楼梯专项模型的训练范围为:
|
||
|
||
```text
|
||
vx = [-1.0, 1.0] m/s
|
||
vy = [-0.5, 0.5] m/s
|
||
yaw_rate = [-1.0, 1.0] rad/s
|
||
```
|
||
|
||
当前 `Go1TerrainConfig` 允许更宽的 `vy` 和 `yaw_rate`,用于通用 RoboGauge
|
||
测试。评估 15k 楼梯模型时,超出上述范围的结果应视为训练分布外表现。
|
||
|
||
### 3.4 上一动作
|
||
|
||
`last_action` 是上一个控制周期模型输出的 12 维原始 action,不是 PD 目标
|
||
位置,也不是实际关节位置。episode 开始时应初始化为全零。
|
||
|
||
## 4. 历史观测
|
||
|
||
Student MoE encoder 使用 5 个连续控制帧:
|
||
|
||
```text
|
||
[obs(t-4), obs(t-3), obs(t-2), obs(t-1), obs(t)]
|
||
```
|
||
|
||
控制周期为 20 ms,因此历史缓冲包含 5 帧、跨度 80 ms。flatten 后 encoder
|
||
实际处理 225 维数据,但该缓冲已经封装在 TorchScript 模型中。
|
||
|
||
调用 `model.reset()` 后,历史缓冲为全零。随后每次调用模型都会丢弃最旧
|
||
一帧并追加当前观测。这与训练和 RoboGauge 中的行为一致。
|
||
|
||
## 5. TorchScript 输出
|
||
|
||
MoE TorchScript 返回:
|
||
|
||
```python
|
||
action, (weights, latent) = model(obs_tensor)
|
||
```
|
||
|
||
各输出含义:
|
||
|
||
| 输出 | 形状 | 含义 |
|
||
| - | - | - |
|
||
| `action` | `[1, 12]` | 确定性关节位置增量动作 |
|
||
| `weights` | `[1, 8]` | 8 个 Student MoE expert 的门控权重 |
|
||
| `latent` | `[1, 32]` | Student encoder 估计的环境隐变量 |
|
||
|
||
实际控制只需要 `action`。`weights` 和 `latent` 用于分析与可视化,可以忽略。
|
||
|
||
将原始 action 转换为关节位置目标:
|
||
|
||
```python
|
||
target_dof_pos = DEFAULT_DOF_POS + 0.25 * action
|
||
```
|
||
|
||
随后使用位置 PD 控制:
|
||
|
||
```text
|
||
tau = Kp * (target_dof_pos - q) - Kd * dq
|
||
Kp = 28.0
|
||
Kd = 0.7
|
||
```
|
||
|
||
实际部署还应按电机能力限制输出力矩。
|
||
|
||
## 6. ONNX 输入输出
|
||
|
||
`policy_15k.onnx` 是无状态模型,不在模型内部保存历史。它与 TorchScript
|
||
的输入接口不同:
|
||
|
||
| 项目 | ONNX 规格 |
|
||
| - | - |
|
||
| 输入名称 | `obs` |
|
||
| 输入形状 | `[1, 225]` |
|
||
| 输出 0 | `actions`,形状 `[1, 12]` |
|
||
| 输出 1 | `weights`,形状 `[1, 8]` |
|
||
| 输出 2 | `latent`,形状 `[1, 32]` |
|
||
|
||
ONNX 输入不是直接对 `[1, 5, 45]` 执行 `reshape`。225 维输入按观测项分组,
|
||
每组内部再按时间从旧到新排列:
|
||
|
||
```text
|
||
[
|
||
angular_velocity(t-4:t), # 5 * 3 = 15
|
||
projected_gravity(t-4:t), # 5 * 3 = 15
|
||
command(t-4:t), # 5 * 3 = 15
|
||
dof_position_error(t-4:t), # 5 * 12 = 60
|
||
dof_velocity(t-4:t), # 5 * 12 = 60
|
||
last_action(t-4:t), # 5 * 12 = 60
|
||
]
|
||
```
|
||
|
||
调用方需要维护 5 帧历史,并在 episode reset 时将历史清零。RoboGauge 当前
|
||
直接使用 TorchScript;ONNX 主要用于其他推理后端。
|
||
|
||
## 7. 最小调用示例
|
||
|
||
```python
|
||
import numpy as np
|
||
import torch
|
||
|
||
from robogauge.utils.math_utils import get_projected_gravity
|
||
|
||
|
||
JOINT_ORDER = [
|
||
"FR_hip", "FR_thigh", "FR_calf",
|
||
"FL_hip", "FL_thigh", "FL_calf",
|
||
"RR_hip", "RR_thigh", "RR_calf",
|
||
"RL_hip", "RL_thigh", "RL_calf",
|
||
]
|
||
|
||
DEFAULT_DOF_POS = np.array([
|
||
-0.1, 0.8, -1.5,
|
||
0.1, 0.8, -1.5,
|
||
-0.1, 1.0, -1.5,
|
||
0.1, 1.0, -1.5,
|
||
], dtype=np.float32)
|
||
|
||
model = torch.jit.load("resources/models/go1/policy_10k.pt", map_location="cpu")
|
||
model.eval()
|
||
|
||
# 每个 episode 开始时执行。
|
||
model.reset()
|
||
last_action = np.zeros(12, dtype=np.float32)
|
||
|
||
|
||
def build_observation(base_ang_vel, base_quat_wxyz, command, q, dq, last_action):
|
||
"""All joint arrays must use the FR, FL, RR, RL policy order."""
|
||
obs = np.concatenate([
|
||
np.asarray(base_ang_vel, dtype=np.float32) * 0.25,
|
||
get_projected_gravity(np.asarray(base_quat_wxyz, dtype=np.float32)),
|
||
np.asarray(command, dtype=np.float32) * np.array([2.0, 2.0, 0.25], dtype=np.float32),
|
||
(np.asarray(q, dtype=np.float32) - DEFAULT_DOF_POS),
|
||
np.asarray(dq, dtype=np.float32) * 0.05,
|
||
np.asarray(last_action, dtype=np.float32),
|
||
]).astype(np.float32)
|
||
assert obs.shape == (45,)
|
||
return np.clip(obs, -100.0, 100.0)
|
||
|
||
|
||
# 以下变量由 IMU、关节编码器和上层速度控制器提供。
|
||
obs = build_observation(
|
||
base_ang_vel=base_ang_vel_body,
|
||
base_quat_wxyz=base_quat_wxyz,
|
||
command=velocity_command,
|
||
q=joint_position,
|
||
dq=joint_velocity,
|
||
last_action=last_action,
|
||
)
|
||
|
||
with torch.inference_mode():
|
||
obs_tensor = torch.from_numpy(obs).unsqueeze(0)
|
||
action_tensor, (weights, latent) = model(obs_tensor)
|
||
|
||
action = action_tensor.squeeze(0).cpu().numpy()
|
||
target_dof_pos = DEFAULT_DOF_POS + 0.25 * action
|
||
last_action = action.copy()
|
||
```
|
||
|
||
## 8. Reset 要求
|
||
|
||
以下情况必须同时重置模型历史和上一动作:
|
||
|
||
- 仿真 episode reset
|
||
- 机器人摔倒后重新站立
|
||
- 策略重新加载
|
||
- 控制器长时间暂停后重新启动
|
||
- 真机急停解除后重新接管
|
||
|
||
```python
|
||
model.reset()
|
||
last_action.fill(0.0)
|
||
```
|
||
|
||
只清零 `last_action` 而不调用 `model.reset()`,会残留上一段运行的历史状态。
|
||
|
||
## 9. 不需要的部署信号
|
||
|
||
部署使用 Student 策略,不需要训练阶段 Teacher/Critic 的 263 维 privileged
|
||
observation。以下信号不进入部署策略:
|
||
|
||
- 机身线速度
|
||
- 足端接触力
|
||
- 电机力矩
|
||
- 关节加速度
|
||
- 187 维地形高度扫描
|
||
- 深度图、相机图像或 LiDAR
|
||
|
||
这些信息仅在训练阶段用于 Teacher/Critic。Student 通过 5 帧本体观测历史估计
|
||
32 维 latent。
|
||
|
||
## 10. RoboGauge 对应实现
|
||
|
||
RoboGauge 已实现上述接口:
|
||
|
||
- 观测构建:`robogauge/tasks/robots/go2/go2.py`
|
||
- Go1 参数和关节顺序:`robogauge/tasks/robots/go1/go1_config.py`
|
||
- MoE 输出处理:`robogauge/tasks/robots/go2/go2_moe.py`
|
||
- 投影重力:`robogauge/utils/math_utils.py`
|
||
|
||
在 RoboGauge 中加载 Go1 MoE 模型时,通常无需手动构建观测:
|
||
|
||
```bash
|
||
MUJOCO_GL=glfw \
|
||
PYTHONPATH=/path/to/RoboGauge \
|
||
python robogauge/scripts/run.py \
|
||
--task-name go1_moe.stairs_fd \
|
||
--experiment-name go1_policy_test \
|
||
--model-path /path/to/policy.pt \
|
||
--level 6 \
|
||
--goals target_pos_velocity
|
||
```
|
||
|
||
## 11. 常见错误检查
|
||
|
||
1. 输入形状应为 `[1, 45]`,不是 `[45]` 或 `[1, 225]`。
|
||
2. 输入 dtype 必须为 `torch.float32`。
|
||
3. 角速度必须在机身坐标系下。
|
||
4. 四元数必须是 `[w, x, y, z]`。
|
||
5. 关节顺序必须是 `FR, FL, RR, RL`。
|
||
6. `last_action` 必须使用未乘 `0.25` 的原始模型输出。
|
||
7. 每次 reset 必须同时调用 `model.reset()` 并清零 `last_action`。
|
||
8. 不要向部署模型输入 privileged observation 或高度扫描。
|