首页/新闻资讯/正文详情

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

发布时间:2026/9/24 11:13:23 来源:云帆数科 栏目:资讯中心
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),仅供参考

相关推荐

A/B测试工具全景图:主流平台一览
A/B测试工具全景图:主流平台一览

A/B测试工具分企业级套件、产品实验平台、开源自建和内置工具四类;选型先定预算与流量规模,再看统计引擎与数据底座,别被名气带偏。你在做A/B测试工具调研,会发现名字一大堆:Optimizely、VWO、Adobe Target&#xff0c… · 2026/9/24 11:13:11

Linux GRUB2 配置与启动故障修复完全指南
Linux GRUB2 配置与启动故障修复完全指南

系统起不来时,屏幕上的报错往往只给出一句零散线索:Operating System not found、Give root password for maintenance。 本文按故障位置分成四类来讲:/etc/fstab 写错导致的挂载失败、GRUB2 配置的修改与加固、GRUB2 引导程序与引导文件损坏… · 2026/9/24 11:13:05

结论与贡献也要双版本烟测:三列表 + A/B 验收表
结论与贡献也要双版本烟测:三列表 + A/B 验收表

千笔-AIWritePaper https://www.aiwritepaper.com 结论与贡献最容易出现两种假完成:一是「贡献条很多」,但主张编号、证据锚点与可口述句对不上;二是结论写得很满,却从未留下「只会堆口号不核证据」的失败对照。claim—evidence—… · 2026/9/24 11:12:52

WCH-LINK与DAP-LINK驱动安装失败?Windows 10/11完整解决指南
WCH-LINK与DAP-LINK驱动安装失败?Windows 10/11完整解决指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 11:44:23

DP83848以太网PHY硬件设计与驱动调试实战指南
DP83848以太网PHY硬件设计与驱动调试实战指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 11:44:10

UDS流控三剑客:BS、STmin与FC帧的硬核解析
UDS流控三剑客:BS、STmin与FC帧的硬核解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 11:44:10

GDPR十年实录:从合规工具到隐私技术协议的演进
GDPR十年实录:从合规工具到隐私技术协议的演进

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 11:44:04

百视通盒子BesTV R3300-L刷机实战:从驱动识别到固件烧录全指南
百视通盒子BesTV R3300-L刷机实战:从驱动识别到固件烧录全指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 11:44:03

800V车载PFC电感选型实战:从参数陷阱到车规量产
800V车载PFC电感选型实战:从参数陷阱到车规量产

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 11:43:57

基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程
基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程

简介:这是一套面向计算机、人工智能、自动化等专业学生与教师的毕业设计级项目资源,围绕YOLOv8实现渔船作业监控系统,可用于毕设、课程设计、大作业或项目立项演示。压缩包共97个文件,约24.21MB,以70个Python源码文件为… · 2026/9/24 0:00:13

1D-CNN时间序列建模实战:从Conv1d原理到工业落地
1D-CNN时间序列建模实战:从Conv1d原理到工业落地

简介:面向时间序列数据建模的一维卷积神经网络完整实现,适合深度学习入门者及需要快速验证时序模型的研究者,能够从音频、文本、传感器或股价等序列中挖掘局部特征与时间依赖。压缩包体积很小,只有3KB,内含3个Python脚… · 2026/9/24 0:00:26

柔软的L:汉语语流中被忽视的舌肌张力控制
柔软的L:汉语语流中被忽视的舌肌张力控制

1. 这个“L”不是字母表里的L,而是舌尖上的L最近在几个方言群和语音教学社群里,反复看到有人发一句:“也说字母L:柔软的长舌”。初看以为是英语发音课笔记,点开才发现全是方言爱好者、播音系学生、语言康复师甚至戏曲演… · 2026/9/24 0:00:44

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码