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 Algorithms"(John 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 JAX"(PPO Agent 的 JAX 紧凑实现)。也就是说,整个 PPO 支持被刻意收敛在单一文件中,便于阅读与二次开发。
该模块对应的核心 API 结构为:
- 模块
dopamine.jax.agents.ppo:包入口,仅包含ppo_agent子模块; - 模块
dopamine.jax.agents.ppo.ppo_agent:PPO 算法实现主体,见 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 时自动包装为 tuple |
observation_shape | 必填(tuple) | 观测形状,构造时assert isinstance(observation_shape, tuple) |
action_limits | 必填 | 动作下/上界对,用于将高斯分布缩放到合法动作区间 |
stack_size | 1 | 状态栈帧数 |
update_horizon | 1 | 回放缓冲的更新视野 |
network | continuous_networks.PPOActorCriticNetwork | PPO 网络结构 |
num_layers/hidden_units/activation | 2 / 64 /'tanh' | 共享的 Actor/Critic 网络层数、隐藏单元数与激活函数 |
update_period | 2048 | 两次 PPO 更新之间的环境步数(收集一整段轨迹的周期) |
num_epochs | 10 | 对同一批轨迹进行梯度更新的轮数 |
batch_size | 64 | minibatch 大小 |
gamma | 0.99 | 折扣因子 |
lambda_ | 0.95 | GAE(广义优势估计)参数 |
epsilon | 0.2 | PPO 裁剪阈值 |
vf_coefficient | 0.5 | critic 损失系数 |
entropy_coefficient | 0.0 | 熵正则系数 |
clip_critic_loss | True | 是否对 critic 损失做裁剪 |
optimizer | 'adam' | 优化器名称(由dqn_agent.create_optimizer创建) |
max_gradient_norm | 0.5 | 全局梯度范数裁剪上限 |
seed | None(取当前时间) | 内部 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_limits(ppo_agent.py#L517-L529),这为替换自定义网络保留了灵活性。
数据管线:on-policy 轨迹收集与顺序采样回放
PPO 是 on-policy 算法,其数据管线与 DQN 的经验回放有本质区别:
_build_replay_buffer()(ppo_agent.py#L583-L596)组合了:accumulator.TransitionAccumulator(accumulator.py),负责按stack_size、update_horizon、gamma累积片段;samplers.SequentialSamplingDistribution(samplers.py),顺序采样(sort_samples=False),保证轨迹时序不被打乱;- 外层
replay_buffer.ReplayBuffer(replay_buffer.py)。
begin_episode/step中,每步都调用select_action(ppo_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、裁剪目标与多轮 minibatch
train()(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_{t+1}) * (1 - terminal_t) - V(s_t) A_t = delta_t + lambda_ * gamma * A_{t+1} * (1 - terminal_t) returns = advantages + q_values # 即 r + gamma*V(s')其中轨迹末端的V(s_{t+1})使用"整网前向"得到的next_q_value(因下一步动作尚未采样),从而把未来回报的估计延伸到当前轨迹之外。这段代码中的terminals掩码保证了跨 episode 边界不泄漏价值。
2. 计算旧策略的对数概率与采样动作
用method=network_def.actor批量计算(state, action)下的log_probability(同样stop_gradient),并采样出sampled_actions供后续统计输出。二者都会在后续作为"旧策略"基准参与裁剪。
3. 构造并打乱 minibatch
create_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_loss=True,采用 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_allowlist='tensorboard'默认只写 tensorboard)。
Actor-Critic 网络设计
默认网络PPOActorCriticNetwork定义在 dopamine/jax/continuous_networks.py#L412-L496,由setup()组合两个子网络:
PPOActorNetwork(continuous_networks.py#L314-L381):状态经若干nn.Dense+ 激活(默认 tanh)后,输出高斯分布的loc(orthogonal(sqrt(0.01))初始化)与状态无关、零初始化的scale_diag(jnp.exp保证正数)。当action_limits非空时,通过_transform_distribution(tfp 的Tanh + Shift/Scalebijector)把分布缩放到动作合法区间;熵在变换之前计算(因变换的 Jacobian 不恒定)。PPOCriticNetwork(continuous_networks.py#L384-L409):相同 MLP 主干,输出 1 维标量价值(orthogonal(1.0)初始化)。
网络输出统一封装为PPOActorOutput(sampled_action, log_probability, entropy)、PPOCriticOutput(q_value)与PPOActorCriticOutput(continuous_networks.py#L55-L73)。初始化时network_def.init(init_key, self.state, init_key)同时初始化两个子网络,而network_def.apply(..., method=network_def.actor/critic)则支持只调用单侧子网络——这正是训练代码里反复使用的调用方式。
对于离散动作(Atari),官方配置切换到networks.PPODiscreteActorCriticNetwork(见 dopamine/jax/networks.py)。
官方 gin 配置实战:MuJoCo 与 Atari
仓库为 PPO 提供了两份开箱即用的配置,分别对应论文附录 A 的表 3(MuJoCo)与表 5(Atari)。
MuJoCo 连续控制配置 dopamine/jax/agents/ppo/configs/ppo.gin
import 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_rate=True,anneal_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_preprocessing=True的 gym 环境包装、legacy gym(v2版本),经由create_continuous_runner.schedule = 'continuous_train'与agent_name = 'ppo'接入 continuous_domains/run_experiment.py 的连续训练循环。
Atari 离散控制配置 dopamine/jax/agents/ppo/configs/ppo_atari.gin
PPOAgent.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各乘 8(1024 = 8 × 128、256 = 8 × 32),使每个训练迭代采样到的样本数与原实现一致。 - 离散网络:切换到
networks.PPODiscreteActorCriticNetwork,动作采样为离散分布。 - 超参数调整:
epsilon降至 0.1、entropy_coefficient提升到 0.01(鼓励探索),num_epochs降为 3,学习率 2.5e-4。 - Atari 预处理:
use_ppo_preprocessing=True的 Atari 环境(灰度、帧堆叠等),入口在 discrete_domains/atari_lib.py。
运行方式
两份配置都依赖create_agent.agent_name = 'ppo'与对应的 runner 绑定。运行连续域(MuJoCo)实验可参考 continuous_domains/train.py:
python -m dopamine.continuous_domains.train \ --base_dir=/tmp/dopamine/ppo \ --gin_files='dopamine/jax/agents/ppo/configs/ppo.gin'Atari 实验对应 discrete_domains/train.py,将gin_files换成ppo_atari.gin并安装 atari 依赖即可。注意ppo.gin依赖use_legacy_gym=True(gym 旧接口),需要相应版本的环境库。
测试与校验
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),仅供参考