尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Dopamine 中的 QuantileNetwork 详解:基于 JAX 的分位数回归网络结构与实战配置

Dopamine 中的 QuantileNetwork 详解:基于 JAX 的分位数回归网络结构与实战配置 机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载导读QuantileNetwork是 Dopamine 研究框架JAX 后端中用于计算智能体收益分位数return quantiles的卷积神经网络是实现 Quantile Regression DQNQR-DQNDabney et al., 2017的核心构件。它把经典 Nature DQN 的卷积骨干与分位数输出层相结合为每个动作输出num_atoms个分位数值。本文结合 docs/api_docs/python/dopamine/jax/networks/QuantileNetwork.md 的 API 说明从源码、训练流程与 gin 配置三个层面完整剖析该网络帮助你在 Atari 2600 与 MinAtar 环境中理解、替换与调优它。一、QuantileNetwork 是什么在分位数回归强化学习中智能体不再只估计状态-动作值函数 Q(s, a) 的期望而是估计整个收益分布的若干分位点。QuantileNetwork正是负责从原始观测如 Atari 游戏帧出发输出这一组分位数值的卷积网络。其官方 API 描述见 QuantileNetwork.md为Convolutional network used to compute the agents return quantiles.在仓库中的实际定义为dopamine/jax/networks.py### Quantile Networks ### gin.configurable class QuantileNetwork(nn.Module): Convolutional network used to compute the agents return quantiles. num_actions: int num_atoms: int inputs_preprocessed: bool False它继承自 Flax 的nn.Module并通过gin.configurable暴露给 gin 配置系统这意味着你可以在.gin文件中直接通过QuantileNetwork.xxx ...覆盖其字段默认值。二、Dataclass 字段参数语义与默认值API 文档的 Attributes 表列出了该类作为 Flaxnn.Moduledataclass的全部字段字段类型含义与取值建议num_actionsint智能体可执行的动作数量。输出层将据此生成num_actions * num_atoms个神经元。在 Atari 环境中通常为 18create_atari_environment会根据游戏自动设置参见 atari_lib.py。num_atomsint收益分布被离散为的分位数个数论文中也称N。Dopamine 默认使用 200与 Rainbow 论文保持一致见下文 gin 配置。取值越大分布表达越精细但计算与显存开销线性增长。inputs_preprocessedbool输入是否已经完成预处理。默认False此时网络内部会把输入除以 255.0 归一化若为True则跳过该步骤适用于已在外部预处理过的输入。parentnn.ModuleFlax dataclass 自动生成的父模块字段框架内部使用。namestrFlax dataclass 自动生成的模块名称字段用于作用域标识框架内部使用。其中num_actions与num_atoms是定义网络拓扑的关键超参数由JaxQuantileAgent在初始化时通过functools.partial(network, num_atomsnum_atoms)注入见 quantile_agent.py因此你在 gin 中只需配置 agent 级的JaxQuantileAgent.num_atoms网络字段会自动同步。三、网络架构逐层拆解QuantileNetwork的__call__逻辑与 Nature DQN 卷积骨干完全一致仅在输出层改为分位数形式。以下是 networks.py 中的完整前向过程nn.compact def __call__(self, x): initializer nn.initializers.variance_scaling( scale1.0 / jnp.sqrt(3.0), modefan_in, distributionuniform ) if not self.inputs_preprocessed: x preprocess_atari_inputs(x) # x.astype(jnp.float32) / 255.0 x nn.Conv(features32, kernel_size(8, 8), strides(4, 4), kernel_initinitializer)(x) x nn.relu(x) x nn.Conv(features64, kernel_size(4, 4), strides(2, 2), kernel_initinitializer)(x) x nn.relu(x) x nn.Conv(features64, kernel_size(3, 3), strides(1, 1), kernel_initinitializer)(x) x nn.relu(x) x x.reshape((-1)) # flatten x nn.Dense(features512, kernel_initinitializer)(x) x nn.relu(x) x nn.Dense(featuresself.num_actions * self.num_atoms, kernel_initinitializer)(x) logits x.reshape((self.num_actions, self.num_atoms)) probabilities nn.softmax(logits) q_values jnp.mean(logits, axis1) return atari_lib.RainbowNetworkType(q_values, logits, probabilities)各层的作用如下输入预处理若inputs_preprocessedFalse调用preprocess_atari_inputs把 uint8 像素帧转为float32并缩放到[0, 1]networks.py。卷积骨干3 层卷积32×8×8 步长 4、64×4×4 步长 2、64×3×3 步长 1每层后接 ReLU与 Nature DQN 相同用于提取游戏画面特征。全连接层展平后接 512 维全连接与 ReLU。分位数输出层Dense(num_actions * num_atoms)把特征映射为num_actions * num_atoms个标量再reshape为(num_actions, num_atoms)的 logits 矩阵——每一行代表某个动作的num_atoms个分位数值。概率与 Q 值对 logits 沿分位数维做softmax得到probabilitiesq_values则取 logits 的均值jnp.mean(logits, axis1)得到每个动作的期望 Q 值近似。初始化策略注意输出层与卷积层使用variance_scaling(scale1/sqrt(3), modefan_in, distributionuniform)与RainbowNetwork一致而非 DQN 的xavier_uniform。这种初始化直接作用于分位数 logits 的数值尺度是保证训练早期数值稳定的细节之一。四、输出类型RainbowNetworkType网络返回atari_lib.RainbowNetworkType(q_values, logits, probabilities)这是一个命名元组定义于 dopamine/discrete_domains/atari_lib.pyRainbowNetworkType collections.namedtuple( c51_network, [q_values, logits, probabilities] )三个字段的语义与形状在num_actionsA、num_atomsN时字段形状含义q_values(A,)每个动作的期望 Q 值分位数 logits 的均值用于动作选择argmax。logits(A, N)每个动作在 N 个分位点上的原始输出训练时用于计算分位数损失。probabilities(A, N)对 logits 做 softmax 后的分布QR-DQN 中 softmax 仅用于输出约定实际损失计算用的是 logits。之所以复用RainbowNetworkType是因为 QR-DQN 与 C51 都采用每动作一个分布的输出约定Dopamine 据此统一了接口便于在两种 agent 之间复用训练/评估代码。五、与 JaxQuantileAgent 的协作从网络到训练QuantileNetwork的实际使用方是 dopamine/jax/agents/quantile/quantile_agent.py 中的JaxQuantileAgent其默认网络即指向它networknetworks.QuantileNetwork, kappa1.0, num_atoms200, gamma0.99,1. 网络注入在__init__中agent 通过functools.partial(network, num_atomsnum_atoms)把num_atoms绑定进网络构造随后在_build_networks_and_optimizer中调用self.network_def.init(rng, xself.state)完成参数初始化quantile_agent.py。2. 目标分布构造训练时target_distribution函数quantile_agent.py通过目标网络对下一状态打分取q_values的 argmax 确定贪心动作next_qt_argmax从logits中取出该动作对应的分位数向量next_logits目标分位数 rewards gamma * next_logits终端状态乘数为 0并用stop_gradient阻断梯度。3. 分位数 Huber 损失train函数quantile_agent.py实现论文 Eq. 9-10计算贝尔曼误差矩阵bellman_errors用kappa默认 1.0截断的 Huber 损失作为基础损失构造分位中点tau_hat (arange(num_atoms) 0.5) / num_atoms分位数损失 |tau_hat - (bellman_errors 0)| * huber_loss对分位数维求和、目标值维求平均。这意味着num_atoms既决定了网络的输出宽度也直接决定了损失计算的矩阵维度是 QR-DQN 的核心超参数。六、gin 配置实战默认参数与 Atari 环境QuantileNetwork在 Atari 上的完整默认配置见 dopamine/jax/agents/quantile/configs/quantile.ginimport dopamine.jax.agents.quantile.quantile_agent import dopamine.discrete_domains.atari_lib import dopamine.discrete_domains.run_experiment # 网络/损失相关核心参数 JaxQuantileAgent.kappa 1.0 # Huber 损失截断点 JaxQuantileAgent.num_atoms 200 # 分位数个数即网络的 num_atoms JaxQuantileAgent.gamma 0.99 JaxQuantileAgent.update_horizon 3 # n-step 更新 JaxQuantileAgent.min_replay_history 20000 JaxQuantileAgent.update_period 4 JaxQuantileAgent.target_update_period 8000 JaxQuantileAgent.epsilon_train 0.01 JaxQuantileAgent.epsilon_eval 0.001 JaxQuantileAgent.epsilon_decay_period 250000 JaxQuantileAgent.replay_scheme prioritized # 优先经验回放 JaxQuantileAgent.optimizer adam # 优化器影响网络参数更新 create_optimizer.learning_rate 0.00005 create_optimizer.eps 0.0003125 # 环境与训练调度 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 jax_quantile Runner.num_iterations 200 Runner.training_steps 250000 Runner.evaluation_steps 125000 Runner.max_steps_per_episode 27000 # 回放缓冲区 ReplayBuffer.max_capacity 1_000_000 ReplayBuffer.batch_size 32 PrioritizedSamplingDistribution.max_capacity 1_000_000关键参数说明num_atoms 200与 RainbowHessel et al., 2018保持一致配置文件头部注释明确说明Hyperparameters follow Dabney et al. (2017) but we modify as necessary to match those used in Rainbow。kappa 1.0分位数 Huber 损失的截断值控制损失对异常值的鲁棒性。replay_scheme prioritizedQR-DQN 默认启用优先经验回放。在_train_step中agent 以sqrt(loss 1e-10)作为优先级更新样本权重并用逆优先级对损失加权quantile_agent.py同时把平均损失写入 TensorBoard 的QuantileLoss标量。运行训练的命令JAX 版本使用dopamine/discrete_domains/train.pypython -m dopamine.discrete_domains.train \ --base_dir/tmp/dopamine/quantile \ --gin_filesdopamine/jax/agents/quantile/configs/quantile.gin七、替换与自定义网络测试用例给出的模板API 文档把QuantileNetwork作为可插拔模块设计JaxQuantileAgent.network接受任何满足输出(num_actions, num_atoms)形状 logits约定的 Flaxnn.Module。这一点被 tests/dopamine/jax/agents/quantile/quantile_agent_test.py 中的MockQuantileNetwork明确验证class MockQuantileNetwork(linen.Module): Custom Jax network used in tests. num_actions: int num_atoms: int inputs_preprocessed: bool False linen.compact def __call__(self, x): ... x linen.Dense( featuresself.num_actions * self.num_atoms, kernel_initcustom_init, bias_initlinen.initializers.ones, )(x) logits x.reshape((self.num_actions, self.num_atoms)) probabilities linen.softmax(logits) qs jnp.mean(logits, axis1) return atari_lib.RainbowNetworkType(qs, logits, probabilities)测试同时校验了输出契约quantile_agent_test.pylogits.shape (num_actions, num_atoms)probabilities.shape logits.shapeq_values.shape (num_actions,)因此若你要替换网络例如换成 Impala 骨干或轻量 MLP只需保证① 是nn.Module② 接受观测输入返回RainbowNetworkType③ 输出形状遵循上述约定。在 gin 中通过JaxQuantileAgent.network your.module即可无缝接入agent 的初始化、回放与损失计算代码无需任何改动。八、轻量变体MinatarQuantileNetwork对于低分辨率环境如 MinAtar 的 10×10 单帧输入Dopamine 提供了对应的轻量实现 dopamine/labs/environments/minatar/minatar_env.py 中的MinatarQuantileNetwork与QuantileNetwork保持相同的输出契约num_actions、num_atoms、inputs_preprocessed仅把卷积骨干替换为适配小尺寸输入的浅层结构。其 gin 配置示例quantile_space_invaders.gin展示了如何在网络之上配置 agentJaxQuantileAgent.observation_shape %minatar_env.SPACE_INVADERS_SHAPE JaxQuantileAgent.observation_dtype %minatar_env.DTYPE JaxQuantileAgent.stack_size 1 JaxQuantileAgent.network minatar_env.MinatarQuantileNetwork JaxQuantileAgent.kappa 1.0 JaxQuantileAgent.num_atoms 200 JaxQuantileAgent.gamma 0.99这证明QuantileNetwork的卷积骨干 分位数输出层模式具备良好的可移植性更换环境只需替换骨干网络输出层与训练逻辑完全复用。九、小结QuantileNetwork是 Dopamine JAX 生态中 QR-DQN 智能体的标准网络它以 Nature DQN 卷积骨干提取特征通过num_actions * num_atoms的全连接层输出每个动作的收益分位数并以RainbowNetworkType统一封装q_values/logits/probabilities。理解它的字段语义num_actions、num_atoms、inputs_preprocessed、初始化策略与输出契约是配置、替换乃至自定义分布强化学习网络的基础。相关代码与证据可继续查阅网络实现dopamine/jax/networks.pyAgent 与损失dopamine/jax/agents/quantile/quantile_agent.py默认配置dopamine/jax/agents/quantile/configs/quantile.gin输出类型定义dopamine/discrete_domains/atari_lib.py契约测试tests/dopamine/jax/agents/quantile/quantile_agent_test.py赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐Dopamine JAX 中 ImplicitQuantileNetworkIQN网络详解分位数嵌入结构、源码实现与训练配置Dopamine JAX 中 ImplicitQuantileNetworkIQN网络详解分位数嵌入结构、源码实现与训练配置 导读 ImplicitQua机器学习深度学习深入解析 Dopamine 中的 Quantile Regression DQN基于 JAX 的分位数回归强化学习智能体深入解析 Dopamine 中的 Quantile Regression DQN基于 JAX 的分位数回归强化学习智能体 导读 本文围绕 Dopamine 研机器学习深度学习Dopamine 框架中的 JAX Quantile DQN基于分位数回归的分布强化学习智能体全解析Dopamine 框架中的 JAX Quantile DQN基于分位数回归的分布强化学习智能体全解析 导读 本文聚焦于 Dopamine 研究框架中 JAX强化学习机器学习深度学习上一篇5分钟跑通 BabelDOC PDF翻译安装、命令与排错指南下一篇uWebSockets日志结构化JSON格式与字段标准化创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表