news 2026/9/24 14:50:06

Dopamine legacy_networks 模块解析:TensorFlow 离散域网络架构与实战配置指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Dopamine legacy_networks 模块解析:TensorFlow 离散域网络架构与实战配置指南
  • 机器学习
  • 深度学习

【免费下载链接】dopamine

Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.

项目地址:https://gitcode.com/gh_mirrors/do/dopamine
点击查看免费下载

导读

dopamine.discrete_domains.legacy_networks是 Dopamine 强化学习框架中负责构建离散动作域(discrete domains)价值网络的核心模块,涵盖了 Atari 2600 上的 Nature DQN、Rainbow(C51)与 Implicit Quantile(IQN)三大卷积网络,以及 Cartpole、Acrobot、LunarLander、MountainCar 等 Gym 经典环境的 DQN / Rainbow / Fourier 基网络。读完本文,你将掌握这些"Legacy(Keras 化之前的经典架构)"网络的设计意图、每个函数与类的输入输出契约、在.gin配置文件中如何引用与参数化它们,以及如何借助maybe_transform_variable_names完成新旧 checkpoint 的变量名映射。

模块定位:Legacy 网络在 Dopamine 中的角色

legacy_networks位于 dopamine/discrete_domains/legacy_networks.py,官方 API 文档将其定义为"Legacy (pre-Keras) network architectures",即Keras 化改造之前就存在的网络架构。虽然命名带 "Legacy",但模块中的网络全部以tf.keras.Model/tf.keras.layers.Layer形式重新实现,是 Dopamine TF 分支中 DQN、Rainbow、IQN 三类智能体的默认价值网络,同时也是 Gym 离散域示例配置的默认网络来源。

从源码结构看,模块分为四大块:

  1. 常量与类型别名NATURE_DQN_OBSERVATION_SHAPENATURE_DQN_DTYPENATURE_DQN_STACK_SIZE三个 Atari 观测常量,以及从 atari_lib.py 复用的三个 namedtuple 类型DQNNetworkTypeRainbowNetworkTypeImplicitQuantileNetworkType
  2. Keras Atari 卷积网络NatureDQNNetworkRainbowNetworkImplicitQuantileNetwork三个tf.keras.Model类;
  3. 通用 Gym 网络BasicDiscreteDomainNetwork全连接子网、FourierBasis特征生成器,以及 Cartpole / Acrobot / LunarLander / MountainCar 的 DQN、Fourier、Rainbow 网络;
  4. checkpoint 兼容工具maybe_transform_variable_names

三个输出类型契约(namedtuple)

网络统一通过 namedtuple 返回输出,定义于 atari_lib.py:

类型字段适用网络
DQNNetworkTypeq_valuesDQN 风格网络
RainbowNetworkTypeq_values, logits, probabilitiesRainbow / C51 风格网络
ImplicitQuantileNetworkTypequantile_values, quantilesIQN 网络

nature_dqn_networkrainbow_networkimplicit_quantile_network三个顶层函数是这些 namedtuple 与网络类之间的薄包装层,它们接收network_type参数(即上述某个 namedtuple 类型)并在前向传播后把结果按对应字段重新构造返回,从而兼容旧式tf.slim风格的函数式调用约定。

Atari 卷积网络三件套

NatureDQNNetwork:经典 DQN 卷积网络

NatureDQNNetwork实现了 Nature 论文(Mnih et al., 2015) 的经典结构,用于计算智能体的 Q 值,构造函数签名与 atari_lib.py 中的常量一一对应:

NatureDQNNetwork(num_actions, name=None)
  • num_actions:int,动作数,决定输出层维度;
  • name:str,网络参数的作用域名称。

网络结构(legacy_networks.py):

类型参数说明
conv1Conv2D32 个 8×8 卷积核,stride 4,padding=same输入为堆叠的 Atari 帧(默认 84×84×4)
conv2Conv2D64 个 4×4 卷积核,stride 2,padding=same
conv3Conv2D64 个 3×3 卷积核,stride 1,padding=same
flattenFlatten展平特征图
dense1Dense512 单元,ReLU
dense2Densenum_actions单元,无激活输出 Q 值

call()中,输入先被tf.cast转为tf.float32并除以 255 归一化到[0, 1],最后返回DQNNetworkType(self.dense2(x))。源码为卷积层统一命名为'Conv'、全连接层命名为'fully_connected',注释明确说明这是为了"使变量名与 tf.slim 变量名/checkpoint 更接近",方便复用旧 checkpoint。

RainbowNetwork:C51 分布价值网络

RainbowNetwork把 Q 值建模为价值分布而非标量,输出层维度变为num_actions * num_atoms,构造函数为:

RainbowNetwork(num_actions, num_atoms, support, name=None)
  • num_atoms:int,价值分布的桶(bucket)数量;
  • supporttf.linspace,Q 值分布的支撑点向量。

关键差异(legacy_networks.py):

  1. 所有层使用VarianceScaling(scale=1.0 / np.sqrt(3.0), mode='fan_in', distribution='uniform')初始化器(C51 论文推荐的小方差初始化);
  2. 前向计算中,dense2输出被 reshape 为[-1, num_actions, num_atoms]得到 logits,经 softmax 得概率分布,再与支撑点做加权求和还原 Q 值:
    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)

ImplicitQuantileNetwork:隐分位数网络

ImplicitQuantileNetwork实现 Dabney et al. (2018) 的 IQN,核心思想是用随机采样的分位数驱动网络输出,构造函数:

ImplicitQuantileNetwork(num_actions, quantile_embedding_dim, name=None)
  • quantile_embedding_dim:int,分位数输入的嵌入维度。

分位数流计算(legacy_networks.py):

  1. 卷积三件套 + flatten 提取状态特征state_net_tiled(沿 batch 维复制num_quantiles份);
  2. [0, 1)均匀采样num_quantiles个分位数,复制到quantile_embedding_dim维;
  3. cos(i * π * tau)(i=1..embedding_dim)做分位数嵌入,这是论文中的分位数特征映射;
  4. 通过延迟创建的dense_quantile层(之所以延迟,是因为其输出单元数依赖输入特征长度,只能在首次前向时确定)与状态特征做逐元素相乘,再经dense1dense2输出quantile_values

call()的签名为call(self, state, num_quantiles)num_quantiles表示每次前向采样的分位数个数。三个 Atari 网络对应的函数式包装分别为nature_dqn_networkrainbow_networkimplicit_quantile_network(见 legacy_networks.py),它们可被.gin直接引用为@legacy_networks.xxx

通用 Gym 网络家族

除 Atari 外,模块为低维 Gym 环境提供了三套网络族,全部通过@gin.configurable暴露给配置系统。

BasicDiscreteDomainNetwork:带归一化的全连接子网

BasicDiscreteDomainNetwork(legacy_networks.py)是 Cartpole / Acrobot / MountainCar 等网络的公共内层模块,定义为tf.keras.layers.Layer

BasicDiscreteDomainNetwork(min_vals, max_vals, num_actions, num_atoms=None, name=None, activation_fn=tf.keras.activations.relu)
  • min_vals/max_vals:与state同形状的最值向量,用于输入归一化;传None则跳过归一化(如 LunarLander);
  • num_atoms=None:为 None 时构造 DQN 风格网络(输出num_actions),否则构造 Rainbow 风格网络(输出num_actions * num_atoms)。

前向时输入被归一化到[-1, 1]

x -= self.min_vals x /= self.max_vals - self.min_vals x = 2.0 * x - 1.0 # Rescale in range [-1, 1]

各环境的归一化边界定义在 gym_lib.py:

环境min_valsmax_vals
Cartpole[-2.4, -5.0, -π/12, -2π][2.4, 5.0, π/12, 2π]
Acrobot[-1, -1, -1, -1, -5, -5][1, 1, 1, 1, 5, 5]
MountainCar[-1.2, -0.07][0.6, 0.07]

FourierDQNNetwork 与 FourierBasis:线性函数逼近

FourierBasis类(legacy_networks.py)实现了 Konidaris, Osentoski & Thomas (2011) 的"Value Function Approximation in Reinforcement Learning using the Fourier Basis"。它使用只含余弦项的基函数(因此系数数量仅为完整傅里叶逼近的一半),通过itertools.product(range(order+1), repeat=nvars)生成所有阶数组合的乘数向量,并剔除第一个全零项(对应常数偏置)。特征计算为:

def compute_features(self, features): scaled = self.scale(features) # 缩放到 [0, 1] return tf.cos(np.pi * tf.matmul(scaled, self.multipliers, transpose_b=True))

FourierDQNNetwork组合 Fourier 特征与无偏置线性层Dense(num_actions, use_bias=False),函数签名:

FourierDQNNetwork(min_vals, max_vals, num_actions, fourier_basis_order=3, name=None)

由于FourierBasis需要输入特征维度才能构造,feature_generator在首次前向时才延迟创建(与dense_quantile同理)。

Cartpole / Acrobot / LunarLander / MountainCar 具体网络

基类/组成归一化边界输出类型
CartpoleDQNNetworkBasicDiscreteDomainNetworkCARTPOLE_*DQNNetworkType(q_values)
CartpoleFourierDQNNetwork继承 FourierDQNNetworkCARTPOLE_*DQNNetworkType(q_values)
CartpoleRainbowNetworkBasicDiscreteDomainNetwork(num_atoms)CARTPOLE_*RainbowNetworkType
AcrobotDQNNetworkBasicDiscreteDomainNetworkACROBOT_*DQNNetworkType
AcrobotFourierDQNNetwork继承 FourierDQNNetworkACROBOT_*DQNNetworkType
AcrobotRainbowNetworkBasicDiscreteDomainNetwork(num_atoms)ACROBOT_*RainbowNetworkType
LunarLanderDQNNetworkBasicDiscreteDomainNetwork(None, None)无归一化DQNNetworkType
MountainCarDQNNetworkBasicDiscreteDomainNetworkMOUNTAINCAR_*DQNNetworkType

Rainbow 风格网络(CartpoleRainbowNetworkAcrobotRainbowNetwork)在call()中重复同样的"reshape logits → softmax → 支撑点加权"流程(见 legacy_networks.py),与 Atari 版RainbowNetwork一致。对应的函数式包装cartpole_dqn_networkcartpole_fourier_dqn_networkcartpole_rainbow_networkacrobot_dqn_networkacrobot_fourier_dqn_networkacrobot_rainbow_network均声明为@gin.configurable,函数签名示例(来自 fourier_dqn_network.md):

dopamine.discrete_domains.legacy_networks.fourier_dqn_network( min_vals, max_vals, num_actions, state, fourier_basis_order=3 )

返回"DQN 风格智能体的 Q 值或 Rainbow 风格智能体的 logits"。

在 .gin 配置中引用与参数化网络

网络类/函数均为@gin.configurable,可直接在.gin文件中按需引用。以 dqn_cartpole.gin 为例:

import dopamine.discrete_domains.gym_lib import dopamine.discrete_domains.legacy_networks import dopamine.discrete_domains.run_experiment import dopamine.tf.agents.dqn.dqn_agent import dopamine.tf.replay_memory.circular_replay_buffer import gin.tf.external_configurables DQNAgent.observation_shape = %gym_lib.CARTPOLE_OBSERVATION_SHAPE DQNAgent.observation_dtype = %gym_lib.CARTPOLE_OBSERVATION_DTYPE DQNAgent.stack_size = %gym_lib.CARTPOLE_STACK_SIZE DQNAgent.network = @legacy_networks.CartpoleDQNNetwork DQNAgent.gamma = 0.99 DQNAgent.update_horizon = 1 DQNAgent.min_replay_history = 500 DQNAgent.update_period = 4 DQNAgent.target_update_period = 100 DQNAgent.epsilon_fn = @dqn_agent.identity_epsilon DQNAgent.tf_device = '/gpu:0' # use '/cpu:*' for non-GPU version DQNAgent.optimizer = @tf.train.AdamOptimizer() tf.train.AdamOptimizer.learning_rate = 0.001 tf.train.AdamOptimizer.epsilon = 0.0003125 create_gym_environment.environment_name = 'CartPole' create_gym_environment.version = 'v0' create_agent.agent_name = 'dqn' Runner.create_environment_fn = @gym_lib.create_gym_environment Runner.num_iterations = 500 Runner.training_steps = 1000 Runner.evaluation_steps = 1000 Runner.max_steps_per_episode = 200 # Default max episode length. WrappedReplayBuffer.replay_capacity = 50000 WrappedReplayBuffer.batch_size = 128

其中DQNAgent.network = @legacy_networks.CartpoleDQNNetwork即把价值网络替换为本章介绍的 Cartpole 专用 DQN 网络。仓库内同类配置还包括:

  • dqn_acrobot.gin →@legacy_networks.AcrobotDQNNetwork
  • dqn_lunarlander.gin →@legacy_networks.LunarLanderDQNNetwork
  • dqn_mountaincar.gin →@legacy_networks.MountainCarDQNNetwork
  • c51_cartpole.gin →@legacy_networks.CartpoleRainbowNetwork,并配套RainbowAgent.num_atoms = 201RainbowAgent.vmax = 100.replay_scheme = 'uniform'
  • c51_acrobot.gin、rainbow_cartpole.gin、rainbow_acrobot.gin → 对应的 Cartpole/Acrobot Rainbow 网络

而 Atari 场景下,dqn_agent.py 与 rainbow_agent.py 在构造函数中默认指定network=legacy_networks.NatureDQNNetwork/legacy_networks.RainbowNetwork,同时复用legacy_networks.NATURE_DQN_OBSERVATION_SHAPENATURE_DQN_DTYPENATURE_DQN_STACK_SIZE三个常量,implicit_quantile_agent.py 则默认使用ImplicitQuantileNetwork。运行入口为 dopamine/discrete_domains/train.py,可通过python -m dopamine.discrete_domains.train --base_dir=... --gin_files=...方式加载上述 gin 配置启动训练。

自定义网络与 checkpoint 兼容

自定义网络的实现约定

从 dqn_agent.py 的network参数文档可见自定义网络的契约:"tf.Keras.Model,期望两个参数num_actionsnetwork_type,对其实例的调用将返回一个网络实例",并以legacy_networks.NatureDQNNetwork为示例。即:自定义网络需要接受与官方网络相同的构造参数,并在call()中返回对应的 namedtuple 输出类型,同时用@gin.configurable装饰以便配置系统实例化。

maybe_transform_variable_names:新旧 checkpoint 变量名映射

Keras 化升级改变了变量命名(例如偏置项由bias变为biases、卷积核由kernel变为weights),maybe_transform_variable_names(legacy_networks.py)用于弥合这一差异:

@gin.configurable(denylist=['variables']) def maybe_transform_variable_names(variables, legacy_checkpoint_load=False):
  • variables:待转换的全部变量列表;
  • legacy_checkpoint_load=True时,把变量名中的bias → biaseskernel → weights,映射为<new_names, var>字典,供tf.compat.v1.train.Saver加载tf.slim时代保存的旧 checkpoint 到 Keras 模型;
  • 否则返回None,即不做任何映射。

该函数在 dqn_agent.py 中通过legacy_networks.maybe_transform_variable_names(...)被实际调用。结合模块中为层统一命名'Conv'/'fully_connected'的做法,可以推断出整套兼容链路的设计意图:让 Keras 新模型的变量名尽量贴近 tf.slim 旧命名,使旧权重能够无损迁移

小结与选型建议

  • Atari 图像输入:默认选择NatureDQNNetwork(DQN)、RainbowNetwork(分布价值)或ImplicitQuantileNetwork(分位数价值),三者共享 32/64/64 卷积骨架,仅在输出头与初始化策略上分道扬镳;
  • 低维 Gym 观测:DQN 优先用CartpoleDQNNetwork等全连接网络;若观测维度较低、希望用线性函数逼近快速验证,可选用CartpoleFourierDQNNetwork/AcrobotFourierDQNNetworkfourier_basis_order默认 3);分布价值算法则对应CartpoleRainbowNetwork/AcrobotRainbowNetwork
  • 迁移旧权重:启用legacy_checkpoint_load=True即可自动完成bias/biaseskernel/weights的变量名映射。

如需深入每个函数与类的完整签名和参数文档,可继续查阅 legacy_networks 模块 API 文档 及其子页面(如 nature_dqn_network、rainbow_network、implicit_quantile_network、cartpole_dqn_network 等),并对照源码 legacy_networks.py 与 atari_lib.py 阅读实现细节。

  • 机器学习
  • 深度学习

【免费下载链接】dopamine

Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.

项目地址:https://gitcode.com/gh_mirrors/do/dopamine
点击查看免费下载
上一篇:Res2Net101_26w_4s.in1k特征提取完全指南:解锁多尺度表示能力
下一篇:量化技术深度剖析:FineTuningLLMs中的8位与4位量化原理

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/24 14:49:43

Argos Translate:一条命令安装,快速上手离线多语言翻译

Argos Translate&#xff1a;一条命令安装&#xff0c;快速上手离线多语言翻译 【免费下载链接】argos-translate Open-source offline translation library written in Python 项目地址: https://gitcode.com/GitHub_Trending/ar/argos-translate Argos Translate 是一…

作者头像 李华
网站建设 2026/9/24 14:48:15

Zonotope几何建模:虚拟电厂分布式资源不确定性聚合方法

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

作者头像 李华