turtlebot3-burger_150.png
turtlebot3-waffle-pi_150.png
turtlebot3-arm_150.png
walking-y2_150.png
turbot3-multi_150.png
turbot3-dl-ros1_150.png
turbot3-ai.png
turbot3-dl-ros2_150.png
turbot3-slam_150.png
turbot3-arm_150.png
turtlebot4-lite_150.png
turtlebot4-pro_150.png
turbot4-dl_150.png
turbot4-ai_150.png
aidriving-racebot_150.png
aidriving-autodrive_150.png
turtlebot-arm_150.png
openmanipulator-x_150.png
Home » OpenDuckMini强化学习框架入门教程 » OpenDuckMini强化学习框架入门教程-ONNX集成

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 模型部署到真实机器人上:

  1. 安装 ONNX Runtime:在机器人上安装 ONNX Runtime(轻量级,支持 ARM)
  2. 加载模型:使用 OnnxInfer 或直接调用 ONNX Runtime API
  3. 传感器处理:实现与训练时相同的观测构建逻辑(使用真实传感器数据)
  4. 控制输出:将策略动作映射为舵机控制信号
# 在机器人上安装 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交流群,微信扫描右侧二维码立即进群交流

群二维码

标签: OpenDuckMini强化学习框架