ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

Dopamine 中 C51 分布投影函数 project_distribution 的完整解析:Eq7 的实现、参数与调用链

Dopamine 中 C51 分布投影函数 project_distribution 的完整解析:Eq7 的实现、参数与调用链 机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载dopamine.agents.rainbow.rainbow_agent.project_distribution是 Dopamine 强化学习框架中分布强化学习Categorical DQN / C51的核心函数它实现 Bellemare et al. (2017) 论文arXiv:1707.06887中方程 (7) 的分布投影操作把一批(support, weights)分布投影到目标支撑集target_support上。本文结合仓库源码从数学原理、TF 与 JAX 两套实现、训练调用链三个层面完整剖析该函数的输入输出约定、逐元素演算过程与工程实现细节帮助读者真正读懂这段不易消化的代码。函数签名与输入输出约定在 TF 实现 中函数签名与 API 文档一致def project_distribution( supports, weights, target_support, validate_argsFalse ):在 JAX 实现 中则省略了validate_args参数JAX 版本默认不做运行时校验。四个输入参数的含义如下参数形状说明supports(batch_size, num_dims)原始分布的支撑点support points即分布定义在哪些取值上weights(batch_size, num_dims)各支撑点上的权重。对 Categorical DQN 而言是概率但并不强制要求是概率不要求求和为 1target_support(num_dims,)投影目标分布的支撑集必须单调递增且等间距Vmin与Vmax分别由该张量的首尾元素推断validate_args标量 bool仅 TF 版本是否对target_support的内容做运行时校验返回形状为(batch_size, num_dims)的张量即投影后的分布。抛出ValueError——当target_support没有维度或supports、weights、target_support形状不兼容时。文档自带的运行示例逐行读懂 Eq7原文档特意给出了一组跑得通的样例输入用来配合源码中的Ex:注释理解supports [[0, 2, 4, 6, 8], # 第 1 个样本的 5 个支撑点 [1, 3, 4, 5, 6]] # 第 2 个样本的 5 个支撑点 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] # 目标支撑集Vmin4, Vmax8这里batch_size 2num_dims 5。投影的本质是把每个样本在[0, 8]区间上的离散分布重新搬到[4, 8]的等间距网格上同时保持质量守恒。这与论文中 Eq7 的符号一一对应delta_z \Delta z相邻支撑点的间距由target_support[1:] - target_support[:-1]的第一个元素得到本例为1clipped_support [\hat{T}_{z_j}]^{V_max}_{V_min}先把支撑点裁剪到[Vmin, Vmax]本例为[[4, 4, 4, 6, 8], [4, 4, 4, 5, 6]]numerator |clipped_support - z_i|每个被投影点与每个目标网格点的绝对距离clipped_quotient [1 - numerator / \Delta z]_0^1距离归一化后裁剪到[0, 1]形成线性插值的分配比例inner_prod clipped_quotient * weights按比例把权重分摊到相邻网格点上最终按\sum_{j0}^{N-1}求和得到投影结果。对第 1 个样本手工验证支撑点0, 2被裁剪到4因此在网格点4处来自0, 2, 4的权重0.1 0.6 0.1 0.8全部落在4上支撑点6恰好落在网格点6上权重0.1支撑点8恰好落在网格点8上权重0.1。最终投影为[0.8, 0.0, 0.1, 0.0, 0.1]与源码注释给出的projection结果完全一致。TF 实现的逐步演算含 Ex: 注释TF 版本在 rainbow_agent.py 中逐步构建计算图关键步骤target_support_deltas target_support[1:] - target_support[:-1] delta_z target_support_deltas[0] # Ex: 1 ... v_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, :] 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]) numerator tf.abs(tiled_support - reshaped_target_support) quotient 1 - (numerator / delta_z) clipped_quotient tf.clip_by_value(quotient, 0, 1) weights weights[:, None, :] inner_prod clipped_quotient * weights projection tf.reduce_sum(inner_prod, 3) projection tf.reshape(projection, [batch_size, num_dims])实现策略是广播式的一次性计算把形状为(batch_size, num_dims)的输入升维到(batch_size, num_dims, num_dims)的距离矩阵每个原始支撑点 × 每个目标网格点利用tf.tile构造出tiled_supportEx 中大小为 2×5×5×5与reshaped_target_support相减取绝对值得到numerator再依次完成归一化、裁剪、乘权重、求和。这种写法虽然内存占用较大(batch_size, num_dims, num_dims)但能在一张计算图中完整表达 Eq7且梯度可以自然回传方便在训练中直接使用。validate_args运行时校验的四条断言TF 版本在validate_argsTrue时会追加四条tf.Assert校验rainbow_agent.pysupports与weights形状一致supports的第二维与target_support形状一致target_support是单维张量target_support严格单调递增target_support_deltas 0target_support等间距所有delta等于delta_z。静态形状检查assert_is_compatible_with、assert_has_rank在构图期完成动态断言则在运行期生效。在 C51/Rainbow 的实际训练路径中该参数默认取False见下文调用链因为target_support是由vmin、vmax、num_atoms三个配置项构造的固定网格保证恒满足上述约束。JAX 实现的函数式写法JAX 版本 语义完全一致但用 JAX 原生算子实现代码更紧凑v_min, v_max target_support[0], target_support[-1] num_dims target_support.shape[0] # N in Eq7 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))注意 JAX 版对delta_z的推导方式不同TF 版从target_support相邻差取值JAX 版直接用(v_max - v_min) / (num_dims - 1)计算——两者在等间距这一前提成立时完全等价。由于 JAX 版输入维度约定为(num_dims,)而非批量的(batch_size, num_dims)批量展开由调用方通过jax.vmap完成见下文因此末尾的jnp.sum(..., -1)配合jnp.squeeze消除单例维度。整体无副作用、可被jax.jit编译便于嵌入可微训练图。在训练流程中的真实调用链TF_build_target_distribution三步构造TF 版 Rainbow/C51 agent 在 rainbow_agent.py 的_build_target_distribution中调用project_distribution该函数注释完整描述了 C51 目标分布的构造流程计算 Bellman 目标支撑集r \gamma Z从回放缓冲区取出rewards将self._support平铺为(batch_size, num_atoms)并用is_terminal_multiplier 1.0 - terminals把终止状态的折扣系数置 0得到target_support rewards gamma_with_terminal * tiled_support选取下一状态最优动作的概率next_qt_argmax tf.argmax(next_target_net_outputs.q_values, axis1)再通过tf.gather_nd取出对应动作的next_probabilities投影回原始支撑集调用project_distribution(target_support, next_probabilities, self._support)即用目标网络的分布做一次回投结果经tf.stop_gradient后作为交叉熵的labels与在线网络所选动作的logits计算softmax_cross_entropy_with_logits损失rainbow_agent.py。JAXtarget_distributionvmap批量展开JAX 版在 rainbow_agent.py 定义了target_distribution用functools.partial(jax.vmap, in_axes(None, 0, 0, 0, None, None))对批量维度自动展开内部同样三步target_support rewards gamma_with_terminal * support→ 按jnp.argmax(q_values)选取next_probabilities→jax.lax.stop_gradient(project_distribution(...))。训练主循环train中直接调用该函数构造targetrainbow_agent.py。在其他 agent 中的复用project_distribution不止服务于基础 Rainbowfull_rainbow完整 Rainbow 实现在构造目标分布时直接复用rainbow_agent.project_distributionSPR agentAtari 100k 基准中的 SPR同样导入并调用该函数。这证明该函数是仓库内所有 C51 式分布强化学习 agent 共享的公共原语。形状约束与常见错误从源码的校验逻辑可以归纳出三条必须满足的形状/取值约束违反即报错或产生错误结果supports与weights形状必须一致均为(batch_size, num_dims)target_support必须是一维、单调递增、等间距JAX 版还要求(num_dims,)单样本形状批量由vmap处理Vmin/Vmax完全由target_support首尾元素决定——若传入的支撑网格不满足等间距TF 版在validate_argsTrue时会触发断言JAX 版则会得到错误的delta_z从而产生数值偏差。实际使用中target_support通常由 agent 的num_atoms、vmin、vmax配置生成如 JAX Rainbow agent 默认num_atoms51, vminNone, vmax10.0见 rainbow_agent.py只要保证(vmax - vmin)能被(num_atoms - 1)整除等间距与单调性即可自动满足。测试与正确性保障仓库为两个实现都配备了单元测试TF 版测试位于 tests/dopamine/tf/agents/rainbow/rainbow_agent_test.py覆盖project_distribution对文档示例输入的计算结果以及validate_args的校验路径JAX 版测试位于 tests/dopamine/jax/agents/rainbow/rainbow_agent_test.py。这些测试直接以supports [[0, 2, 4, 6, 8], ...]这类文档示例作为输入断言投影结果确保 TF 与 JAX 两套实现、以及文档描述三者行为一致是理解该函数行为的最快验证入口。小结project_distribution是 C51 分布强化学习的搬运工它把 Bellman 更新产生的任意分布通过线性插值无损地投影回固定网格支撑集上从而让分布式的价值学习能够与标准的交叉熵损失平滑衔接。理解它的关键是把握三点target_support的等间距网格约定、Vmin/Vmax从网格端点推断、以及裁剪 → 距离归一化 → 裁剪 → 加权求和的 Eq7 四步流水线。无论是阅读 TF 版的广播式实现还是 JAX 版的函数式实现本文给出的逐元素演算都能帮助你快速验证推导。赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐深度解析Dopamine框架中的分布式价值函数Rainbow算法实现指南深度解析Dopamine框架中的分布式价值函数Rainbow算法实现指南 Dopamine是一个专门为强化学习算法快速原型开发而设计的研究框架由Google强化学习机器学习深度学习MyTinySTL中的函数调用invoke函数实现MyTinySTL中的函数调用invoke函数实现 在C编程中函数调用是最基本的操作之一。但当面对函数指针、成员函数指针、仿函数Functor等多种标准库GyroFlow导出慢3步让M1 Mac硬编提速GyroFlow导出慢3步让M1 Mac硬编提速 导出5分钟的4K GoPro素材GyroFlow的进度条在90%之后磨蹭十分钟风扇拉满活动监视器里CP视频处理桌面应用音视频创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表