十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

Dopamine JAX NoisyNetwork 解析:基于 Fortunato et al. (2018) 的参数化噪声网络实现与工程实践

Dopamine JAX NoisyNetwork 解析:基于 Fortunato et al. (2018) 的参数化噪声网络实现与工程实践 机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载导读本文聚焦于 Dopamine 强化学习框架 JAX 版网络模块中的dopamine.jax.networks.NoisyNetwork类。该类实现了 Fortunato et al. (2018) 提出的参数化噪声网络NoisyNet是完整 Rainbow 智能体FullRainbow中替代 ε-greedy 探索的核心组件。读完本文你将掌握NoisyNet 的理论动机学习式探索替代 ε-greedy、Dopamine 中分解式高斯噪声Factored Gaussian Noise的完整 JAX 实现细节、rng_key与eval_mode两个核心属性的作用以及如何通过 gin 配置在 FullRainbow 中开关噪声网络。NoisyNetwork 在 Dopamine 中的定位NoisyNetwork定义于 dopamine/jax/networks.py位于该文件的 Noisy Nets for FullRainbowNetwork 小节紧随 Rainbow 网络实现之后是 FullRainbowNetwork 的前置依赖。其 API 文档页面即本文所依据的 NoisyNetwork.md将该类定位为 Noisy Network from Fortunato et al. (2018)并记录了四个属性rng_key、eval_mode、parent、name后两者为 Flax dataclass 字段。从源码结构看NoisyNetwork 通过 feature_layer 这一工厂函数被接入网络主干def feature_layer(key, noisy, eval_modeFalse): Network feature layer depending on whether noisy_nets are used on or not. def noisy_net(x, features): return NoisyNetwork(rng_keykey, eval_modeeval_mode)(x, features) def dense_net(x, features): return nn.Dense(features, kernel_initnn.initializers.xavier_uniform())(x) return noisy_net if noisy else dense_net即当noisyTrue时隐藏层与输出层均为噪声层否则退化为普通nn.Dense采用 xavier_uniform 初始化。这一抽象让同一套FullRainbowNetwork代码可以在有噪声与无噪声两种模式下无缝切换。两个核心属性rng_key 与 eval_modeAPI 文档明确记载了该类的两个业务属性理解它们对正确使用 NoisyNetwork 至关重要属性类型说明rng_keyjax.Array文档标注为jax.interpreters.xla.DeviceArrayJAX 随机数生成器密钥用于在每次前向传播时采样噪声eval_modebool默认False是否在评估阶段关闭噪声parent/nameFlax dataclass 字段Flaxnn.Module的自动生成字段用于模块树管理在 networks.py 中对应声明为rng_key: jax.Array eval_mode: bool Falseeval_mode的核心作用是确定性评估在训练时注入噪声以驱动探索在评估时必须关闭噪声以获得稳定的策略输出。其底层逻辑在__call__中直接体现见下节。前向传播均值与方差的参数化__call__方法networks.py#L494-L529实现了论文中的核心公式891011nn.compact def __call__(self, x, features, biasTrue, kernel_initNone): def mu_init(key, shape): # Initialization of mean noise parameters (Section 3.2) low -1 / jnp.power(x.shape[0], 0.5) high 1 / jnp.power(x.shape[0], 0.5) return jax.random.uniform(key, minvallow, maxvalhigh, shapeshape) def sigma_init(key, shape, dtypejnp.float32): # Initialization of sigma noise parameters (Section 3.2) return jnp.ones(shape, dtype) * (0.1 / onp.sqrt(x.shape[0])) if self.eval_mode: # Turn off noise during evaluation w_epsilon onp.zeros(shape(x.shape[0], features), dtypeonp.float32) b_epsilon onp.zeros(shape(features,), dtypeonp.float32) else: # Factored gaussian noise in (10) and (11) in Fortunato et al. (2018). rng_p, rng_q jax.random.split(self.rng_key, num2) p NoisyNetwork.sample_noise(rng_p, [x.shape[0], 1]) q NoisyNetwork.sample_noise(rng_q, [1, features]) f_p NoisyNetwork.f(p) f_q NoisyNetwork.f(q) w_epsilon f_p * f_q b_epsilon jnp.squeeze(f_q) # See (8) and (9) in Fortunato et al. (2018) for output computation. w_mu self.param(kernel_mu, mu_init, (x.shape[0], features)) w_sigma self.param(kernel_sigma, sigma_init, (x.shape[0], features)) w w_mu jnp.multiply(w_sigma, w_epsilon) ret jnp.matmul(x, w) b_mu self.param(bias_mu, mu_init, (features,)) b_sigma self.param(bias_sigma, sigma_init, (features,)) b b_mu jnp.multiply(b_sigma, b_epsilon) return jnp.where(bias, ret b, ret)可以逐行拆解出如下设计要点权重参数化可学习参数被拆成均值kernel_mu、bias_mu与标准差kernel_sigma、bias_sigma两组最终权重为w w_mu w_sigma * w_epsilon对应论文公式89。噪声项epsilon随前向传播即时采样而mu/sigma参数随梯度下降更新——这正是噪声由网络自己学习learning the exploration的含义。分解式高斯噪声Factored Gaussian Noise训练分支使用论文公式1011的高效近似——从rng_key用jax.random.split分出两个子密钥rng_p、rng_q分别采样p形状[x.shape[0], 1]与q形状[1, features]经f变换后外积f_p * f_q得到权重噪声w_epsilon偏置噪声b_epsilon squeeze(f_q)。相比对每个元素独立采样该分解把采样量从input × features降到input features大幅减少随机数生成开销。静态方法fnetworks.py#L489-L492f(x) sign(x) * |x|^0.5即对标准正态采样施加符号保持的平方根压缩是论文公式1011中的核心变换将高斯噪声映射为均值为 0、方差为 1 的噪声分布。静态方法sample_noisenetworks.py#L485-L487jax.random.normal(key, shape)的标准正态采样封装便于测试中替换。参数初始化论文 Section 3.2mu按±1/sqrt(input_dim)均匀初始化sigma初始化为常量0.1/sqrt(input_dim)。注意到eval_modeTrue时w_epsilon、b_epsilon被直接置零等价于仅使用均值权重做确定性前向这就是评估阶段关掉噪声的实现。bias开关返回前用jnp.where(bias, ret b, ret)决定是否加偏置为上层复用如 SPR 中无偏置变体预留了接口。在 FullRainbowNetwork 中的接入位置在 FullRainbowNetwork.call中NoisyNetwork 被用在三个位置隐藏层networks.py#L588-L590x net(x, features512)即 512 维单隐层是噪声层Dueling 优势头networks.py#L593adv net(x, featuresself.num_actions * self.num_atoms)Dueling 价值头networks.py#L594value net(x, featuresself.num_atoms)。三个分支共享同一个rng_key由调用方传入或内部jax.random.PRNGKey(int(time.time() * 1e6))生成但注意每个分支是独立调用feature_layer返回的闭包因此噪声采样互不干扰。若duelingFalse则仅用单个噪声层输出num_actions * num_atoms维 logits 后 reshapenetworks.py#L598-L600。从配置到运行FullRainbow 中的开关与实践NoisyNetwork的启用由JaxFullRainbowAgent.noisy参数控制full_rainbow_agent.py#L254默认True。在 full_rainbow.gin 中可以看到标准启用方式import dopamine.jax.agents.full_rainbow.full_rainbow_agent import dopamine.jax.agents.dqn.dqn_agent import dopamine.jax.networks import dopamine.discrete_domains.atari_lib import dopamine.discrete_domains.run_experiment JaxDQNAgent.gamma 0.99 JaxDQNAgent.update_horizon 3 JaxDQNAgent.min_replay_history 20_000 # agent steps JaxDQNAgent.update_period 4 JaxDQNAgent.target_update_period 8_000 # agent steps JaxDQNAgent.epsilon_train 0.01 JaxDQNAgent.epsilon_eval 0.001 JaxDQNAgent.epsilon_decay_period 250_000 # agent steps JaxDQNAgent.optimizer adam JaxFullRainbowAgent.noisy True JaxFullRainbowAgent.dueling True JaxFullRainbowAgent.double_dqn True JaxFullRainbowAgent.num_atoms 51 JaxFullRainbowAgent.vmax 10. JaxFullRainbowAgent.replay_scheme prioritized create_optimizer.learning_rate 0.0000625 create_optimizer.eps 0.00015 atari_lib.create_atari_environment.game_name Pong atari_lib.create_atari_environment.sticky_actions True create_runner.schedule continuous_train create_agent.agent_name full_rainbow create_agent.debug_mode True Runner.num_iterations 200 Runner.training_steps 250_000 # agent steps Runner.evaluation_steps 125_000 # agent steps Runner.max_steps_per_episode 27_000 # agent steps ReplayBuffer.max_capacity 1_000_000 ReplayBuffer.batch_size 32 PrioritizedSamplingDistribution.max_capacity 1_000_000一个容易被忽视的关键联动位于 full_rainbow_agent.py#L335epsilon_fnzero_epsilon if self._noisy else epsilon_fn,也就是说当noisyTrue时智能体的 ε 被强制恒定为 0——因为探索已经完全交给噪声网络ε-greedy 被彻底关闭反之当noisyFalse时回退到默认的线性衰减 ε 探索。这一联动在多个仓库配置中得到印证OTRainbow.gin 中显式注释 Dont use noisy networks, dueling DQN, and double DQN并设置JaxFullRainbowAgent.noisy FalseDrQ.gin 与 DrQ_eps.gin 同样关闭 noisyDrQ 的 Efficient DQN 设置而 SPR.gin、DER.gin 以及 moes 的 full_rainbow.gin 则保持noisy True。离线 RL 场景则完全弃用噪声网络在 offline_rainbow_agent.py#L174 和 offline_classy_cql_agent.py#L306 中均硬编码noisyFalse并注明 No need for noisy networks for offline RL.——离线学习不需要在线探索这一取舍直接验证了 NoisyNet 作为在线探索机制的定位。相应地offline_rl/jax/networks.py#L226 中的NoisyRainbowNetwork也将noisy默认置为False并注释 No exploration in offline RL, kept for compatibility.eval_mode 在训练/评估流程中的传递eval_mode从智能体一路传递到网络层。在 JaxFullRainbowAgent 的select_action中full_rainbow_agent.py#L58-L73eval_mode作为显式参数参与动作选择Dopamine 测试中也有相应约定在 full_rainbow_agent_test.py#L96-L107 中测试智能体创建后立即设置agent.eval_mode True并配合epsilon_fnlambda w, x, y, z: 0.0保证 non-random action choices随后在 testStepEval 中验证 eval 模式下不进行训练。这印证了评估阶段的动作选择必须是确定性的噪声网络在eval_modeTrue时贡献零扰动从而与 ε-greedy 体系中的epsilon_eval 0.001保持一致的评估语义。变体实现SPR 中的定制 NoisyNetwork仓库还提供了一版为 SPRSimulated Policy Learning定制的 NoisyNetwork位于 dopamine/labs/atari_100k/spr_networks.py#L47-L118。该实现与主版本的关键差异在于eval_mode由构造属性改为__call__参数因为 SPR 需要对同一网络反复调用带噪声/不带噪声两种模式用于增强后的观测与原始观测类级属性无法满足这种动态切换需求features成为构造参数默认 512采样形状改为[x.shape[-1], 1]与[1, self.features]参数命名与初始化微调kernel/kernell、bias/biass四组参数sigma初始化为0.5/sqrt(x.shape[-1])eval_mode的实现方式不同不是跳过采样而是用jnp.where(eval_mode, zeros, epsilon)将噪声张量置零。这展示了在同一框架内针对不同训练范式在线 Rainbow 与自监督增强的 SPR对同一理论的差异化工程实现。SPR 配置SPR.gin#L33中JaxFullRainbowAgent.noisy True表明该变体依然沿用噪声探索。在 MoE 架构中的复用NoisyNetwork 还被 MoEMixture-of-Experts实验室实验复用。在 dopamine/labs/moes/architectures/networks.py 中专家网络通过base_networks.NoisyNetwork(rng_keyself.rng_key, eval_modeself.eval_mode)构造噪声层networks.py#L47-L49并根据noisy标志选择走噪声层还是普通 Dense 层networks.py#L61-L64。这进一步验证了NoisyNetwork是一个与具体网络拓扑解耦的、可插拔的线性层组件——只要传入rng_key与eval_mode即可嵌入任意 Flax 模块。小结何时使用 NoisyNetwork结合上述源码证据可以给出如下工程结论在线强化学习如 FullRainbow 在 Atari 上的训练默认noisyTrue噪声网络同时承担探索与函数逼近双重任务此时 ε-greedy 被zero_epsilon关闭评估阶段框架通过eval_mode自动去噪保证评估动作确定可复现离线 RL探索无意义仓库实现统一将noisy置为False以节省算力并保持兼容需要动态开关噪声的网络如 SPR 的自监督增强流程参考 spr_networks.py 的变体将eval_mode移入__call__参数。如需深入了解其上层行为可继续阅读 FullRainbowNetwork 的完整实现、JaxFullRainbowAgent 的构造逻辑以及 full_rainbow_agent_test.py 中关于网络输出形状与 eval 模式的测试用例。赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐如何用 fuels-rs 的 WalletsConfig 3 步搞定多资产测试钱包配置如何用 fuels rs 的 WalletsConfig 3 步搞定多资产测试钱包配置 写合约测试时你大概率会遇到这样的场景一个用例需要 3 个钱包其中每机器学习深度学习LeetCode-Go62. Unique Paths 网格路径计数的动态规划解法与 Go 实现详解LeetCode Go62. Unique Paths 网格路径计数的动态规划解法与 Go 实现详解 本文基于 LeetCode Go 仓库中第 62 题 U机器学习深度学习Dopamine JAX 经典控制环境 Rainbow 网络ClassicControlRainbowNetwork 架构解析与实战配置Dopamine JAX 经典控制环境 Rainbow 网络ClassicControlRainbowNetwork 架构解析与实战配置 ClassicCon机器学习深度学习上一篇B站视频下载终极指南5分钟掌握BilibiliDown跨平台免费下载神器下一篇EdgeDB 1.0 Alpha 7Lalande发布详解数据库级配置、RFC 1000 迁移与 CLI 变革创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表