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强化学习框架入门教程-Runner 架构

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 训练函数会:

  1. 加载保存的参数作为初始策略
  2. 继续训练直到达到目标步数
  3. 如果总步数未达到 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交流群,微信扫描右侧二维码立即进群交流

群二维码

标签: OpenDuckMini强化学习框架