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
恢复时:
- PPO 训练器加载 Orbax 检查点中的参数
- 继续训练直到达到
num_timesteps - 注意:
num_timesteps指总训练步数,恢复训练时如果之前已训练了 100M 步,设置--num_timesteps 200000000会继续训练 100M 步
最佳实践
- 训练完成后:选择最新或奖励最高的 ONNX 模型进行部署
- 回滚:保存多个检查点,方便回滚到之前的训练状态
- ONNX 兼容性:opset=11 确保与 Isaac Lab、ROS 等框架兼容
- 批量测试:可使用 ONNX Runtime 批量测试多个 ONNX 模型,选择性能最佳者
纠错,疑问,交流: 请进入讨论区或 请点击进入页面,扫码加入微信群或Q群进行交流
获取最新文章: 扫一扫加入“创客智造”公众号
欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流


















