OpenDuckMini强化学习框架入门教程-Runner 架构
纠错,疑问,交流: 请进入讨论区或 请点击进入页面,扫码加入微信群或Q群进行交流
获取最新文章: 扫一扫加入“创客智造”公众号
欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流
Runner 架构
- 理解Runner架构,包含数据流,基类,实现类,检查点系统,训练配置
概述
Runner 是训练系统的核心调度器,负责协调环境、策略、随机化和日志记录等组件。Open Duck Playground 采用双层 Runner 架构:通用的 BaseRunner 基类和机器人特定的 OpenDuckMiniV2Runner。
Runner 数据流
OpenDuckMiniV2Runner
│
┌────────────┴────────────┐
│ 初始化 │
│ env = Joystick(...) │
│ randomizer = domain_ │
│ randomize │
│ obs_size = ... │
│ action_size = ... │
└────────────┬────────────┘
│
┌────────────┴────────────┐
│ BaseRunner │
│ │
│ train(): │
│ 1. 配置 PPO 参数 │
│ 2. 创建网络工厂 │
│ 3. 调用 train() │
│ (来自 brax.agents.ppo) │
│ │
│ 回调: │
│ progress_fn → TensorBoard│
│ policy_params_fn → │
│ 检查点 + ONNX │
└─────────────────────────┘
BaseRunner 类
定义在 playground/common/runner.py 中。
初始化
class BaseRunner(ABC):
def __init__(self, args: argparse.Namespace) -> None:
self.args = args
self.output_dir = Path.cwd() / Path(args.output_dir)
# TensorBoard 日志
self.writer = SummaryWriter(log_dir=self.output_dir)
# JAX 编译缓存(加速 GPU 编译)
os.makedirs(".tmp", exist_ok=True)
jax.config.update("jax_compilation_cache_dir", ".tmp/jax_cache")
jax.config.update("jax_persistent_cache_min_entry_size_bytes", -1)
jax.config.update("jax_persistent_cache_min_compile_time_secs", 0)
JAX 编译缓存的配置对加速重复训练至关重要。首次训练时 JAX 需要编译所有计算图,后续恢复训练或微调时可直接使用缓存。
progress_fn — 进度回调
训练过程中定期回调,用于记录和打印训练指标:
def progress_callback(self, num_steps: int, metrics: dict) -> None:
# 记录所有指标到 TensorBoard
for metric_name, metric_value in metrics.items():
self.writer.add_scalar(metric_name, metric_value, num_steps)
# 打印关键指标
print(f"STEP: {num_steps} reward: {metrics['eval/episode_reward']} ...")
policy_params_fn — 策略参数回调
定期保存检查点和导出 ONNX 模型:
def policy_params_fn(self, current_step, make_policy, params):
# 1. 使用 Orbax 保存 JAX 参数
orbax_checkpointer = ocp.PyTreeCheckpointer()
save_args = orbax_utils.save_args_from_target(params)
timestamp = datetime.now().strftime("%Y_%m_%d_%H%M%S")
checkpoint_path = f"{self.output_dir}/{timestamp}_{current_step}"
orbax_checkpointer.save(checkpoint_path, params, force=True, save_args=save_args)
# 2. 导出 ONNX 模型
onnx_path = f"{self.output_dir}/{timestamp}_{current_step}.onnx"
export_onnx(params, self.action_size, self.ppo_params, self.obs_size, onnx_path)
train() — 训练主循环
def train(self) -> None:
# 1. 加载 PPO 配置
self.ppo_params = locomotion_params.brax_ppo_config(...)
# 2. 配置网络工厂
if "network_factory" in self.ppo_params:
network_factory = functools.partial(
ppo_networks.make_ppo_networks,
**self.ppo_params.network_factory
)
# 3. 配置训练函数
train_fn = functools.partial(
ppo.train,
**self.ppo_training_params,
network_factory=network_factory,
randomization_fn=self.randomizer, # 域随机化
progress_fn=self.progress_callback, # 进度回调
policy_params_fn=self.policy_params_fn, # 检查点保存
restore_checkpoint_path=self.restore_checkpoint_path,
)
# 4. 启动训练
_, params, _ = train_fn(
environment=self.env,
eval_env=self.eval_env,
wrap_env_fn=wrapper.wrap_for_brax_training,
)
OpenDuckMiniV2Runner 类
定义在 playground/open_duck_mini_v2/runner.py 中。
初始化
class OpenDuckMiniV2Runner(BaseRunner):
def __init__(self, args):
super().__init__(args)
# 选择环境
available_envs = {
"joystick": (joystick, joystick.Joystick),
"standing": (standing, standing.Standing),
}
self.env_file = available_envs[args.env]
# 创建环境
self.env_config = self.env_file[0].default_config()
self.env = self.env_file[1](task=args.task)
self.eval_env = self.env_file[1](task=args.task)
# 配置训练组件
self.randomizer = randomize.domain_randomize # 域随机化函数
self.action_size = self.env.action_size # 动作空间维度
self.obs_size = self.env.observation_size["state"][0] # 观测维度
# 检查点恢复
self.restore_checkpoint_path = args.restore_checkpoint_path
main() 入口
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--output_dir", type=str, default="checkpoints")
parser.add_argument("--num_timesteps", type=int, default=150000000)
parser.add_argument("--env", type=str, default="joystick")
parser.add_argument("--task", type=str, default="flat_terrain")
parser.add_argument("--restore_checkpoint_path", type=str, default=None)
args = parser.parse_args()
runner = OpenDuckMiniV2Runner(args)
runner.train()
检查点系统
保存格式
检查点使用 Orbax 格式保存,这是一种为 JAX 参数定制的序列化格式:
checkpoints/
├── 2025_01_15_123456_1000000/ # Orbax 检查点目录
│ ├── checkpoint # 参数数据
│ └── ...
├── 2025_01_15_123456_1000000.onnx # 同一步导出的 ONNX 模型
├── events.out.tfevents.* # TensorBoard 事件文件
└── ...
恢复训练
从检查点恢复训练:
uv run playground/open_duck_mini_v2/runner.py \
--restore_checkpoint_path checkpoints/2025_01_15_123456_1000000
恢复时,PPO 训练函数会:
- 加载保存的参数作为初始策略
- 继续训练直到达到目标步数
- 如果总步数未达到
num_timesteps,会继续训练剩余步数
训练配置
PPO 参数
训练配置来自 mujoco_playground.config.locomotion_params.brax_ppo_config("BerkeleyHumanoidJoystickFlatTerrain")。主要参数包括:
| 参数 | 典型值 | 说明 |
|---|---|---|
num_timesteps |
150M-300M | 总训练步数 |
num_envs |
2048 | 并行环境数量 |
batch_size |
512 | 批次大小 |
unroll_length |
10 | 轨迹长度 |
learning_rate |
3e-4 | 学习率 |
clip_epsilon |
0.3 | PPO clip 阈值 |
entropy_cost |
1e-3 | 熵正则化权重 |
gae_lambda |
0.95 | GAE 衰减因子 |
network_factory.policy_hidden_layer_sizes |
[256, 256] | 策略网络隐层大小 |
记录指标
训练过程中 TensorBoard 记录的主要指标:
- eval/episode_reward:评估回合总奖励(主要收敛指标)
- eval/episode_reward_std:评估奖励标准差
- loss/policy_loss:策略损失函数值
- loss/value_loss:价值网络损失
- loss/total_loss:总损失
- metrics/xxx:环境特定指标(各奖励分量、摆动峰值等)
纠错,疑问,交流: 请进入讨论区或 请点击进入页面,扫码加入微信群或Q群进行交流
获取最新文章: 扫一扫加入“创客智造”公众号
欢迎加入我们的openduckmini交流群,微信扫描右侧二维码立即进群交流


















