尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Dopamine 分布投影详解:深入解析 `project_distribution` 与 C51 算法 Eq7 的实现

Dopamine 分布投影详解:深入解析 `project_distribution` 与 C51 算法 Eq7 的实现 Dopamine 分布投影详解深入解析project_distribution与 C51 算法 Eq7 的实现【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine导读本文以 Dopamine 框架中dopamine.tf.agents.rainbow.rainbow_agent.project_distribution函数为对象完整讲解分布强化学习Distributional RL中分布投影distribution projection这一核心操作它基于 C51 论文Bellemare et al., 2017中的 Eq7 公式将一个支持点集合上的离散概率分布搬运到另一组支持点上。读完本文你将理解该函数的四个参数、批处理计算流程、源码中每一行 TensorFlow 运算的含义、可选的参数校验机制以及它如何被 Rainbow 智能体的目标分布构建_build_target_distribution所调用并掌握 JAX 版本实现的差异与对应测试用例。背景为什么需要分布投影Dopamine 的 TF 版 Rainbow 智能体dopamine/tf/agents/rainbow/rainbow_agent.py是一个简化版 Rainbow它从原始 Rainbow 论文Hessel et al., 2018中实现了三个对 Atari 游戏性能影响最大的组件n-step 更新update_horizon优先经验回放prioritized replayreplay_schemeprioritized分布强化学习distributional RL即 C51 风格的价值分布。与普通 DQN 直接回归 Q 值标量不同分布强化学习让网络输出一个离散的回报分布Dopamine 默认用num_atoms51个均匀间隔的支持点support覆盖[vmin, vmax]区间默认vmin-10.0、vmax10.0见 rainbow_agent.py 中构造函数默认参数以及第 125 行self._support tf.linspace(vmin, vmax, num_atoms)。训练时我们需要构造目标分布用贝尔曼算子r γ·Z生成下一状态的价值分布但该分布的支持点经过奖励缩放和折扣后与网络自身的固定支持点不再对齐。此时就需要project_distribution把 (support, weights) 表示的分布投影回目标支持点上——这正是 C51 论文的 Eq7 所定义的运算。该函数在源码注释中明确说明rainbow_agent.pyProjects a batch of (support, weights) onto target_support. Based on equation (7) in (Bellemare et al., 2017).函数签名与参数详解函数签名定义在 dopamine/tf/agents/rainbow/rainbow_agent.pydef project_distribution(supports, weights, target_support, validate_argsFalse):参数类型与形状含义supportsTensor形状(batch_size, num_dims)定义分布的支持点集合每个样本一行。weightsTensor形状(batch_size, num_dims)原始支持点上的权重。对 CategoricalDQN 智能体而言这些权重通常是概率但并不强制要求归一化。target_supportTensor形状(num_dims)投影目标分布的支持点。必须单调递增Vmin/Vmax分别取该张量的首元素和末元素各点之间必须等间距。validate_argsbool默认False是否在运行时通过tf.Assert校验target_support的内容单调性、等间距、形状兼容性。返回值形状为(batch_size, num_dims)的 Tensor即一批(support, weights)投影到target_support上的结果。可能抛出的异常当target_support没有维度标量或supports、weights、target_support的形状不兼容时抛出ValueError。一个贯穿全文的运行示例源码与 API 文档都使用了同一组示例输入来讲解算法见 rainbow_agent.pysupports [[0, 2, 4, 6, 8], [1, 3, 4, 5, 6]] weights [[0.1, 0.6, 0.1, 0.1, 0.1], [0.1, 0.2, 0.5, 0.1, 0.1]] target_support [4, 5, 6, 7, 8]其中batch_size 2num_dims 5v_min 4v_max 8delta_z 1。下文每一步中间结果都以此例为基础。源码逐行拆解Eq7 的 TensorFlow 实现project_distribution的实现位于 rainbow_agent.py下面按执行顺序逐步解读。1. 准备阶段提取 delta_z 与静态形状校验target_support_deltas target_support[1:] - target_support[:-1] # delta_z \Delta z in Eq7. delta_z target_support_deltas[0] validate_deps [] supports.shape.assert_is_compatible_with(weights.shape) supports[0].shape.assert_is_compatible_with(target_support.shape) target_support.shape.assert_has_rank(1)delta_z是相邻支持点的间距对应 Eq7 中的Δz示例中为1。三条assert_is_compatible_with是静态形状检查supports与weights形状必须兼容supports的第一行必须与target_support形状兼容target_support必须是一维向量。2. 可选校验validate_args 开启时的运行时断言当validate_argsTrue时会追加 5 个tf.Assertrainbow_agent.pysupports与weights形状完全相同supports第二维与target_support长度相同target_support只有一维target_support严格单调递增target_support_deltas 0target_support各点等间距所有target_support_deltas都等于delta_z。这些断言会通过tf.control_dependencies挂到计算图上运行时若违反会抛出tf.errors.InvalidArgumentError断言失败。3. 裁剪支持点clipped_supportv_min, v_max target_support[0], target_support[-1] # Ex: 4, 8 batch_size tf.shape(supports)[0] # Ex: 2 num_dims tf.shape(target_support)[0] # Ex: 5 clipped_support tf.clip_by_value(supports, v_min, v_max)[:, None, :]对应 Eq7 中的[T̂ z_j]^{V_max}_{V_min}把支持点裁剪到[v_min, v_max]区间内然后增加一个维度便于后续广播。示例输出形状(batch_size, 1, num_dims)clipped_support [[[ 4. 4. 4. 6. 8.]], [[ 4. 4. 4. 5. 6.]]]4. 广播构造每个目标点 vs 每个原支持点的距离矩阵tiled_support tf.tile([clipped_support], [1, 1, num_dims, 1]) reshaped_target_support tf.tile(target_support[:, None], [batch_size, 1]) reshaped_target_support tf.reshape( reshaped_target_support, [batch_size, num_dims, 1] )tiled_support把裁剪后的支持点复制num_dims份形状变为(1, batch_size, num_dims, num_dims)reshaped_target_support把目标支持点转成(batch_size, num_dims, 1)。二者广播相减后每个(b, i, j)位置都代表第 i 个目标点与第 j 个原始支持点的距离这是实现 Eq7 中|T̂ z_j − z_i|的关键。5. 计算线性插值系数numerator / quotient / clipped_quotientnumerator tf.abs(tiled_support - reshaped_target_support) quotient 1 - (numerator / delta_z) clipped_quotient tf.clip_by_value(quotient, 0, 1)numerator即|clipped_support − z_i|示例中第一个样本的第 0 行[0, 0, 0, 2, 4]表示目标点 4 到原支持点[4,4,4,6,8]的距离quotient是1 − numerator/Δzclipped_quotient把商裁剪到[0, 1]对应 Eq7 中的[1 − |T̂ z_j − z_i|/Δz]_0^1。直观理解这个值就是原支持点 j 的权重按线性距离分配给目标点 i 的比例——距离越近分配越多超过一个Δz则为 0。6. 加权求和inner_prod → projectionweights weights[:, None, :] # (batch_size, 1, num_dims) inner_prod clipped_quotient * weights # 逐元素乘 projection tf.reduce_sum(inner_prod, 3) # 对原支持点维求和 projection tf.reshape(projection, [batch_size, num_dims])inner_prod是 Eq7 中的Σ_j clipped_quotient · p_j(x, π(x))即每个目标点接收到的来自所有原支持点的加权贡献最后沿原支持点维度求和并 reshape 回(batch_size, num_dims)。示例最终输出与测试用例 rainbow_agent_test.py 中testExampleFromCodeComments的期望完全一致projection [[0.8, 0.0, 0.1, 0.0, 0.1], [0.8, 0.1, 0.1, 0.0, 0.0]]以第一行为例权重[0.1, 0.6, 0.1, 0.1, 0.1]分布在支持点[0,2,4,6,8]上其中0.6落在点 2 上距目标点 4 的距离为 2恰好一个Δz的整数倍于是按线性插值规则 0.6 全部投影到目标点 4点 6 上的 0.1 投影到目标点 6点 8 上的 0.1 投影到目标点 8而点 0 上的 0.1 因为超出[v_min, v_max]范围被裁剪后全部落向最近的目标点 4。最终[0.10.60.1, 0, 0.1, 0, 0.1] [0.8, 0, 0.1, 0, 0.1]且总和保持为 1。在 RainbowAgent 中的调用链目标分布如何构建project_distribution是分布 RL 训练回路中的关键一环。在 TF 版RainbowAgent中它被_build_target_distribution调用rainbow_agent.py完整流程为从回放缓冲取rewards将支持点tiled_support平铺到整个 batch计算带终止标志的折扣因子gamma_with_terminal cumulative_gamma * (1 - terminal)从而得到贝尔曼目标支持点target_support rewards gamma_with_terminal * tiled_support终止状态下该值为 0用目标网络输出挑选使期望值最大的动作next_qt_argmax取出对应的下一状态概率next_probabilities调用project_distribution(target_support, next_probabilities, self._support)把贝尔曼目标分布投影回原始支持点。随后在_build_train_oprainbow_agent.py中该目标分布经tf.stop_gradient后作为 softmax 交叉熵的标签与在线网络输出的 logits 计算损失在 prioritized 方案下损失还叠加1/sqrt(probs 1e-10)的重要性采样权重并回写优先级sqrt(loss 1e-10)。值得一提的是该函数并非 TF 智能体专用JAX 版 Rainbowdopamine/jax/agents/rainbow/rainbow_agent.py、JAX 版 Full Rainbowdopamine/jax/agents/full_rainbow/full_rainbow_agent.py以及 Atari 100k 的 SPR 智能体dopamine/labs/atari_100k/spr_agent.py都实现了同名同语义的投影函数说明该运算在分布 RL 家族中是通用基础设施。JAX 版本的等价实现JAX 版project_distribution在 dopamine/jax/agents/rainbow/rainbow_agent.py 中实现逻辑完全等价但更简洁省略了校验与形状广播的显式中间张量v_min, v_max target_support[0], target_support[-1] num_dims target_support.shape[0] delta_z (v_max - v_min) / (num_dims - 1) clipped_support jnp.clip(supports, v_min, v_max) numerator jnp.abs(clipped_support - target_support[:, None]) quotient 1 - (numerator / delta_z) clipped_quotient jnp.clip(quotient, 0, 1) inner_prod clipped_quotient * weights return jnp.squeeze(jnp.sum(inner_prod, -1))两处实现的核心差异delta_z的求法不同TF 版取target_support相邻差分的首元素JAX 版直接按等间距假设计算(v_max − v_min) / (num_dims − 1)。二者在支持点等间距时结果一致。缺少validate_argsJAX 版没有参数校验开关且 JAX 的静态形状检查也更宽松因此调用方需自行保证输入满足单调递增、等间距的前提。测试验证行为由测试用例锁定project_distribution的正确性由 tests/dopamine/tf/agents/rainbow/rainbow_agent_test.py 中的一整套用例覆盖主要分为两类形状与参数校验类均断言抛出ValueError或运行时断言失败testInconsistentSupportsAndWeightssupports与weights第二维不一致testInconsistentSupportsAndTargetSupportsupports与target_support维度不匹配testZeroDimensionalTargetSupporttarget_support为标量testMultiDimensionalTargetSupporttarget_support为二维张量testProjectWithNonMonotonicTargetSupporttarget_support非单调递增如[8, 7, 6, 5, 4]testProjectNewSupportHasInconsistentDeltasktarget_support不等间距如[3, 4, 6, 7, 8]。数值正确性类对投影结果做assertAllClosetestProjectSingleIdenticalDistribution支持点不变时投影即恒等testProjectSingleDifferentDistribution、testProjectFromNonMonotonicSupport支持点平移/乱序时权重按距离重新分配testExampleFromCodeComments即上文示例期望输出[[0.8, 0, 0.1, 0, 0.1], [0.8, 0.1, 0.1, 0, 0]]testProjectBatchOfDifferentDistributions/testProjectBatchOfDifferentDistributionsWithLargerDelta验证 batch 处理与更大Δz支持点间隔为 4下的分配正确性testUsingPlaceholders验证通过tf.placeholder动态喂数据时的行为。这些测试同时印证了两点工程细节其一校验断言在validate_argsTrue时通过tf.Assert实现运行时违反会抛tf.errors.InvalidArgumentError其二投影结果逐行求和保持为 1权重为概率时即该变换是保质量的mass-preserving。使用注意事项保证 target_support 等间距且单调递增delta_z直接取相邻差分的首元素若后续点间距不一致投影结果将不满足 Eq7 的定义运行时断言仅在validate_argsTrue时触发生产环境建议自行保证。weights 不必是概率文档明确说明虽然 CategoricalDQN 中权重是概率但函数并不要求归一化若传入非归一化权重输出只是按相同规则线性分配的加权结果。越界支持点会被裁剪所有超出[v_min, v_max]的原始支持点都会被裁剪到边界对应质量会被集中到最近的目标点如示例中支持点 0 的质量全部流向目标点 4。选择正确的vmin/vmax它们决定价值分布的覆盖范围在RainbowAgent中通过num_atoms、vmin、vmax构造参数控制rainbow_agent.py默认num_atoms51、vmin-vmax-10.0、vmax10.0与 C51 论文保持一致。批处理形状约定supports、weights必须是(batch_size, num_dims)target_support必须是(num_dims)三者任何一处维度不匹配都会在构图期静态检查或运行期断言被捕获。小结project_distribution是 Dopamine 中分布强化学习算法C51 / Rainbow / Full Rainbow / SPR共用的质量搬运工具它以 C51 论文 Eq7 为数学基础通过裁剪 → 距离矩阵 → 线性插值 → 加权求和四步把贝尔曼算子作用后的分布无损地投影回网络输出支持点上。理解它就理解了分布 RL 训练中目标分布构造的核心环节也能读懂 RainbowAgent 的训练回路与 JAX 版实现之间的对应关系。【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表