首页/新闻资讯/正文详情

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

发布时间:2026/9/24 19:04:18 来源:云帆数科 栏目:资讯中心
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),仅供参考

相关推荐

OPC UA双认证兼容方案:匿名与用户名密码在C# .NET 8中的实现
OPC UA双认证兼容方案:匿名与用户名密码在C# .NET 8中的实现

搞OPC UA上位机开发的同学,十有八九都遇到过这种场景:客户端用默认配置去连服务器,结果要么连不上,要么连上了什么数据都读不到。尤其是到了C# .NET 8下面跑OPC UA,一上来就踩坑的,往往不是加密算法配得不对… · 2026/9/24 19:04:12

C# OPC UA客户端双认证方案:避开匿名登录陷阱的实战指南
C# OPC UA客户端双认证方案:避开匿名登录陷阱的实战指南

去年做一个设备数据采集项目时,我踩过一个印象特别深的坑:PLC 侧的 OPC UA 服务器是设备厂商调好的,我这边要写一个 C# 上位机服务去对接。开发阶段图省事,客户端连接全部走匿名登录(AnonymousIdentityToken&#xff0… · 2026/9/24 19:04:06

Windows宝塔面板部署Python项目:Nginx反向代理与502排查实战
Windows宝塔面板部署Python项目:Nginx反向代理与502排查实战

上个月同事扔给我一个Flask项目,说要在Windows服务器的宝塔面板上部署,我一开始觉得这不是有手就行?结果从配Nginx到看502,整整折腾了大半个晚上。最气人的是,配置文件明明改了,nginx -t也提示正常&#xf… · 2026/9/24 19:04:06

AI原生数据治理选型指南:五大平台能力分化与决策框架
AI原生数据治理选型指南:五大平台能力分化与决策框架

1. 当数据治理撞上AI原生,选型逻辑为什么突然变了过去几年做数据治理,大家聊得最多的是元数据采集覆盖率、血缘解析准确率、数据质量规则跑批时长这些指标。但从2025年下半年开始,我陆续参与了几个大型企业的数据平台升级评审,发现… · 2026/9/24 19:36:40

GNG生长型神经气体网络:自适应聚类的动态拓扑解法
GNG生长型神经气体网络:自适应聚类的动态拓扑解法

1. 什么是GNG生长型神经气体网络?它为什么能甩开K-means和DBSCAN几条街? “GNG生长型神经气体网络”——光看这名字,很多人第一反应是:又一个拗口的学术黑话。但如果你正在处理客户分群、异常检测、传感器数据压缩,或者… · 2026/9/24 19:36:40

MySQL状态查看与Navicat连接排查:从SHOW STATUS到Access denied实战
MySQL状态查看与Navicat连接排查:从SHOW STATUS到Access denied实战

说实话,这节MySQL课的后两节,信息量比前面几节加起来都大。老师先带我们把SHOW STATUS过了一遍,然后现场演示了 Navicat 链接 MySQL 的完整流程,下课的时候还有一半人卡在 Access denied 上——包括我。回来我花了一整个晚上把课堂… · 2026/9/24 19:36:40

acore-db-app:Python封装库,让AzerothCore数据库操作化繁为简
acore-db-app:Python封装库,让AzerothCore数据库操作化繁为简

维护AzerothCore服务端的朋友应该都有过这种经历:开发到后期,各种数据修复、批量任务、跨库同步的需求接踵而来,每天不是在写SQL,就是在写连接数据库的Python脚本。我自己的痛点是,pymysql裸用起来倒是不难&#xff0c… · 2026/9/24 19:36:40

MySQL 状态查看与 Navicat 连接失败排查指南
MySQL 状态查看与 Navicat 连接失败排查指南

上午后两节课,正好讲到了MySQL状态和Navicat链接MySQL,这两块其实都是日常开发里最高频的操作:一个是判断数据库到底健不健康,一个是让你从黑窗口里解放出来。如果你刚装好MySQL不知道下一步干什么,或者被Navicat连接时… · 2026/9/24 19:36:40

Python智慧教室源码实战:专注度分析、作弊检测与动态点名
Python智慧教室源码实战:专注度分析、作弊检测与动态点名

简介:这是一套面向教育技术开发者与Python学习者的智慧教室综合实践源码,围绕课堂专注度分析、考试作弊检测与动态点名三大场景展开,适合希望将计算机视觉、自然语言处理落地到教学管理的中级开发者参考。压缩包共218个文件、约17.04MB&#… · 2026/9/24 19:36:34

基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程
基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程

简介:这是一套面向计算机、人工智能、自动化等专业学生与教师的毕业设计级项目资源,围绕YOLOv8实现渔船作业监控系统,可用于毕设、课程设计、大作业或项目立项演示。压缩包共97个文件,约24.21MB,以70个Python源码文件为… · 2026/9/24 0:00:13

1D-CNN时间序列建模实战:从Conv1d原理到工业落地
1D-CNN时间序列建模实战:从Conv1d原理到工业落地

简介:面向时间序列数据建模的一维卷积神经网络完整实现,适合深度学习入门者及需要快速验证时序模型的研究者,能够从音频、文本、传感器或股价等序列中挖掘局部特征与时间依赖。压缩包体积很小,只有3KB,内含3个Python脚… · 2026/9/24 0:00:26

柔软的L:汉语语流中被忽视的舌肌张力控制
柔软的L:汉语语流中被忽视的舌肌张力控制

1. 这个“L”不是字母表里的L,而是舌尖上的L最近在几个方言群和语音教学社群里,反复看到有人发一句:“也说字母L:柔软的长舌”。初看以为是英语发音课笔记,点开才发现全是方言爱好者、播音系学生、语言康复师甚至戏曲演… · 2026/9/24 0:00:44

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码