OpenDuckMini强化学习框架入门教程-ONNX集成
纠错,疑问,交流: 请进入讨论区或 请点击进入页面,扫码加入微信群或Q群进行交流
获取最新文章: 扫一扫加入“创客智造”公众号
欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流
ONNX 集成
- 理解ONNX集成,包含OnnxInfer类,ONNX模型格式,从JAX到ONNX的转换,ONNX推理的性能,与MuJoCo推理的集成,导出与验证,部署到真实机器人和兼容性等
概述
ONNX(Open Neural Network Exchange)为 Open Duck Playground 提供了标准化的模型部署格式。通过 ONNX 格式,训练好的策略可以脱离 JAX/TensorFlow 环境,在轻量级的 ONNX Runtime 上运行,这为部署到边缘设备(如树莓派、Jetson Nano 等)提供了可能。
OnnxInfer 类
定义在 playground/common/onnx_infer.py 中,提供简洁的推理接口:
class OnnxInfer:
def __init__(self, onnx_model_path, input_name="obs", awd=False):
"""
Args:
onnx_model_path: ONNX 模型文件路径
input_name: 输入张量名称(默认为 "obs")
awd: 是否使用"always wait for data"模式
"""
self.ort_session = onnxruntime.InferenceSession(
self.onnx_model_path,
providers=["CPUExecutionProvider"]
)
self.input_name = input_name
self.awd = awd
def infer(self, inputs):
"""执行推理"""
if self.awd:
# AWD 模式:输入为单个 numpy 数组,自动添加 batch 维度
outputs = self.ort_session.run(
None, {self.input_name: [inputs]}
)
return outputs[0][0]
else:
# 标准模式:输入已包含 batch 维度
outputs = self.ort_session.run(
None, {self.input_name: inputs.astype("float32")}
)
return outputs[0]
两种推理模式
AWD 模式 (awd=True):
- 输入:形状
(obs_size,)的一维数组 - 内部自动包装为
(1, obs_size)并添加 batch 维度 - 输出:取 batch 中的第一个结果
(outputs[0][0]) - 用于 MuJoCo 推理(单环境)
标准模式 (awd=False):
- 输入:形状
(batch_size, obs_size)的二维数组 - 输出:完整批量结果
- 可用于批量测试
ONNX 模型格式
输入输出
| 属性 | 值 |
|---|---|
| 输入名称 | obs |
| 输入类型 | tensor(float) |
| 输入形状 | (1, obs_size) |
| 输出名称 | continuous_actions |
| 输出类型 | tensor(float) |
| 输出形状 | (1, action_size) |
| opset 版本 | 11 |
模型架构(绘制的计算图)
obs (1 x 46)
│
▼
Sub (减均值)
│
▼
Div (除标准差)
│
▼
MatMul (hidden_0, W0) ──→ Add (b0) ──→ Swish
│
▼
MatMul (hidden_1, W1) ──→ Add (b1) ──→ Swish
│
▼
MatMul (hidden_2, W2) ──→ Add (b2)
│
▼
Split (2 x 14)
│
▼
Gather (取前 14 作为均值)
│
▼
Tanh
│
▼
continuous_actions (1 x 14)
从 JAX 到 ONNX 的转换
转换流程
JAX 训练参数
│
▼
步骤 1: 提取策略网络参数
params[1].policy.params
├── MLP_0.hidden_0: {kernel, bias}
├── MLP_0.hidden_1: {kernel, bias}
└── MLP_0.hidden_2: {kernel, bias}
│
▼
步骤 2: 提取观测归一化参数
params[0].mean["state"] (46 维均值)
params[0].std["state"] (46 维标准差)
│
▼
步骤 3: 创建等效 TensorFlow Keras 模型
MLP( [256, 256, 28] )
│
▼
步骤 4: 权重复制
JAX kernel → TF Dense.kernel
JAX bias → TF Dense.bias
│
▼
步骤 5: tf2onnx 转换
tf2onnx.convert.from_keras(
model, opset=11, output_path="model.onnx"
)
策略网络参数
训练好的策略参数以 JAX Flax 格式存储:
params[0] = { "mean": {"state": mean_array}, "std": {"state": std_array} }
params[1] = {
"policy": {
"params": {
"MLP_0": {
"hidden_0": {"kernel": ..., "bias": ...},
"hidden_1": {"kernel": ..., "bias": ...},
"hidden_2": {"kernel": ..., "bias": ...},
}
}
}
}
ONNX 推理的性能
在 CPU 上,ONNX Runtime 推理非常高效:
观测维度: 46
动作维度: 14
网络结构: [256, 256, 28]
单次推理时间: ~0.1 ms
推理 FPS: ~10000
这使得 ONNX 模型非常适合实时控制应用(通常 50-100 Hz 即可)。
与 MuJoCo 推理的集成
在 MjInfer 中,ONNX 推理的集成非常简洁:
# 加载模型
self.policy = OnnxInfer(onnx_model_path, awd=True)
# 推理循环(50 Hz)
obs = self.get_obs(data, commands) # 构建 46 维观测
action = self.policy.infer(obs) # ONNX 推理 → 14 维动作
导出与验证
验证导出正确性
可以使用 ONNX Runtime 直接加载和测试导出的模型:
import onnxruntime as ort
import numpy as np
# 加载 ONNX 模型
session = ort.InferenceSession("model.onnx")
# 准备测试输入
obs = np.random.randn(1, 46).astype(np.float32)
# 推理
outputs = session.run(None, {"obs": obs})
action = outputs[0] # shape: (1, 14)
性能基准测试
onnx_infer.py 的 __main__ 部分包含简单的基准测试:
if __name__ == "__main__":
oi = OnnxInfer(args.onnx_model_path, awd=True)
times = []
for _ in range(1000):
inputs = np.random.uniform(size=obs_size).astype(np.float32)
start = time.time()
oi.infer(inputs)
times.append(time.time() - start)
print("Average time:", sum(times) / len(times))
print("Average fps:", 1 / (sum(times) / len(times)))
部署到真实机器人
将 ONNX 模型部署到真实机器人上:
- 安装 ONNX Runtime:在机器人上安装 ONNX Runtime(轻量级,支持 ARM)
- 加载模型:使用
OnnxInfer或直接调用 ONNX Runtime API - 传感器处理:实现与训练时相同的观测构建逻辑(使用真实传感器数据)
- 控制输出:将策略动作映射为舵机控制信号
# 在机器人上安装 ONNX Runtime
pip install onnxruntime
# 推理代码示例(简化)
import onnxruntime
import numpy as np
session = onnxruntime.InferenceSession("policy.onnx")
def get_action(imu_data, joint_angles, command):
obs = build_observation(imu_data, joint_angles, command)
action = session.run(None, {"obs": obs.reshape(1, -1)})[0][0]
return action
兼容性
ONNX opset 11 确保与以下框架的兼容性:
| 框架 | 兼容性 |
|---|---|
| Isaac Lab | ✅ Opset 11 是推荐版本 |
| ROS 2 | ✅ 通过 onnxruntime-cpp 集成 |
| TensorRT | ✅ 可从 ONNX 转换为 TensorRT 引擎 |
| OpenVINO | ✅ 支持 ONNX 模型导入 |
| CoreML | ✅ 支持 ONNX 模型导入 |
纠错,疑问,交流: 请进入讨论区或 请点击进入页面,扫码加入微信群或Q群进行交流
获取最新文章: 扫一扫加入“创客智造”公众号
欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流


















