- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
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_args=False ):在 JAX 实现 中则省略了validate_args参数(JAX 版本默认不做运行时校验)。四个输入参数的含义如下:
| 参数 | 形状 | 说明 |
|---|---|---|
supports | (batch_size, num_dims) | 原始分布的支撑点(support points),即分布定义在哪些取值上 |
weights | (batch_size, num_dims) | 各支撑点上的权重。对 Categorical DQN 而言是概率,但并不强制要求是概率(不要求求和为 1) |
target_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] # 目标支撑集,Vmin=4, Vmax=8这里batch_size = 2,num_dims = 5。投影的本质是:把每个样本在[0, 8]区间上的离散分布,重新"搬到"[4, 8]的等间距网格上,同时保持质量守恒。这与论文中 Eq7 的符号一一对应:
delta_z = \Delta z:相邻支撑点的间距,由target_support[1:] - target_support[:-1]的第一个元素得到,本例为1;clipped_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_{j=0}^{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_support(Ex 中大小为 2×5×5×5)与reshaped_target_support,相减取绝对值得到numerator,再依次完成归一化、裁剪、乘权重、求和。这种写法虽然内存占用较大((batch_size, num_dims, num_dims)),但能在一张计算图中完整表达 Eq7,且梯度可以自然回传,方便在训练中直接使用。
validate_args:运行时校验的四条断言
TF 版本在validate_args=True时会追加四条tf.Assert校验(rainbow_agent.py):
supports与weights形状一致;supports的第二维与target_support形状一致;target_support是单维张量;target_support严格单调递增(target_support_deltas > 0);target_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, axis=1),再通过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)。
JAX:target_distribution+vmap批量展开
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中直接调用该函数构造target(rainbow_agent.py)。
在其他 agent 中的复用
project_distribution不止服务于基础 Rainbow:
- full_rainbow(完整 Rainbow 实现)在构造目标分布时直接复用
rainbow_agent.project_distribution; - SPR agent(Atari 100k 基准中的 SPR)同样导入并调用该函数。
这证明该函数是仓库内所有 C51 式分布强化学习 agent 共享的公共原语。
形状约束与常见错误
从源码的校验逻辑可以归纳出三条必须满足的形状/取值约束,违反即报错或产生错误结果:
supports与weights形状必须一致,均为(batch_size, num_dims);target_support必须是一维、单调递增、等间距(JAX 版还要求(num_dims,)单样本形状,批量由vmap处理);Vmin/Vmax完全由target_support首尾元素决定——若传入的支撑网格不满足等间距,TF 版在validate_args=True时会触发断言,JAX 版则会得到错误的delta_z从而产生数值偏差。
实际使用中,target_support通常由 agent 的num_atoms、vmin、vmax配置生成(如 JAX Rainbow agent 默认num_atoms=51, vmin=None, vmax=10.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 版的函数式实现,本文给出的逐元素演算都能帮助你快速验证推导。
- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
相关推荐
深度解析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),仅供参考