YAOTU INSIGHTS

深入解析 Dopamine 中的 RainbowAgent:简化版 Rainbow 算法实现与源码剖析

深入解析 Dopamine 中的 RainbowAgent:简化版 Rainbow 算法实现与源码剖析
深入解析 Dopamine 中的 RainbowAgent简化版 Rainbow 算法实现与源码剖析【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine导读RainbowAgent是 Google Dopamine 研究框架TensorFlow 版本中基于 Rainbow 算法Hessel et al., 2018的精简实现它去掉了原论文中价值函数网络的诸多复杂组件只保留了被实验证明对 Atari 游戏表现提升最显著的三项核心技术n 步更新n-step updates、优先经验回放prioritized replay与分布式强化学习distributional RL。本文以 RainbowAgent API 文档 为核心骨架结合 rainbow_agent.py 源码、RainbowNetwork 网络实现 与 rainbow.gin 配置完整讲解该 Agent 的类结构、全部构造参数、C51 目标分布投影算法与训练损失计算过程并给出可直接运行的 gin 配置说明帮助你掌握如何在 Dopamine 中配置、定制和深入理解这个分布式的 Rainbow Agent。RainbowAgent 是什么简化版 Rainbow 的三项核心组件根据 rainbow_agent 模块文档 的定义RainbowAgent是一个紧凑的、经过简化的 Rainbow agent 实现A compact implementation of a simplified Rainbow agent它继承自DQNAgent代码位于 dopamine/tf/agents/rainbow/rainbow_agent.py。完整版 Rainbow 论文组合了六项改进Double DQN、Prioritized Experience Replay、Dueling Networks、Multi-step Learning、Distributional RL、Noisy Nets而 Dopamine 的这份实现只选取了其中三个被实验证明对 Atari 游戏智能体性能影响最显著的组件n 步更新n-step updates用累计 n 步的回报来更新价值函数对应构造参数update_horizon优先经验回放prioritized replay按时序差分误差的重要性对经验样本加权采样对应构造参数replay_scheme分布式强化学习distributional RL不再预测期望值 $Q(s,a)$而是直接预测回报的价值分布即 C51 算法Bellemare et al., 2017对应num_atoms、vmin、vmax等参数。同时源码开头的模块注释rainbow_agent.py明确说明该实现还刻意去掉了一些次要的超参数选择以保持实现的简洁性将优先回放的重要性采样指数 beta 固定为 0.5而不是像原论文那样从 0.4 线性增长到 1.0移除了 alpha 参数在论文中 alpha 全程固定为 0.5。这两个简化都基于作者在 Asterix、Pong、Q*Bert、Seaquest、Space Invaders 五款游戏上的实验观察详见 rainbow_agent.py 中的注释固定 beta0.5 除 Pong 外表现与动态调度相当甚至更好而直接使用tf.sqrt()作为损失加权也符合面对平方损失取平方根的直观语义。类签名与继承关系RainbowAgent在源码中的定义方式为gin.configurable class RainbowAgent(dqn_agent.DQNAgent): A compact implementation of a simplified Rainbow agent.rainbow_agent.py需要注意两点继承自DQNAgent见 DQNAgent API 文档因此它复用了 DQN 的整套框架epsilon 贪心探索、目标网络target network周期性同步、经验回放训练循环、sess/tf_device等图执行机制。你可以把 RainbowAgent 理解为在 DQN 骨架之上替换了回放缓冲区与训练目标的变体。带有gin.configurable装饰器这是 Dopamine 使用 gin-config 做超参数配置的关键——所有构造参数都可以直接在.gin文件中以RainbowAgent.xxx ...的形式覆盖无需修改任何 Python 代码。这也是下文 gin 配置实战 的基础。构造函数参数详解从sess到summary_writing_frequencyRainbowAgent.__init__的完整签名定义在 rainbow_agent.py其参数大部分直接透传给父类DQNAgent.__init__见 rainbow_agent.py。下表按类别整理全部参数及其默认值、含义参数默认值含义sess必填tf.compat.v1.Session用于执行 ops 的 TensorFlow 会话num_actions必填智能体在任意状态下可采取的动作数量整数observation_shapeNATURE_DQN_OBSERVATION_SHAPE观测形状若为单个整数则假定为 2D 正方形图像observation_dtypeNATURE_DQN_DTYPE观测数据类型连续输入应设为tf.float32stack_sizeNATURE_DQN_STACK_SIZE状态栈中叠加的帧数Atari 默认为 4 帧networklegacy_networks.RainbowNetwork生成网络的 Keras 模型类要求接收(num_actions, num_atoms, support, network_type)四个参数num_atoms51价值分布的分桶atom数量vminNone价值分布支撑集下限为None时自动取-vmax与 C51 一致vmax10.0价值分布支撑集上限gamma0.99折扣因子update_horizon1更新时的回报视野即 n 步更新中的nmin_replay_history20000开始训练价值函数前需要积累的经验转移数量update_period4DQN 两次更新之间的周期agent stepstarget_update_period8000目标网络更新周期agent stepsepsilon_fnlinearly_decaying_epsilon训练期 epsilon 衰减函数接收(decay_period, step, warmup_steps, epsilon)四个参数epsilon_train0.01训练期 epsilon 最终衰减到的值epsilon_eval0.001评估时的 epsilon 值epsilon_decay_period250000epsilon 衰减调度长度replay_schemeprioritized回放内存采样方案prioritized或uniformtf_device/cpu:*Agent 图执行所在的 TensorFlow 设备use_stagingFalse为True时使用 staging area 预取下一个训练批次可提速约 30%optimizerAdamOptimizer(learning_rate0.00025, epsilon0.0003125)用于训练价值函数的优化器summary_writerNone输出训练统计的 SummaryWriter为None时关闭摘要写入summary_writing_frequency500摘要写入频率值越小训练越慢构造阶段的关键内部逻辑构造函数中除了透传参数还完成了三件与分布强化学习直接相关的初始化工作rainbow_agent.py# 防御性地把 vmax 转成 float避免某些工具把 float 转成 int vmax float(vmax) self._num_atoms num_atoms # 若未指定 vmin则类似 C51 设为 -vmax vmin vmin if vmin else -vmax self._support tf.linspace(vmin, vmax, num_atoms) self._replay_scheme replay_scheme self.optimizer optimizer其中self._support是在[vmin, vmax]区间上等间距采样的num_atoms个支撑点例如默认配置下num_atoms51vmax10就是[-10, -9.6, ..., 9.6, 10]共 51 个点。整个价值分布就定义在这组支撑点上网络输出的是每个支撑点对应的概率质量而_build_target_distribution与project_distribution则负责构造并投影目标分布。网络结构RainbowNetwork 如何输出价值分布RainbowAgent默认使用的网络是 legacy_networks.py 中的RainbowNetwork它继承了经典的 Nature DQN 卷积主干但最后一层输出从每个动作一个 Q 值改成了每个动作 × 每个 atom 一个 logitconv1Conv2D32 个 8×8 卷积核步长 4paddingsameReLUconv2Conv2D64 个 4×4 卷积核步长 2ReLUconv3Conv2D64 个 3×3 卷积核步长 1ReLUflattendense1512 维全连接层ReLUdense2输出维度为num_actions * num_atoms的全连接层无激活即每个动作对应num_atoms51个 logits。在前向call中legacy_networks.py输入图像先tf.cast为 float32 并除以 255 归一化然后依次经过上述卷积与全连接层最后通过 reshape 得到(batch_size, num_actions, num_atoms)的输出再经 softmax 得到每个动作上的回报分布、并依据支撑点加权求和还原出期望 Q 值。该网络在RainbowAgent._create_network中被实例化def _create_network(self, name): network self.network( self.num_actions, self._num_atoms, self._support, namename) return networkrainbow_agent.py注意这里传给 Keras 模型的name会被用于创建变量作用域——在线网络与目标网络各自实例化一份从而拥有独立的参数集合这正是 DQN 目标网络机制的实现基础。优先经验回放WrappedPrioritizedReplayBuffer 与采样方案RainbowAgent._build_replay_bufferrainbow_agent.py统一使用prioritized_replay_buffer.WrappedPrioritizedReplayBuffer实现见 dopamine/tf/replay_memory/prioritized_replay_buffer.py并用replay_scheme参数区分两种行为prioritized经验按 TD 误差优先级采样新经验以当前最大记录优先级入队见_store_transition中self._replay.memory.sum_tree.max_recorded_priority的用法rainbow_agent.pyuniform所有经验优先级设为相同的 1.0退化为均匀采样与普通 DQN 回放一致。此外该缓冲区构造时传入了update_horizon与gamma这意味着 n 步回报的累计是在回放缓冲区内部完成的——取出的每条经验都携带了 n 步折现的回报与对应的累计折扣因子cumulative_gamma这正是_build_target_distribution中gamma_with_terminal self.cumulative_gamma * is_terminal_multiplier一行rainbow_agent.py的由来。若传入非法方案既不是uniform也不是prioritized_build_replay_buffer会抛出ValueError: Invalid replay scheme: ...这是源码中明确给出的校验行为。核心算法C51 目标分布构建与投影Eq7_build_target_distribution构造贝尔曼目标分布分布强化学习的关键在于把回报的分布而不是回报的期望作为回归目标。_build_target_distributionrainbow_agent.py按照 Bellemare et al. (2017) 的 C51 算法构造目标分布流程如下计算目标支撑点target_support rewards gamma_with_terminal * tiled_support。即对每个批次样本把支撑点平铺成batch_size × num_atoms的矩阵加上经终止状态修正的折扣奖励若为终止状态is_terminal_multiplier为 0目标支撑点退化为全 0。选出最优动作的概率分布利用目标网络输出的 Q 值期望通过tf.argmax找到每个样本期望值最高的动作再用tf.gather_nd从目标网络的probabilities中取出该动作对应的分布next_probabilitiesrainbow_agent.py。投影到原支撑点由于奖励使目标支撑点发生了平移它们不再落在原始num_atoms个等距支撑点上因此调用project_distribution完成投影rainbow_agent.py。project_distributionEq7 的逐行实现project_distribution(supports, weights, target_support, validate_argsFalse)是模块级的独立函数API 文档源码见 rainbow_agent.py它把一批(support, weights)分布投影到target_support上。源码注释用一组非常直观的示例supports[[0,2,4,6,8],[1,3,4,5,6]]weights 为对应概率target_support[4,5,6,7,8]配合Ex:前缀逐行标注中间张量的形态非常适合精读。其核心步骤对应论文中的公式 (7)由target_support的首尾元素推断v_min、v_max并计算等距间隔delta_z将supports裁剪到[v_min, v_max]区间tf.clip_by_value平铺成batch_size × num_dims × num_dims × 1的张量计算每个原始支撑点与每个目标支撑点的距离得到numerator |clipped_support - z_i|进而得quotient 1 - numerator / delta_z再裁剪到[0, 1]用权重矩阵加权求和inner_prod clipped_quotient * weights沿最后一维tf.reduce_sum得到投影后的分布。函数对输入形状有明确约束rainbow_agent.pysupports与weights形状兼容、supports的每行与target_support形状一致、target_support必须是一维。若validate_argsTrue还会额外断言 target_support 单调递增、等距分布等条件rainbow_agent.py形状不兼容时抛出ValueError。_build_train_op交叉熵损失与优先级更新训练目标建立在投影得到的分布之上rainbow_agent.pytarget_distribution tf.stop_gradient(self._build_target_distribution()) ... chosen_action_logits tf.gather_nd( self._replay_net_outputs.logits, reshaped_actions) loss tf.nn.softmax_cross_entropy_with_logits( labelstarget_distribution, logitschosen_action_logits)即用在线网络对实际采取动作输出的 logits 与投影后的目标分布做softmax 交叉熵——注意这里labels是概率分布而非 one-hot 标签这正是分布 RL 与普通 DQN 在损失函数上的本质区别。当replay_scheme prioritized时还有两个附加步骤rainbow_agent.py重要性采样加权loss_weights 1.0 / tf.sqrt(probs 1e-10)再除以最大值归一化。这里的tf.sqrt等价于论文中的 $\alpha 0.5$、$\beta 0.5$ 组合即实现刻意把 beta 固定为 0.5、去掉 alpha 参数的直接体现更新优先级update_priorities_op self._replay.tf_set_priority(self._replay.indices, tf.sqrt(loss 1e-10))用当前损失加 1e-10 避免 0 优先级导致1.0/0.0NaN回写回放缓冲区中对应样本的优先级实现高误差样本被更频繁采样的闭环。在uniform方案下update_priorities_op退化为tf.no_op()。若配置了summary_writer还会在Losses作用域下写入CrossEntropyLoss标量摘要rainbow_agent.py。gin 配置实战以 rainbow.gin 为例Dopamine 采用 gin-config 管理实验配置。以 rainbow.gin 为例它展示了如何完整地实例化并调优一个 RainbowAgent默认跑 Atari Pong# 导入所需模块必填声明可配置项的来源 import dopamine.tf.agents.rainbow.rainbow_agent import dopamine.discrete_domains.atari_lib import dopamine.discrete_domains.run_experiment import dopamine.tf.replay_memory.prioritized_replay_buffer import gin.tf.external_configurables # Agent 超参数遵循 Hessel et al., 2018仅 sticky_actions 与论文不同 RainbowAgent.num_atoms 51 RainbowAgent.vmax 10. RainbowAgent.gamma 0.99 RainbowAgent.update_horizon 3 # n 步更新的 n RainbowAgent.min_replay_history 20000 # agent steps RainbowAgent.update_period 4 RainbowAgent.target_update_period 8000 # agent steps RainbowAgent.epsilon_train 0.01 RainbowAgent.epsilon_eval 0.001 RainbowAgent.epsilon_decay_period 250000 # agent steps RainbowAgent.replay_scheme prioritized RainbowAgent.tf_device /gpu:0 # 无 GPU 时改用 /cpu:* RainbowAgent.optimizer tf.train.AdamOptimizer() # 注意以下学习率/epsilon 与 C51 配置不同 tf.train.AdamOptimizer.learning_rate 0.0000625 tf.train.AdamOptimizer.epsilon 0.00015 # 环境与运行配置 atari_lib.create_atari_environment.game_name Pong # 以 0.25 概率使用 sticky actionsMachado et al., 2017 的建议 atari_lib.create_atari_environment.sticky_actions True create_agent.agent_name rainbow Runner.num_iterations 200 Runner.training_steps 250000 # agent steps Runner.evaluation_steps 125000 # agent steps Runner.max_steps_per_episode 27000 # agent steps # 回放缓冲区配置 WrappedPrioritizedReplayBuffer.replay_capacity 1000000 WrappedPrioritizedReplayBuffer.batch_size 32要点解读RainbowAgent.optimizer tf.train.AdamOptimizer()的语法表示引用一个由 gin 管理的可配置实例其learning_rate与epsilon随后单独覆盖。默认代码构造器中的 Adamlr0.00025, eps0.0003125会被此配置替换为论文使用的 Adamlr0.0000625, eps0.00015。create_agent.agent_name rainbow告诉 run_experiment.py 通过 gin 查找并创建RainbowAgent——这也解释了为什么RainbowAgent必须标注gin.configurable。对比仓库中的 c51.gin 可以看出C51 变体仅把RainbowAgent.network指向 Categorical 网络、update_horizon保持 1即不做 n 步、replay_scheme改为uniform其余参数基本不变这从配置层面印证了三项组件可以独立开关的设计。仓库还提供了面向不同环境的变体rainbow_cartpole.gin、rainbow_acrobot.gin 以及用于性能分析的 rainbow_profiling.gin。运行方式与 Dopamine 的标准训练入口一致例如python -m dopamine.discrete_domains.train \ --base_dir/tmp/dopamine/runs \ --gin_filesdopamine/tf/agents/rainbow/configs/rainbow.gin测试与验证源码层的可信依据仓库在 tests/dopamine/tf/agents/rainbow/rainbow_agent_test.py 中对上述实现做了系统验证可作为理解实现语义的参考project_distribution有一系列独立的单元测试rainbow_agent_test.py覆盖了支撑点平移、裁剪边界、权重归一化、极端输入等场景直接对应project_distribution的输入输出约定RainbowAgentTest通过构造RainbowAgent(sess, num_actions4)等调用验证 Agent 能正常构建rainbow_agent_test.py并验证了构造参数校验逻辑如非法replay_scheme会抛错。小结何时选择 RainbowAgent如何进一步定制总结起来Dopamine 的RainbowAgent适合作为在 DQN 骨架上叠加三大改进的参考实现与实验起点默认配置即给出与 Hessel et al. (2018) 对齐的超参数且每一项改进都能通过update_horizon、replay_scheme、num_atoms/vmin/vmax独立调节方便做消融实验。若你希望继续深入阅读 rainbow_agent.py 中_build_target_distribution与project_distribution的完整逐行注释内含运行示例张量可彻底理解 Eq7 的数值过程对比 DQNAgent 文档 可理清哪些机制来自父类目标网络、epsilon 探索、训练循环查看 RainbowNetwork 可了解如何自定义网络结构如替换为 Dueling 结构并通过RainbowAgent.network参数注入。需要说明的是本文所述均以当前仓库中的 TensorFlow 1.x 版本实现为准仓库还提供了独立的 JAX 版本见 dopamine/jax/agents/rainbow其 API 与 TF 版存在差异使用前请根据实际安装的依赖选择对应版本。【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考