ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Dopamine 中 PPO Agent 的 JAX 实现:从 PPOAgent 源码到 gin 配置实战

Dopamine 中 PPO Agent 的 JAX 实现:从 PPOAgent 源码到 gin 配置实战 Dopamine 中 PPO Agent 的 JAX 实现从 PPOAgent 源码到 gin 配置实战【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine导读本文围绕 Dopamine 仓库中 docs/api_docs/python/dopamine/jax/agents/ppo.md 所定义的dopamine.jax.agents.ppo模块展开它是Proximal Policy Optimization AlgorithmsJohn Schulman 等人arXiv:1707.06347在 JAX 下的一套紧凑实现。文章将带你逐层拆解PPOAgent的类结构、GAE 优势估计与裁剪式策略更新的底层计算、Actor/Critic 网络设计与两类官方 gin 配置MuJoCo 与 Atari使读者既能直接用官方超参数复现实验也能按需改造网络与训练流程。模块概览dopamine.jax.agents.ppo在 dopamine/jax/agents/ppo/init.py 中该包只暴露一个子模块ppo_agent其 docstring 原文为Compact implementation of a PPO agent in JAXPPO Agent 的 JAX 紧凑实现。也就是说整个 PPO 支持被刻意收敛在单一文件中便于阅读与二次开发。该模块对应的核心 API 结构为模块dopamine.jax.agents.ppo包入口仅包含ppo_agent子模块模块dopamine.jax.agents.ppo.ppo_agentPPO 算法实现主体见 dopamine/jax/agents/ppo/ppo_agent.py类dopamine.jax.agents.ppo.ppo_agent.PPOAgent继承自JaxDQNAgent见 dopamine/jax/agents/dqn/dqn_agent.py是面向训练/评估使用的对外 Agent 类。从继承关系可以推断PPO 复用了 Dopamine 中 DQN Agent 的骨架如 checkpoint/bundle 机制、summary writer、collector 分发等基础设施但把Q 学习式的离线更新替换为 on-policy 的轨迹采样与多轮 minibatch 更新这是理解该实现的关键切入点。PPOAgent构造参数与默认值全解析PPOAgent定义在 dopamine/jax/agents/ppo/ppo_agent.py#L416-L772构造函数带gin.configurable装饰器所有超参数均可通过 gin 文件覆盖。其默认参数与含义如下参数默认值含义action_shape必填int 或 tuple动作空间维度传入 int 时自动包装为 tupleobservation_shape必填tuple观测形状构造时assert isinstance(observation_shape, tuple)action_limits必填动作下/上界对用于将高斯分布缩放到合法动作区间stack_size1状态栈帧数update_horizon1回放缓冲的更新视野networkcontinuous_networks.PPOActorCriticNetworkPPO 网络结构num_layers/hidden_units/activation2 / 64 /tanh共享的 Actor/Critic 网络层数、隐藏单元数与激活函数update_period2048两次 PPO 更新之间的环境步数收集一整段轨迹的周期num_epochs10对同一批轨迹进行梯度更新的轮数batch_size64minibatch 大小gamma0.99折扣因子lambda_0.95GAE广义优势估计参数epsilon0.2PPO 裁剪阈值vf_coefficient0.5critic 损失系数entropy_coefficient0.0熵正则系数clip_critic_lossTrue是否对 critic 损失做裁剪optimizeradam优化器名称由dqn_agent.create_optimizer创建max_gradient_norm0.5全局梯度范数裁剪上限seedNone取当前时间内部 RNG 种子int(time.time() * 1e6)值得注意的两处实现细节学习率等优化器参数不在这里PPOAgent只接收optimizer名称实际学习率、eps、退火等由 dqn_agent.py 中的create_optimizer通过 gin 配置见下文配置节。网络构造分叉当network.__name__ PPOActorCriticNetwork时会传入num_layers/hidden_units/activation等结构参数否则仅传action_shape与action_limitsppo_agent.py#L517-L529这为替换自定义网络保留了灵活性。数据管线on-policy 轨迹收集与顺序采样回放PPO 是 on-policy 算法其数据管线与 DQN 的经验回放有本质区别_build_replay_buffer()ppo_agent.py#L583-L596组合了accumulator.TransitionAccumulatoraccumulator.py负责按stack_size、update_horizon、gamma累积片段samplers.SequentialSamplingDistributionsamplers.py顺序采样sort_samplesFalse保证轨迹时序不被打乱外层replay_buffer.ReplayBufferreplay_buffer.py。begin_episode/step中每步都调用select_actionppo_agent.py#L391-L413从策略网络获得高斯分布的均值和方差jax.random.split出新 key 后采样动作并stop_gradient阻止梯度回流。_train_step()ppo_agent.py#L651-L706是触发点只有当回放缓冲累计转移数add_count update_period时才执行训练训练完self._replay.clear()清空缓冲并把最后观测重新记录进缓冲以保证轨迹连续性self._record_observation(self._last_observation)。这正是收集一段、更新一次的 on-policy 节奏。核心训练逻辑GAE、裁剪目标与多轮 minibatchtrain()ppo_agent.py#L46-L190)是 PPO 更新的总控可分为四个阶段1. 计算价值与 GAE先对整段轨迹做jax.vmap批量前向取出 critic 的q_value并stop_gradient随后calculate_advantages_and_returns()ppo_agent.py#L193-L220按论文公式 (11)(12) 从后向前递推delta_t r_t gamma * V(s_{t1}) * (1 - terminal_t) - V(s_t) A_t delta_t lambda_ * gamma * A_{t1} * (1 - terminal_t) returns advantages q_values # 即 r gamma*V(s)其中轨迹末端的V(s_{t1})使用整网前向得到的next_q_value因下一步动作尚未采样从而把未来回报的估计延伸到当前轨迹之外。这段代码中的terminals掩码保证了跨 episode 边界不泄漏价值。2. 计算旧策略的对数概率与采样动作用methodnetwork_def.actor批量计算(state, action)下的log_probability同样stop_gradient并采样出sampled_actions供后续统计输出。二者都会在后续作为旧策略基准参与裁剪。3. 构造并打乱 minibatchcreate_minibatches_and_shuffle()ppo_agent.py#L223-L266要求states.shape[0] % batch_size 0否则直接 assert 失败把整段轨迹按batch_size切块num_batches n // batch_size再用jax.random.permutation打乱批块顺序——批内时序保持连续。4. 多轮 minibatch 更新外层循环num_epochs次、内层遍历全部 minibatch调用train_minibatch()ppo_agent.py#L269-L388带jax.jit且将网络/优化器/标量超参声明为static_argnames。其损失函数实现了 PPO 的两个核心目标Actor 裁剪目标论文公式 (7)ratio jnp.exp(log_probability - old_log_probability) actor_loss jnp.mean( -jnp.minimum( ratio * advantages, jnp.clip(ratio, 1.0 - epsilon, 1.0 epsilon) * advantages, ) )Critic 损失默认clip_critic_lossTrue采用 PPO 实现细节 的mse_loss。熵正则entropy_loss jnp.mean(actor_output.entropy)最终总损失为actor_loss vf_coefficient * critic_loss - entropy_coefficient * entropy_loss。此外在loss_fn之外还有两处关键预处理minibatch 内优势归一化(advantages - mean) / (std 1e-8)以及优化器链optax.clip_by_global_norm(max_gradient_norm)的全局梯度裁剪ppo_agent.py#L575-L580。训练返回的loss_stats包含Losses/Combined、Losses/Actor、Losses/Critic、Losses/Entropy以及Values/SampledAction{i}等标量由_train_step写入 TensorBoard 或经collector_dispatcher分发collector_allowlisttensorboard默认只写 tensorboard。Actor-Critic 网络设计默认网络PPOActorCriticNetwork定义在 dopamine/jax/continuous_networks.py#L412-L496由setup()组合两个子网络PPOActorNetworkcontinuous_networks.py#L314-L381状态经若干nn.Dense 激活默认 tanh后输出高斯分布的locorthogonal(sqrt(0.01))初始化与状态无关、零初始化的scale_diagjnp.exp保证正数。当action_limits非空时通过_transform_distributiontfp 的Tanh Shift/Scalebijector把分布缩放到动作合法区间熵在变换之前计算因变换的 Jacobian 不恒定。PPOCriticNetworkcontinuous_networks.py#L384-L409相同 MLP 主干输出 1 维标量价值orthogonal(1.0)初始化。网络输出统一封装为PPOActorOutput(sampled_action, log_probability, entropy)、PPOCriticOutput(q_value)与PPOActorCriticOutputcontinuous_networks.py#L55-L73。初始化时network_def.init(init_key, self.state, init_key)同时初始化两个子网络而network_def.apply(..., methodnetwork_def.actor/critic)则支持只调用单侧子网络——这正是训练代码里反复使用的调用方式。对于离散动作Atari官方配置切换到networks.PPODiscreteActorCriticNetwork见 dopamine/jax/networks.py。官方 gin 配置实战MuJoCo 与 Atari仓库为 PPO 提供了两份开箱即用的配置分别对应论文附录 A 的表 3MuJoCo与表 5Atari。MuJoCo 连续控制配置 dopamine/jax/agents/ppo/configs/ppo.ginimport dopamine.continuous_domains.run_experiment import dopamine.discrete_domains.gym_lib import dopamine.jax.agents.ppo.ppo_agent import dopamine.jax.agents.dqn.dqn_agent import dopamine.jax.continuous_networks import dopamine.jax.replay_memory.replay_buffer PPOAgent.network continuous_networks.PPOActorCriticNetwork PPOAgent.num_layers 2 PPOAgent.hidden_units 64 PPOAgent.activation tanh PPOAgent.update_period 2048 PPOAgent.optimizer adam PPOAgent.max_gradient_norm 0.5 create_optimizer.learning_rate 3e-4 create_optimizer.eps 1e-5 create_optimizer.anneal_learning_rate True create_optimizer.anneal_steps 160_000 # 500 iterations * 10 epochs * 2048 timesteps / 64 batches PPOAgent.num_epochs 10 PPOAgent.batch_size 64 PPOAgent.gamma 0.99 PPOAgent.lambda_ 0.95 PPOAgent.epsilon 0.2 PPOAgent.vf_coefficient 0.5 PPOAgent.entropy_coefficient 0.0 PPOAgent.clip_critic_loss True PPOAgent.seed None # Seed with the current time create_gym_environment.environment_name HalfCheetah create_gym_environment.version v2 create_gym_environment.use_legacy_gym True create_gym_environment.use_ppo_preprocessing True create_continuous_runner.schedule continuous_train create_continuous_agent.agent_name ppo ContinuousTrainRunner.create_environment_fn gym_lib.create_gym_environment ContinuousRunner.num_iterations 500 ContinuousRunner.training_steps 2048 ContinuousRunner.max_steps_per_episode None ReplayBuffer.max_capacity 2048 ReplayBuffer.batch_size 2048要点解读学习率退火anneal_learning_rateTrueanneal_steps 160_000的注释精确给出了推导500 iterations × 10 epochs × 2048 timesteps / 64 batches——即整个实验的梯度更新次数学习率在该步数内线性衰减到 0。回放缓冲即轨迹桶ReplayBuffer.max_capacity 2048与update_period 2048、training_steps 2048三者一致说明缓冲只装当前这一段轨迹配合SequentialSamplingDistribution保持时序。环境侧使用use_ppo_preprocessingTrue的 gym 环境包装、legacy gymv2版本经由create_continuous_runner.schedule continuous_train与agent_name ppo接入 continuous_domains/run_experiment.py 的连续训练循环。Atari 离散控制配置 dopamine/jax/agents/ppo/configs/ppo_atari.ginPPOAgent.network networks.PPODiscreteActorCriticNetwork PPOAgent.update_period 1024 # 8 * 128 PPOAgent.optimizer adam PPOAgent.max_gradient_norm 0.5 create_optimizer.learning_rate 2.5e-4 create_optimizer.eps 1e-5 create_optimizer.anneal_learning_rate True create_optimizer.anneal_steps 117_600 # 980 iterations * 3 epochs * 10240 timesteps / 256 batches PPOAgent.num_epochs 3 PPOAgent.batch_size 256 # 8 * 32 PPOAgent.gamma 0.99 PPOAgent.lambda_ 0.95 PPOAgent.epsilon 0.1 PPOAgent.vf_coefficient 0.5 PPOAgent.entropy_coefficient 0.01 PPOAgent.clip_critic_loss True atari_lib.create_atari_environment.game_name Pong atari_lib.create_atari_environment.use_ppo_preprocessing True create_runner.schedule continuous_train create_agent.agent_name ppo Runner.num_iterations 980 Runner.training_steps 10240 ReplayBuffer.max_capacity 1024 ReplayBuffer.batch_size 1024与 MuJoCo 配置的关键差异单 actor 等价补偿配置文件注释明确指出原论文使用 8 个并行 actor而本仓库是单 actor 实现因此把batch_size和update_period各乘 81024 8 × 128、256 8 × 32使每个训练迭代采样到的样本数与原实现一致。离散网络切换到networks.PPODiscreteActorCriticNetwork动作采样为离散分布。超参数调整epsilon降至 0.1、entropy_coefficient提升到 0.01鼓励探索num_epochs降为 3学习率 2.5e-4。Atari 预处理use_ppo_preprocessingTrue的 Atari 环境灰度、帧堆叠等入口在 discrete_domains/atari_lib.py。运行方式两份配置都依赖create_agent.agent_name ppo与对应的 runner 绑定。运行连续域MuJoCo实验可参考 continuous_domains/train.pypython -m dopamine.continuous_domains.train \ --base_dir/tmp/dopamine/ppo \ --gin_filesdopamine/jax/agents/ppo/configs/ppo.ginAtari 实验对应 discrete_domains/train.py将gin_files换成ppo_atari.gin并安装 atari 依赖即可。注意ppo.gin依赖use_legacy_gymTruegym 旧接口需要相应版本的环境库。测试与校验PPOAgent 单元测试 覆盖了 Agent 的构造create_agent仅需action_shape、action_limits、observation_shape、update_period、seed即可实例化、参数存取、训练步行为等配合 losses_test.py、continuous_networks_test.py 可以交叉验证 GAE、裁剪损失与网络前向的正确性。若你修改了 PPO 相关实现运行python -m tests.dopamine.jax.agents.ppo.ppo_agent_test是基本的回归手段。小结Dopamine 的 PPO JAX 实现以单文件ppo_agent.py承载完整算法继承JaxDQNAgent复用工程基建用TransitionAccumulator SequentialSamplingDistribution构建 on-policy 轨迹管线以 GAE 估计优势、裁剪式目标更新策略、可选裁剪的 critic 损失与全局梯度裁剪保证训练稳定两份 gin 配置则忠实复刻论文的 MuJoCo/Atari 超参数含单 actor 补偿可直接用于复现或作为调参起点。对于希望深入 PPO 实现细节或在 JAX 生态中定制 RL 算法的研究者这份实现是结构清晰、易于改造的参考范本。【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表