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

Dopamine JAX 中的 huber_loss:分段平滑损失函数的实现、原理与在 DQN / 分布强化学习中的应用

发布时间:2026/9/24 15:08:03 来源:云帆数科 栏目:资讯中心
Dopamine JAX 中的 huber_loss:分段平滑损失函数的实现、原理与在 DQN / 分布强化学习中的应用
机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载导读dopamine.jax.losses.huber_loss是 Dopamine 强化学习框架 JAX 分支中定义的分段损失函数用于在 TD 误差回归时同时兼顾 MSE 的快速收敛与 MAE 对异常值的鲁棒性。本指南以官方 API 文档 docs/api_docs/python/dopamine/jax/losses/huber_loss.md 为核心完整讲解该函数的数学定义、参数语义、源码实现位于 dopamine/jax/losses.py并结合 DQN、Quantile、IQN 等 JAX Agent 的训练调用链与单元测试说明如何在实际强化学习训练中使用并选择delta阈值。读完本文你将掌握 Huber loss 在 Dopamine JAX 中的精确行为、如何通过loss_type切换损失函数以及它为何是分布强化学习QR-DQN / IQN的默认底座。一、函数签名与官方定义huber_loss位于dopamine/jax/losses.py其完整签名如下def huber_loss( targets: jnp.ndarray, predictions: jnp.ndarray, delta: float 1.0 ) - jnp.ndarray:参数类型含义targetsjnp.ndarray目标值Target values在强化学习中通常是贝尔曼目标Bellman targetpredictionsjnp.ndarray预测值Prediction values通常是 Q 网络对采样动作的输出deltafloat默认1.0阈值Threshold决定误差从二次区切换为线性区的分界点返回值Huber loss一个与输入形状一致的jnp.ndarray逐元素计算不做均值归约。设x |targets - predictions|官方定义的分段公式为当x delta时0.5 * x^2当x delta时0.5 * delta^2 delta * (x - delta)即小误差区域使用平方误差二次、处处可导、梯度随误差减小而衰减大误差区域使用线性误差梯度恒为delta不会因异常样本产生爆炸性梯度。0.5 * delta^2 delta * (x - delta)这一线性形式保证了在x delta处函数值连续代入x delta得0.5 * delta^2与二次分支平滑衔接。二、源码实现逐行解析huber_loss的实现极为精简仅有三行核心代码dopamine/jax/losses.pyx jnp.abs(targets - predictions) return jnp.where(x delta, 0.5 * x**2, 0.5 * delta**2 delta * (x - delta))实现要点第一步计算逐元素绝对误差x第二步使用jnp.where(condition, on_true, on_false)按元素选择满足x delta的位置取二次分支其余位置取线性分支整个函数基于jax.numpy构建天然支持JIT 编译、自动微分jax.grad/jax.value_and_grad与vmap 向量化因此可以直接嵌入被jax.jit装饰的训练函数中参与反向传播。在 dopamine/jax/losses.py 中还定义了两个同模块损失函数可一并对比理解设计取向mse_loss(targets, predictions)jnp.power((targets - predictions), 2)恒为二次损失梯度随误差线性增大softmax_cross_entropy_loss_with_logits(labels, logits)-jnp.sum(labels * nn.log_softmax(logits))用于策略/分类目标。三、单元测试中的数值行为验证测试文件 tests/dopamine/jax/losses_test.py 使用parameterized参数化测试给出了四组可直接验证数值的用例恰好覆盖了delta的全部关键情形测试用例targetspredictionsdelta期望输出说明BelowDelta1d1.00.01.00.5x1 delta二次分支0.5 * 1^2AboveDelta1d1.00.00.50.375x1 delta线性分支0.5*0.25 0.5*0.5MixedArraysDefaultDeltaones(5)[0,1,2,3,4]默认1.0[0.5, 0.0, 0.5, 1.5, 2.5]数组混合自动逐元素分流MixedArraysSetDeltaones(5)[0,1,2,3,4]2.0[0.5, 0.0, 0.5, 2.0, 4.0]增大delta扩大二次区以MixedArraysDefaultDelta为例x [1, 0, 1, 2, 3]前三个元素x 1走二次分支得到[0.5, 0, 0.5]后两个元素x 2, 3走线性分支分别得到0.5 1*(2-1) 1.5与0.5 1*(3-1) 2.5。对比最后一组可见将delta从1.0增大到2.0后x2的元素从线性区回到二次区输出从1.5变为2.0直观展示了阈值对损失形状的调控作用。四、在 DQN JAX Agent 中的集成loss_type 切换huber_loss在 Dopamine JAX 中最直接的消费方是 DQN Agent。训练函数train通过loss_type参数选择损失函数dopamine/jax/agents/dqn/dqn_agent.pydef loss_fn(params, target): def q_online(state): return network_def.apply(params, state) q_values jax.vmap(q_online)(states).q_values q_values jnp.squeeze(q_values) replay_chosen_q jax.vmap(lambda x, y: x[y])(q_values, actions) if loss_type huber: return jnp.mean(jax.vmap(losses.huber_loss)(target, replay_chosen_q)) return jnp.mean(jax.vmap(losses.mse_loss)(target, replay_chosen_q))关键调用链在线网络对批量状态输出 Q 值jax.vmap(lambda x, y: x[y])取出每个样本实际执行动作对应的 Q 值replay_chosen_q目标网络计算 TD 目标target target_q(...)即R_t γ^N * max_a Q(s, a)见 dopamine/jax/agents/dqn/dqn_agent.pyjax.vmap(losses.huber_loss)(target, replay_chosen_q)对批内每个样本逐元素计算 Huber loss再jnp.mean得到标量损失损失经jax.value_and_grad反向传播更新在线网络参数。loss_type是JaxDQNAgent.__init__的构造参数默认值为msedopamine/jax/agents/dqn/dqn_agent.py文档注释明确说明其语义为whether to use Huber or MSE loss during training。因此训练时只需将loss_typehuber传入JaxDQNAgent即可在 TD 回归中启用带阈值的平滑损失。在 dopamine/labs/tandem_dqn/tandem_dqn_agent.py 中Tandem DQN 甚至直接将默认值设为loss_typehuber说明该模式已被实验室 Agent 作为默认训练配置使用。五、在分布强化学习中的核心地位Huber loss 更深层的价值体现在分布强化学习Distributional RL中它被用作量化回归quantile regression的分位数 Huber 损失quantile Huber loss的底座。QR-DQNQuantile Agent在 dopamine/jax/agents/quantile/quantile_agent.py 中以内联方式实现分位数 Huber 损失huber_loss (jnp.abs(bellman_errors) kappa).astype(jnp.float32) * 0.5 * bellman_errors**2 \ (jnp.abs(bellman_errors) kappa).astype(jnp.float32) * kappa * (jnp.abs(bellman_errors) - 0.5 * kappa) tau_bellman_diff jnp.abs(tau_hat[None, :, None] - (bellman_errors 0).astype(jnp.float32)) quantile_huber_loss tau_bellman_diff * huber_loss其损失形状batch_size x num_atoms x num_atomskappa即阈值等价于huber_loss中的delta默认取值由 Agent 构造参数传入。IQNImplicit Quantile Agent在 dopamine/jax/agents/implicit_quantile/implicit_quantile_agent.py 中显式拆分为两个分支huber_loss_case_one (jnp.abs(bellman_errors) kappa).astype(jnp.float32) * 0.5 * bellman_errors**2 huber_loss_case_two (jnp.abs(bellman_errors) kappa).astype(jnp.float32) * kappa * (jnp.abs(bellman_errors) - 0.5 * kappa) huber_loss huber_loss_case_one huber_loss_case_two quantile_huber_loss jnp.abs(quantiles - jax.lax.stop_gradient((bellman_errors 0).astype(jnp.float32))) * huber_loss / kappa注意 IQN 在此处将 Huber 损失除以kappa归一化这是其与 dopamine/tf/agents/implicit_quantile/implicit_quantile_agent.py 中 TF 版实现一致的处理方式。这些内联实现与losses.huber_loss在分段逻辑上完全同构——同样以|误差| 阈值划分二次区与线性区说明huber_loss是分布强化学习中误差鲁棒化的标准模板。此外在 dopamine/jax/agents/full_rainbow/full_rainbow_agent.py 中Full Rainbow Agent 通过losses.mse_loss if mse_loss else losses.huber_loss在 MSE 与 Huber 之间切换进一步印证了该损失函数在 Dopamine JAX 各 Agent 中的通用性。六、为什么强化学习倾向于使用 Huber loss结合上述实现可以从三个角度理解huber_loss在 TD 学习中的工程价值对异常 TD 误差的鲁棒性训练初期或探索阶段可能出现远超正常范围的贝尔曼误差MSE 的二次梯度会放大这些异常样本的更新幅度导致训练震荡Huber 损失在线性区梯度恒定为delta天然抑制了大误差样本的主导作用。小误差区域的精细学习当误差小于delta时保持二次形式梯度随误差缩小而衰减有利于在接近收敛时进行精细的参数调整。处处可微相比 MAE 在零点不可导Huber 损失在x delta处函数值连续且两侧导数均为delta可与jax.value_and_grad、optax优化器无缝配合。七、使用建议与调参指引默认阈值delta默认为1.0。当 TD 目标与预测值尺度远大于 1如未归一化的奖励累加时可考虑增大delta以扩大二次区反之若误差普遍很小可适当减小delta以获得更强的异常值抑制。通过 Agent 参数切换在 DQN 系列中传loss_typehuber默认mse在分布强化学习中kappa即阈值作为构造参数传入 Agent例如 QR-DQN / IQN 的kappa参数。逐元素语义huber_loss不做均值归约返回与输入同形状的数组在实际训练中需要配合jnp.meanDQN或jnp.sum分布强化学习对分位数维求和完成归约详见各 Agent 的loss_fn。向量化批量计算批量训练时统一使用jax.vmap(losses.huber_loss)(target, replay_chosen_q)对批次逐样本计算Dopamine 各 JAX Agent如 dopamine/labs/atari_100k/spr_agent.py、dopamine/labs/offline_rl/jax/offline_rainbow_agent.py均采用这一模式。八、延伸阅读路径官方 API 文档docs/api_docs/python/dopamine/jax/losses/huber_loss.md损失函数完整实现含mse_loss、softmax_cross_entropy_loss_with_logitsdopamine/jax/losses.py数值验证测试tests/dopamine/jax/losses_test.pyDQN 中的loss_type切换与调用链dopamine/jax/agents/dqn/dqn_agent.py分位数 Huber 损失在 QR-DQN / IQN 中的实现dopamine/jax/agents/quantile/quantile_agent.py、dopamine/jax/agents/implicit_quantile/implicit_quantile_agent.pyTF 分支的对应实现dopamine/tf/agents/implicit_quantile/implicit_quantile_agent.py赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐Dopamine JAX 损失函数深度解析softmax_cross_entropy_loss_with_logits 与分布强化学习Dopamine JAX 损失函数深度解析softmax_cross_entropy_loss_with_logits 与分布强化学习 本篇技术指南围绕 Do机器学习深度学习OpenObserve 日志查询过滤延迟调优P95 从 520ms 到 45ms改了 4 处OpenObserve 日志查询过滤延迟调优P95 从 520ms 到 45ms改了 4 处 我们把一条四条件日志查询放到 OpenObserve 生产集群机器学习深度学习Dopamine 框架中的 JAX Quantile DQN基于分位数回归的分布强化学习智能体全解析Dopamine 框架中的 JAX Quantile DQN基于分位数回归的分布强化学习智能体全解析 导读 本文聚焦于 Dopamine 研究框架中 JAX强化学习机器学习深度学习上一篇ImageDedup深度学习图像去重完整指南构建高效的智能图像查重系统下一篇杜比大喇叭β版安装与使用指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

相关推荐

Yii 2 后向兼容性(BC)承诺解读:从补丁发布到类、接口与常量变更的完整兼容性清单
Yii 2 后向兼容性(BC)承诺解读:从补丁发布到类、接口与常量变更的完整兼容性清单

Yii 2 后向兼容性(BC)承诺解读:从补丁发布到类、接口与常量变更的完整兼容性清单 【免费下载链接】yii2 Yii 2: The Fast, Secure and Professional PHP Framework 项目地址: https://gitcode.com/gh_mirrors/yi/yii2 Yii 2 以「快速、… · 2026/9/24 15:08:03

MouseClick光标美化功能完全指南:9款精美光标主题一键安装与应用
MouseClick光标美化功能完全指南:9款精美光标主题一键安装与应用

MouseClick光标美化功能完全指南:9款精美光标主题一键安装与应用 【免费下载链接】MouseClick 🖱️ MouseClick 🖱️ 是一款功能强大的鼠标连点器和管理工具,采用 Qt Widget 开发 ,具备跨平台兼容性 。软件界面美观 &a… · 2026/9/24 15:08:03

Python整理《电力系统分析》复习题:docx格式清洗与公式保留指南
Python整理《电力系统分析》复习题:docx格式清洗与公式保留指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 15:07:57

2026企业AI办公工具选型指南:框架、平台盘点与落地场景
2026企业AI办公工具选型指南:框架、平台盘点与落地场景

企业引入AI办公工具时,很容易陷入以功能清单判断产品价值的误区。不少管理者会横向罗列各家平台的能力项,用功能数量多少作为取舍依据,或是单纯以采购成本、品牌声量决定选型方向。这种评估方式容易造成AI工具上线之后,难以融入现… · 2026/9/24 15:31:47

Hive Web Scrape Tool 深度指南:基于 Playwright Stealth 的无头浏览器网页内容提取与 SSRF 防护
Hive Web Scrape Tool 深度指南:基于 Playwright Stealth 的无头浏览器网页内容提取与 SSRF 防护

人工智能AI Agent多智能体MCP 服务工具调用浏览器控制 【免费下载链接】hive Multi-Agent Harness for Production AI 项目地址: https://gitcode.com/gh_mirrors/hive48/hive 点击查看 免费下载 导读 web_scrape 是 Hive 多 Agent 生产框架(hive_tool… · 2026/9/24 15:31:35

AI 生成的对比表格怎样转成可计算 Excel?
AI 生成的对比表格怎样转成可计算 Excel?

把 AI 回答里的表格粘进 Excel 后,表面看像一张表,却常常无法求和、筛选或透视。原因通常不在 Excel,而在输入:Markdown 表格只是文本结构;货币符号、百分号、千分位逗号和空格也可能让数字被当作文本。最省事的方式是… · 2026/9/24 15:31:35

2026 HUAWEI HiCar 认证新变化,车载设备研发必看要
2026 HUAWEI HiCar 认证新变化,车载设备研发必看要

​2026 年是 HiCar 认证变化较多的一年。V6.0.0 规范已经全面落地,HarmonyOS NEXT 生态全面铺开,安全要求提升了多个等级。这些变化叠加在一起,对做前装车机、后装盒子、车载应用的厂商都产生了实际影响。这篇文章把 2026 年较为关键的几个变… · 2026/9/24 15:31:35

奥赛一本通 1467 Radio Transmission
奥赛一本通 1467 Radio Transmission

1467 Radio Transmission 题目大意 给定一个字符串,求一个长度尽可能短的串,使得原先的串是这个短串重复若干次之后的子串。 知识要点 KMP 解题思路 首先,求解的这个短串一定可以是原串的前缀,如果不是前缀的话,将这个… · 2026/9/24 15:31:29

STM32无DAC怎么办?用PWM加RC滤波实现低成本模拟输出
STM32无DAC怎么办?用PWM加RC滤波实现低成本模拟输出

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/24 15:31:29

基于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

了解更多?预约专属演示

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

企业微信二维码