- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
RainbowNetwork是 Dopamine 项目(gh_mirrors/do/dopamine)中用于分布强化学习(Distributional RL)的核心卷积网络,其官方 API 文档定义为一句话:"The convolutional network used to compute agent's return distributions."(用于计算 Agent 回报分布的卷积网络)。它基于 C51(Bellemare et al., 2017)思想,不再让网络输出单个 Q 值标量,而是为每个动作输出一整个回报分布。本文将以 RainbowNetwork 官方 API 文档 为骨架,结合仓库中 TF 与 JAX 两套实现、C51/Rainbow Agent 及其 Gin 配置、单元测试,完整拆解它的网络结构、输出语义、与 Agent 的集成方式及配置参数,帮助读者理解并复现这一经典分布强化学习网络。
RainbowNetwork 解决的问题:从 Q 值到回报分布
传统 DQN(Nature DQN)网络只输出每个动作的期望 Q 值q_values。而 RainbowNetwork 输出的是:
logits:形状为[num_actions, num_atoms]的未归一化对数概率;probabilities:对logits做 softmax 后得到的每个原子的概率;q_values:用分布支撑support与probabilities加权求和得到的期望值。
其中num_atoms表示把回报值域[vmin, vmax]均匀切分成多少个"桶"(原子)。这样,每个动作的回报不再是单一标量,而是一组离散概率质量,能够刻画回报的不确定性,这也是 C51 系列算法(以及 Rainbow、C51 等变体)取得显著性能提升的数学基础。
在仓库中,这一输出签名被固化成一个命名元组。见 dopamine/discrete_domains/atari_lib.py:
DQNNetworkType = collections.namedtuple('dqn_network', ['q_values']) RainbowNetworkType = collections.namedtuple( 'c51_network', ['q_values', 'logits', 'probabilities'] ) ImplicitQuantileNetworkType = collections.namedtuple( 'iqn_network', ['quantile_values', 'quantiles'] )RainbowNetworkType(内部命名为c51_network,可见其与 C51 算法的渊源)同时被 dopamine/discrete_domains/legacy_networks.py 复用,作为 TF 网络的标准返回类型;JAX 版网络也统一返回该类型,保证了跨后端接口一致。
网络架构:Nature 风格卷积骨干 + 分布输出头
TensorFlow 实现(tf.keras.Model)
当前仓库中 TF 版RainbowNetwork类位于 dopamine/discrete_domains/legacy_networks.py,继承tf.keras.Model。其构造参数为:
| 参数 | 含义 |
|---|---|
num_actions | 动作空间大小,决定输出头第一维 |
num_atoms | 回报分布原子数(桶数) |
support | 分布的支撑,tf.linspace生成,长度等于num_atoms |
name | 网络参数作用域名称 |
逐层结构如下(均在__init__中定义,activation_fn = tf.keras.activations.relu):
conv1:Conv2D,32 个 8×8 卷积核,步长 4,padding='same',ReLU;conv2:Conv2D,64 个 4×4 卷积核,步长 2,padding='same',ReLU;conv3:Conv2D,64 个 3×3 卷积核,步长 1,padding='same',ReLU;flatten:展平;dense1:全连接 512 单元,ReLU;dense2:全连接num_actions * num_atoms单元(无激活函数),作为分布输出头。
所有层使用统一的kernel_initializer:
def kernel_initializer(): return tf.keras.initializers.VarianceScaling( scale=1.0 / np.sqrt(3.0), mode='fan_in', distribution='uniform' )即方差缩放初始化(scale = 1/sqrt(3)、fan_in模式、均匀分布),与 Nature DQN 论文的初始化惯例保持一致。
call方法的前向计算(legacy_networks.py):
x = tf.cast(state, tf.float32) x = x / 255 # 像素归一化到 [0, 1] x = self.conv1(x) # 32 x 8x8, stride 4 x = self.conv2(x) # 64 x 4x4, stride 2 x = self.conv3(x) # 64 x 3x3, stride 1 x = self.flatten(x) x = self.dense1(x) # 512, ReLU x = self.dense2(x) # num_actions * num_atoms logits = tf.reshape(x, [-1, self.num_actions, self.num_atoms]) probabilities = tf.keras.activations.softmax(logits) q_values = tf.reduce_sum(self.support * probabilities, axis=2) return RainbowNetworkType(q_values, logits, probabilities)关键点:最后一层输出被 reshape 成[batch, num_actions, num_atoms],沿最后一个维度做 softmax(对每个动作内部的原子做归一化),再与support逐元素相乘、沿原子轴求和,得到每个动作的期望 Q 值。
JAX 实现(Flax Linen)
JAX 版RainbowNetwork位于 dopamine/jax/networks.py,是flax.linen.Module,声明为@gin.configurable,字段为:
num_actions: intnum_atoms: intinputs_preprocessed: bool = False(若输入已预处理,可跳过归一化)
其前向函数签名是__call__(self, x, support)——注意support作为运行时参数传入,而 TF 版将其保存在构造器中。前向过程与 TF 版一一对应:
if not self.inputs_preprocessed: x = preprocess_atari_inputs(x) # x.astype(jnp.float32) / 255.0 x = nn.Conv(features=32, kernel_size=(8, 8), strides=(4, 4), kernel_init=initializer)(x) x = nn.relu(x) x = nn.Conv(features=64, kernel_size=(4, 4), strides=(2, 2), kernel_init=initializer)(x) x = nn.relu(x) x = nn.Conv(features=64, kernel_size=(3, 3), strides=(1, 1), kernel_init=initializer)(x) x = nn.relu(x) x = x.reshape((-1)) # flatten x = nn.Dense(features=512, kernel_init=initializer)(x) x = nn.relu(x) x = nn.Dense(features=self.num_actions * self.num_atoms, kernel_init=initializer)(x) logits = x.reshape((self.num_actions, self.num_atoms)) probabilities = nn.softmax(logits) q_values = jnp.sum(support * probabilities, axis=1) return atari_lib.RainbowNetworkType(q_values, logits, probabilities)初始化器同样是variance_scaling(scale=1.0/jnp.sqrt(3.0), mode='fan_in', distribution='uniform'),与 TF 版完全对齐。其中preprocess_atari_inputs定义在 dopamine/jax/networks.py:
def preprocess_atari_inputs(x): """Input normalization for Atari 2600 input frames.""" return x.astype(jnp.float32) / 255.0输入规格:Nature DQN 观测
无论 TF 还是 JAX 版本,网络输入都是 Atari 2600 预处理后的观测帧堆叠。相关常量定义在 dopamine/discrete_domains/atari_lib.py:
NATURE_DQN_OBSERVATION_SHAPE = (84, 84) # Size of downscaled Atari 2600 frame. NATURE_DQN_STACK_SIZE = 4 # Number of frames in the state stack.即输入形状为(84, 84, 4)(84×84 灰度帧、4 帧堆叠)。这也是JaxRainbowAgent与RainbowAgent的默认observation_shape与stack_size(见下文)。
与 Agent 的集成:support 如何构造
JAX 版:JaxRainbowAgent
JaxRainbowAgent位于 dopamine/jax/agents/rainbow/rainbow_agent.py,默认参数为:
network=networks.RainbowNetwork, num_atoms=51, vmin=None, vmax=10.0,在__init__中,它先根据vmin/vmax/num_atoms生成支撑向量,再用functools.partial把num_atoms绑定进网络工厂:
vmax = float(vmax) self._num_atoms = num_atoms # If vmin is not specified, set it to -vmax similar to C51. vmin = vmin if vmin else -vmax self._support = jnp.linspace(vmin, vmax, num_atoms) self._replay_scheme = replay_scheme ... network=functools.partial(network, num_atoms=num_atoms),随后在_build_networks_and_optimizer(rainbow_agent.py)中,把support作为初始化与推理的输入传入:
self.online_params = self.network_def.init( rng, x=state, support=self._support )动作选择(select_action,rainbow_agent.py)使用network_def.apply(params, state, support).q_values取 argmax,即在分布输出之上取期望后做贪心选择。
TensorFlow 版:RainbowAgent
TF 版RainbowAgent位于 dopamine/tf/agents/rainbow/rainbow_agent.py,默认network=legacy_networks.RainbowNetwork、num_atoms=51、vmin=None、vmax=10.0,即复用上述legacy_networks.RainbowNetwork类。其网络工厂约定为"输入(num_actions, num_atoms, support),返回(state, network_type)"的生成函数,见 legacy_networks.py 中的rainbow_network工厂:
model = RainbowNetwork(num_actions, num_atoms, support) net = model(state) return network_type(net.q_values, net.logits, net.probabilities)核心参数与 Gin 配置示例
JAX 版 C51/Rainbow 的配置文件位于 dopamine/jax/agents/rainbow/configs 目录。以 rainbow.gin 为例:
import dopamine.jax.replay_memory.replay_buffer import dopamine.jax.replay_memory.samplers import dopamine.discrete_domains.atari_lib import dopamine.discrete_domains.run_experiment JaxRainbowAgent.num_atoms = 51 JaxRainbowAgent.vmax = 10. JaxRainbowAgent.gamma = 0.99 JaxRainbowAgent.update_horizon = 3 JaxRainbowAgent.min_replay_history = 20000 # agent steps JaxRainbowAgent.update_period = 4 JaxRainbowAgent.target_update_period = 8000 # agent stepsC51 变体见 c51.gin,与 Rainbow 的主要区别是update_horizon = 1(不采用 n-step)。参数含义汇总:
| 参数 | 默认值 | 说明 |
|---|---|---|
num_atoms | 51 | 回报分布原子数,决定输出头大小num_actions × num_atoms |
vmax | 10.0 | 支撑上界,support = linspace(-vmax, vmax, num_atoms) |
vmin | None | 支撑下界,为 None 时取-vmax(与 C51 一致) |
gamma | 0.99 | 折扣因子 |
update_horizon | C51 为 1,Rainbow 为 3 | n 步回报的步数 |
min_replay_history | 20000 | 开始训练前需要积累的转移样本数 |
update_period | 4 | 每多少个 agent 步做一次梯度更新 |
target_update_period | 8000 | 目标网络同步周期 |
实际训练中,vmin/vmax的选取需与游戏回报量级匹配;在经典控制(CartPole、Acrobot、LunarLander、MountainCar)的配置中(如 c51_cartpole.gin),num_atoms常被调高到 201、vmax到 100,并改用networks.ClassicControlRainbowNetwork(同为分布输出,但用 MLP 骨干并支持min_vals/max_vals归一化,见 dopamine/jax/networks.py)。
训练侧:C51 目标分布与分布投影
logits/probabilities不仅用于动作选择,还直接参与损失计算。在 dopamine/jax/agents/rainbow/rainbow_agent.py 的target_distribution中,仓库按 C51 方式构造目标分布:
- 目标支撑为
rewards + gamma_with_terminal * support(终止状态将目标支撑置 0); - 用目标网络对下一状态打分,取
q_values最大的动作对应的probabilities作为下一步分布; - 最后调用
project_distribution把目标分布投影回原始支撑。
project_distribution(rainbow_agent.py)在注释中明确标注其实现基于 Bellemare et al. (2017) 论文的公式 (7):先对目标支撑做裁剪(clip 到[vmin, vmax]),再按三角形核(triangular kernel)把概率质量分配到相邻两个原子。网络在线分支的logits经 softmax 后与投影目标分布计算交叉熵损失,从而端到端学习回报分布。
测试验证:输出形状与训练路径
仓库测试对RainbowNetwork的输出契约有直接验证:
- tests/dopamine/jax/networks_test.py:在
testOutputShape中把networks.RainbowNetwork(以及FullRainbowNetwork、ClassicControlRainbowNetwork、QuantileNetwork)纳入参数化测试,以num_actions、num_atoms实例化网络并断言输出张量形状符合RainbowNetworkType契约。 - tests/dopamine/jax/agents/rainbow/rainbow_agent_test.py 与 tests/dopamine/tf/agents/rainbow/rainbow_agent_test.py:分别用 Mock 网络复刻"
logitsreshape → softmax →q_values = sum(support * probabilities)"的输出逻辑,验证 Agent 在给定分布输出下能正确选择使 Q 值最大的动作。
这两类测试共同印证了:无论后端是 TF 还是 JAX,RainbowNetwork的"三输出"(q_values、logits、probabilities)契约都是 Agent 训练与决策的稳定接口。
与其他 Atari 网络的定位对比
在同级 API 文档中,Dopamine 还提供了两个姊妹网络,三者的差异正体现在输出签名上(定义见 atari_lib.py):
| 网络 | 输出类型 | 输出内容 | 对应算法 |
|---|---|---|---|
| NatureDQNNetwork | DQNNetworkType | 仅q_values标量 | DQN |
RainbowNetwork | RainbowNetworkType | q_values、logits、probabilities | C51 / Rainbow 分布强化学习 |
| ImplicitQuantileNetwork | ImplicitQuantileNetworkType | 分位数取值与分位数位置 | IQN |
RainbowNetwork是所有分布值类(distributional value-based)算法在 Atari 上的默认骨干:JaxRainbowAgent、TFRainbowAgent以及FullRainbowNetwork(dopamine/jax/networks.py,支持开关 distributional 与 noisy nets)都直接或间接复用它定义的分布输出范式。MinAtar 环境的MinatarRainbowNetwork(dopamine/labs/environments/minatar/minatar_env.py)也采用同样的"logits→probabilities→q_values"三输出结构,说明该网络模式在轻量环境上的通用性。
小结
从 API 文档的一句话定义出发,可以完整梳理出RainbowNetwork在 Dopamine 中的全貌:一个 Nature 风格的卷积骨干(3 层卷积 + 512 全连接)加一个num_actions × num_atoms的分布输出头;输入为 84×84×4 的 Atari 帧堆叠;输出q_values/logits/probabilities三要素,支撑值由num_atoms/vmin/vmax决定;TF 与 JAX 两套实现在结构、初始化与归一化策略上完全对齐。它是 C51/Rainbow 系列 Agent 计算回报分布的核心组件,也是理解 Dopamine 分布强化学习体系的最佳切入点——掌握它的输入输出契约与配置参数,即可在此基础上扩展自定义的分布值网络或新的分布强化学习算法。
- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
相关推荐
Wagtail 0.8.5 版本发布解读:11 项 Bug 修复及其在当前代码库中的演变
Wagtail 0.8.5 版本发布解读:11 项 Bug 修复及其在当前代码库中的演变 Wagtail 0.8.5 是 2015 年 2 月 17 日发布的维
机器学习深度学习Lotus生成式与判别式模型对比:如何选择最适合的方案
Lotus生成式与判别式模型对比:如何选择最适合的方案 Lotus是一个基于扩散技术的视觉基础模型,专注于高质量密集预测任务。本文将深入对比Lotus中的生成式
深度卷积神经网络AlexNet解析与实现
深度卷积神经网络AlexNet解析与实现 引言 在深度学习发展历程中,AlexNet是一个里程碑式的模型。2012年,AlexNet在ImageNet图像识别挑
人工智能深度学习机器学习教程
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考