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强化学习框架入门教程-检查点与导出

OpenDuckMini强化学习框架入门教程-检查点与导出

纠错,疑问,交流: 请进入讨论区请点击进入页面,扫码加入微信群或Q群进行交流

获取最新文章: 扫一扫加入“创客智造”公众号

欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流

群二维码

检查点与导出

  • 理解检查点和导出,包含检查点保存,ONNX 模型导出,推理使用 ONNX 模型,从检查点恢复训练等

概述

训练过程中,系统的检查点和模型导出机制确保训练成果被持久化保存,并能够转换为可在真实机器人上部署的格式。Open Duck Playground 使用 Orbax 保存 JAX 参数检查点,并通过 TensorFlow → tf2onnx 管线导出为 ONNX 格式。

检查点保存

保存时机

检查点通过 policy_params_fn 回调定期保存。该回调由 Brax 的 PPO 训练函数调用,保存频率取决于训练配置。

保存格式

使用 Orbax(基于 Google 的 Orbax 库)保存 JAX 参数:

orbax_checkpointer = ocp.PyTreeCheckpointer()
save_args = orbax_utils.save_args_from_target(params)

# 命名格式: <日期>_<步数>
d = datetime.now().strftime("%Y_%m_%d_%H%M%S")
path = f"{self.output_dir}/{d}_{current_step}"

orbax_checkpointer.save(path, params, force=True, save_args=save_args)

输出目录结构

checkpoints/
├── 2025_01_15_123456_1000000/        # Orbax 检查点目录
│   ├── checkpoint                    # 参数数据文件
│   └── ...                           # 元数据文件
├── 2025_01_15_123456_1000000.onnx    # 对应 ONNX 模型
├── 2025_01_16_091234_2000000/
│   ├── checkpoint
│   └── ...
├── 2025_01_16_091234_2000000.onnx
├── events.out.tfevents.*             # TensorBoard 事件文件
└── ...

ONNX 模型导出

export_onnx.py 中的 export_onnx() 函数实现了 JAX 参数到 ONNX 模型的转换。

导出管线

JAX 参数 (params)
    │
    ▼
提取策略网络参数 (params[1].policy.params)
    │
    ▼
创建 TensorFlow 模型副本
    │
    ▼
将 JAX 权重复制到 TensorFlow 层
    │
    ▼
使用 tf2onnx 转换为 ONNX 格式
    │
    ▼
ONNX 模型文件 (.onnx)

实现细节

1. 创建 MLP 模型

使用 TensorFlow 的 Keras 构建与训练时相同的 MLP 架构:

class MLP(tf.keras.Model):
    def __init__(self, layer_sizes, activation=tf.nn.relu, mean_std=None):
        self.mlp_block = tf.keras.Sequential()
        for size in layer_sizes:
            self.mlp_block.add(layers.Dense(size, activation=activation))
        
        # 状态归一化
        if mean_std is not None:
            self.mean = tf.Variable(mean_std[0], trainable=False)
            self.std = tf.Variable(mean_std[1], trainable=False)
    
    def call(self, inputs):
        if self.mean is not None:
            inputs = (inputs - self.mean) / self.std
        logits = self.mlp_block(inputs)
        loc, _ = tf.split(logits, 2, axis=-1)  # 只取均值
        return tf.tanh(loc)  # 输出在 [-1, 1] 范围

2. 权重迁移

从 JAX 参数提取网络权重并迁移到 TensorFlow 模型:

def transfer_weights(jax_params, tf_model):
    for layer_name, layer_params in jax_params.items():
        tf_layer = tf_model.get_layer("MLP_0").get_layer(name=layer_name)
        if isinstance(tf_layer, tf.keras.layers.Dense):
            kernel = np.array(layer_params["kernel"])
            bias = np.array(layer_params["bias"])
            tf_layer.set_weights([kernel, bias])

JAX 参数结构:

params[1].policy.params
    ├── MLP_0
    │   ├── hidden_0: {kernel, bias}
    │   ├── hidden_1: {kernel, bias}
    │   └── hidden_2: {kernel, bias}
    └── (其他层)

3. ONNX 转换

使用 tf2onnx 将 Keras 模型转换为 ONNX:

# opset=11 以兼容 Isaac Lab
model_proto, _ = tf2onnx.convert.from_keras(
    tf_policy_network,
    input_signature=[tf.TensorSpec(shape=(1, obs_size), dtype=tf.float32, name="obs")],
    opset=11,
    output_path=output_path,
)

转换参数:

参数 说明
opset 11 ONNX opset 版本,与 Isaac Lab 兼容
输入名 "obs" 观测输入张量名称
输出名 "continuous_actions" 连续动作输出名称
输入形状 (1, obs_size) 批次大小为 1,obs_size 为观测维度

推理使用 ONNX 模型

导出后的 ONNX 模型通过 OnnxInfer 类加载和使用:

from playground.common.onnx_infer import OnnxInfer

policy = OnnxInfer("model.onnx", awd=True)
action = policy.infer(obs)

ONNX 模型结构

Input: obs (float32, shape: [1, obs_size])
    │
    ▼
Normalization: (obs - mean) / std
    │
    ▼
MLP Network:
  Dense(256) → Swish activation
  Dense(256) → Swish activation
  Dense(28)  → split into loc(14) + scale(14)
    │
    ▼
Output: tanh(loc) (float32, shape: [1, 14])

从检查点恢复训练

要恢复之前中断的训练:

uv run playground/open_duck_mini_v2/runner.py \
    --restore_checkpoint_path checkpoints/2025_01_15_123456_1000000

恢复时:

  1. PPO 训练器加载 Orbax 检查点中的参数
  2. 继续训练直到达到 num_timesteps
  3. 注意:num_timesteps总训练步数,恢复训练时如果之前已训练了 100M 步,设置 --num_timesteps 200000000 会继续训练 100M 步

最佳实践

  1. 训练完成后:选择最新或奖励最高的 ONNX 模型进行部署
  2. 回滚:保存多个检查点,方便回滚到之前的训练状态
  3. ONNX 兼容性:opset=11 确保与 Isaac Lab、ROS 等框架兼容
  4. 批量测试:可使用 ONNX Runtime 批量测试多个 ONNX 模型,选择性能最佳者

纠错,疑问,交流: 请进入讨论区请点击进入页面,扫码加入微信群或Q群进行交流

获取最新文章: 扫一扫加入“创客智造”公众号

欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流

群二维码

标签: OpenDuckMini强化学习框架