ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

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

Dopamine legacy_networks 模块解析:TensorFlow 离散域网络架构与实战配置指南 机器学习深度学习【免费下载链接】dopamineDopamine 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、RainbowC51与 Implicit QuantileIQN三大卷积网络以及 Cartpole、Acrobot、LunarLander、MountainCar 等 Gym 经典环境的 DQN / Rainbow / Fourier 基网络。读完本文你将掌握这些LegacyKeras 化之前的经典架构网络的设计意图、每个函数与类的输入输出契约、在.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、ImplicitQuantileNetworkTypeKeras 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类型字段适用网络DQNNetworkTypeq_valuesDQN 风格网络RainbowNetworkTypeq_values, logits, probabilitiesRainbow / C51 风格网络ImplicitQuantileNetworkTypequantile_values, quantilesIQN 网络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, nameNone)num_actionsint动作数决定输出层维度namestr网络参数的作用域名称。网络结构legacy_networks.py层类型参数说明conv1Conv2D32 个 8×8 卷积核stride 4paddingsame输入为堆叠的 Atari 帧默认 84×84×4conv2Conv2D64 个 4×4 卷积核stride 2paddingsameconv3Conv2D64 个 3×3 卷积核stride 1paddingsameflattenFlatten—展平特征图dense1Dense512 单元ReLUdense2Densenum_actions单元无激活输出 Q 值在call()中输入先被tf.cast转为tf.float32并除以 255 归一化到[0, 1]最后返回DQNNetworkType(self.dense2(x))。源码为卷积层统一命名为Conv、全连接层命名为fully_connected注释明确说明这是为了使变量名与 tf.slim 变量名/checkpoint 更接近方便复用旧 checkpoint。RainbowNetworkC51 分布价值网络RainbowNetwork把 Q 值建模为价值分布而非标量输出层维度变为num_actions * num_atoms构造函数为RainbowNetwork(num_actions, num_atoms, support, nameNone)num_atomsint价值分布的桶bucket数量supporttf.linspaceQ 值分布的支撑点向量。关键差异legacy_networks.py所有层使用VarianceScaling(scale1.0 / np.sqrt(3.0), modefan_in, distributionuniform)初始化器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, axis2) return RainbowNetworkType(q_values, logits, probabilities)ImplicitQuantileNetwork隐分位数网络ImplicitQuantileNetwork实现 Dabney et al. (2018) 的 IQN核心思想是用随机采样的分位数驱动网络输出构造函数ImplicitQuantileNetwork(num_actions, quantile_embedding_dim, nameNone)quantile_embedding_dimint分位数输入的嵌入维度。分位数流计算legacy_networks.py卷积三件套 flatten 提取状态特征state_net_tiled沿 batch 维复制num_quantiles份从[0, 1)均匀采样num_quantiles个分位数复制到quantile_embedding_dim维用cos(i * π * tau)i1..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带归一化的全连接子网BasicDiscreteDomainNetworklegacy_networks.py是 Cartpole / Acrobot / MountainCar 等网络的公共内层模块定义为tf.keras.layers.LayerBasicDiscreteDomainNetwork(min_vals, max_vals, num_actions, num_atomsNone, nameNone, activation_fntf.keras.activations.relu)min_vals/max_vals与state同形状的最值向量用于输入归一化传None则跳过归一化如 LunarLandernum_atomsNone为 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_valsCartpole[-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(order1), repeatnvars)生成所有阶数组合的乘数向量并剔除第一个全零项对应常数偏置。特征计算为def compute_features(self, features): scaled self.scale(features) # 缩放到 [0, 1] return tf.cos(np.pi * tf.matmul(scaled, self.multipliers, transpose_bTrue))FourierDQNNetwork组合 Fourier 特征与无偏置线性层Dense(num_actions, use_biasFalse)函数签名FourierDQNNetwork(min_vals, max_vals, num_actions, fourier_basis_order3, nameNone)由于FourierBasis需要输入特征维度才能构造feature_generator在首次前向时才延迟创建与dense_quantile同理。Cartpole / Acrobot / LunarLander / MountainCar 具体网络类基类/组成归一化边界输出类型CartpoleDQNNetworkBasicDiscreteDomainNetworkCARTPOLE_*DQNNetworkType(q_values)CartpoleFourierDQNNetwork继承 FourierDQNNetworkCARTPOLE_*DQNNetworkType(q_values)CartpoleRainbowNetworkBasicDiscreteDomainNetwork(num_atoms)CARTPOLE_*RainbowNetworkTypeAcrobotDQNNetworkBasicDiscreteDomainNetworkACROBOT_*DQNNetworkTypeAcrobotFourierDQNNetwork继承 FourierDQNNetworkACROBOT_*DQNNetworkTypeAcrobotRainbowNetworkBasicDiscreteDomainNetwork(num_atoms)ACROBOT_*RainbowNetworkTypeLunarLanderDQNNetworkBasicDiscreteDomainNetwork(None, None)无归一化DQNNetworkTypeMountainCarDQNNetworkBasicDiscreteDomainNetworkMOUNTAINCAR_*DQNNetworkTypeRainbow 风格网络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.mddopamine.discrete_domains.legacy_networks.fourier_dqn_network( min_vals, max_vals, num_actions, state, fourier_basis_order3 )返回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.AcrobotDQNNetworkdqn_lunarlander.gin →legacy_networks.LunarLanderDQNNetworkdqn_mountaincar.gin →legacy_networks.MountainCarDQNNetworkc51_cartpole.gin →legacy_networks.CartpoleRainbowNetwork并配套RainbowAgent.num_atoms 201、RainbowAgent.vmax 100.、replay_scheme uniformc51_acrobot.gin、rainbow_cartpole.gin、rainbow_acrobot.gin → 对应的 Cartpole/Acrobot Rainbow 网络而 Atari 场景下dqn_agent.py 与 rainbow_agent.py 在构造函数中默认指定networklegacy_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变为weightsmaybe_transform_variable_nameslegacy_networks.py用于弥合这一差异gin.configurable(denylist[variables]) def maybe_transform_variable_names(variables, legacy_checkpoint_loadFalse):variables待转换的全部变量列表legacy_checkpoint_loadTrue时把变量名中的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 图像输入默认选择NatureDQNNetworkDQN、RainbowNetwork分布价值或ImplicitQuantileNetwork分位数价值三者共享 32/64/64 卷积骨架仅在输出头与初始化策略上分道扬镳低维 Gym 观测DQN 优先用CartpoleDQNNetwork等全连接网络若观测维度较低、希望用线性函数逼近快速验证可选用CartpoleFourierDQNNetwork/AcrobotFourierDQNNetworkfourier_basis_order默认 3分布价值算法则对应CartpoleRainbowNetwork/AcrobotRainbowNetwork迁移旧权重启用legacy_checkpoint_loadTrue即可自动完成bias/biases、kernel/weights的变量名映射。如需深入每个函数与类的完整签名和参数文档可继续查阅 legacy_networks 模块 API 文档 及其子页面如 nature_dqn_network、rainbow_network、implicit_quantile_network、cartpole_dqn_network 等并对照源码 legacy_networks.py 与 atari_lib.py 阅读实现细节。赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐LunaTranslator 使用指南视觉小说翻译从取词到排障LunaTranslator 使用指南视觉小说翻译从取词到排障 玩日文GalGame卡在每句对话上LunaTranslator值得一试。这是一款开源的视觉小机器学习深度学习AptosCore 网络模块深度解析AptosNet 架构、组件与配置指南AptosCore 网络模块深度解析AptosNet 架构、组件与配置指南 导读 AptosNet 是 Aptos 生态中任意两个节点之间通信的主协议专门服区块链Web3如何用 fuels-rs 的 WalletsConfig 3 步搞定多资产测试钱包配置如何用 fuels rs 的 WalletsConfig 3 步搞定多资产测试钱包配置 写合约测试时你大概率会遇到这样的场景一个用例需要 3 个钱包其中每机器学习深度学习上一篇Res2Net101_26w_4s.in1k特征提取完全指南解锁多尺度表示能力下一篇量化技术深度剖析FineTuningLLMs中的8位与4位量化原理创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表