- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
导读
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 离散域示例配置的默认网络来源。
从源码结构看,模块分为四大块:
- 常量与类型别名:
NATURE_DQN_OBSERVATION_SHAPE、NATURE_DQN_DTYPE、NATURE_DQN_STACK_SIZE三个 Atari 观测常量,以及从 atari_lib.py 复用的三个 namedtuple 类型DQNNetworkType、RainbowNetworkType、ImplicitQuantileNetworkType; - Keras Atari 卷积网络:
NatureDQNNetwork、RainbowNetwork、ImplicitQuantileNetwork三个tf.keras.Model类; - 通用 Gym 网络:
BasicDiscreteDomainNetwork全连接子网、FourierBasis特征生成器,以及 Cartpole / Acrobot / LunarLander / MountainCar 的 DQN、Fourier、Rainbow 网络; - checkpoint 兼容工具:
maybe_transform_variable_names。
三个输出类型契约(namedtuple)
网络统一通过 namedtuple 返回输出,定义于 atari_lib.py:
| 类型 | 字段 | 适用网络 |
|---|---|---|
DQNNetworkType | q_values | DQN 风格网络 |
RainbowNetworkType | q_values, logits, probabilities | Rainbow / C51 风格网络 |
ImplicitQuantileNetworkType | quantile_values, quantiles | IQN 网络 |
nature_dqn_network、rainbow_network、implicit_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):
| 层 | 类型 | 参数 | 说明 |
|---|---|---|---|
conv1 | Conv2D | 32 个 8×8 卷积核,stride 4,padding=same | 输入为堆叠的 Atari 帧(默认 84×84×4) |
conv2 | Conv2D | 64 个 4×4 卷积核,stride 2,padding=same | |
conv3 | Conv2D | 64 个 3×3 卷积核,stride 1,padding=same | |
flatten | Flatten | — | 展平特征图 |
dense1 | Dense | 512 单元,ReLU | |
dense2 | Dense | num_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)数量;support:tf.linspace,Q 值分布的支撑点向量。
关键差异(legacy_networks.py):
- 所有层使用
VarianceScaling(scale=1.0 / np.sqrt(3.0), mode='fan_in', distribution='uniform')初始化器(C51 论文推荐的小方差初始化); - 前向计算中,
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):
- 卷积三件套 + flatten 提取状态特征
state_net_tiled(沿 batch 维复制num_quantiles份); - 从
[0, 1)均匀采样num_quantiles个分位数,复制到quantile_embedding_dim维; - 用
cos(i * π * tau)(i=1..embedding_dim)做分位数嵌入,这是论文中的分位数特征映射; - 通过延迟创建的
dense_quantile层(之所以延迟,是因为其输出单元数依赖输入特征长度,只能在首次前向时确定)与状态特征做逐元素相乘,再经dense1、dense2输出quantile_values。
call()的签名为call(self, state, num_quantiles),num_quantiles表示每次前向采样的分位数个数。三个 Atari 网络对应的函数式包装分别为nature_dqn_network、rainbow_network、implicit_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_vals | max_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 具体网络
| 类 | 基类/组成 | 归一化边界 | 输出类型 |
|---|---|---|---|
CartpoleDQNNetwork | BasicDiscreteDomainNetwork | CARTPOLE_* | DQNNetworkType(q_values) |
CartpoleFourierDQNNetwork | 继承 FourierDQNNetwork | CARTPOLE_* | DQNNetworkType(q_values) |
CartpoleRainbowNetwork | BasicDiscreteDomainNetwork(num_atoms) | CARTPOLE_* | RainbowNetworkType |
AcrobotDQNNetwork | BasicDiscreteDomainNetwork | ACROBOT_* | DQNNetworkType |
AcrobotFourierDQNNetwork | 继承 FourierDQNNetwork | ACROBOT_* | DQNNetworkType |
AcrobotRainbowNetwork | BasicDiscreteDomainNetwork(num_atoms) | ACROBOT_* | RainbowNetworkType |
LunarLanderDQNNetwork | BasicDiscreteDomainNetwork(None, None) | 无归一化 | DQNNetworkType |
MountainCarDQNNetwork | BasicDiscreteDomainNetwork | MOUNTAINCAR_* | DQNNetworkType |
Rainbow 风格网络(CartpoleRainbowNetwork、AcrobotRainbowNetwork)在call()中重复同样的"reshape logits → softmax → 支撑点加权"流程(见 legacy_networks.py),与 Atari 版RainbowNetwork一致。对应的函数式包装cartpole_dqn_network、cartpole_fourier_dqn_network、cartpole_rainbow_network、acrobot_dqn_network、acrobot_fourier_dqn_network、acrobot_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 = 201、RainbowAgent.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_SHAPE、NATURE_DQN_DTYPE、NATURE_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_actions与network_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 → biases、kernel → 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/AcrobotFourierDQNNetwork(fourier_basis_order默认 3);分布价值算法则对应CartpoleRainbowNetwork/AcrobotRainbowNetwork; - 迁移旧权重:启用
legacy_checkpoint_load=True即可自动完成bias/biases、kernel/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.
相关推荐
LunaTranslator 使用指南:视觉小说翻译从取词到排障
LunaTranslator 使用指南:视觉小说翻译从取词到排障 玩日文GalGame卡在每句对话上,LunaTranslator值得一试。这是一款开源的视觉小
机器学习深度学习AptosCore 网络模块深度解析:AptosNet 架构、组件与配置指南
AptosCore 网络模块深度解析:AptosNet 架构、组件与配置指南 导读 AptosNet 是 Aptos 生态中任意两个节点之间通信的主协议,专门服
区块链Web3如何用 fuels-rs 的 WalletsConfig 3 步搞定多资产测试钱包配置
如何用 fuels rs 的 WalletsConfig 3 步搞定多资产测试钱包配置 写合约测试时,你大概率会遇到这样的场景:一个用例需要 3 个钱包,其中每
机器学习深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考