☰
GradNorm:动态平衡多任务学习梯度权重的实战指南
2026/9/26 4:01:46 网站建设 项目流程

如果你同时训练过两个任务,多半遇到过这样的局面:一个任务的损失稳定下降,另一个却像没睡醒一样原地打转。这背后不是优化器的问题,而是多个任务在共享网络里“抢梯度”,谁的量级大谁就赢。GradNorm 是我从双任务训练里摸到的一套解法,它把“每个任务该分多少梯度”变成了网络自己在每个 step 都能调节的权重,而不是靠人肉调参。下面我会把原理、实现、排障都过一遍,文里的代码和参数是简化过的生产实践版本,适合正在做多任务学习、手动权重调得想骂人的朋友。

1. 先看病根:多任务训练里损失权重为什么不能用手拍

1.1 损失量级差只是表象

很多人一开始会想:把两个 loss 分别归一化到同一个量级不就行了吗?我试过,确实能解决一部分问题,但只是把最浅的一层遮住了。分类任务的交叉熵很容易飘到 3 到 5,回归任务的 Smooth L1 可能只有 0.1 到 0.5。你直接把两个 loss 加起来,回归任务基本等于不存在;你把回归任务乘上 100,分类任务又开始乱跳。这里的问题在于“量级差”不是一个固定常数,它随训练阶段、batch 组成、任务难度实时变化,你没有办法用一个静态系数把两个 loss 永久对齐。

真正麻烦的是动态变化。训练一开始,任务 A 的 loss 大、梯度也大,你给它调了一个比较小的权重;训练到中期,任务 A 收敛了,loss 降到和任务 B 差不多,但它的梯度范数可能仍然不小,原来的权重现在反而让 A 继续压制 B。手动调权重的本质是拿人眼去盯两条 loss 曲线,再凭经验改系数,而多任务训练里 loss 和梯度并不是同步变化的关系,所以盯着 loss 调权重经常南辕北辙。

1.2 梯度冲突才是真问题

就算两个 loss 量级一致,共享网络内部还是会打架。以我做过的一个语义分割加深度估计任务为例,两个 head 从同一个编码器拿特征,分割任务希望特征保留清晰的类别边界,深度估计任务希望特征对连续深度更敏感。这两个期望本身就有张力,如果分割任务回传的梯度范数比深度任务大一个数量级,那么编码器每个 step 的更新方向基本由分割任务说了算,深度任务虽然在反向传播,却相当于在一个不断被改动的特征空间里追一个移动靶。

所以我后来养成了一个习惯:不看 loss 量级,而是直接看共享层的梯度范数。GradNorm 的核心恰恰就是这一点——它不试图让两个 loss 相等,而是让两个任务在共享层上的梯度范数达到一个动态平衡。

2. GradNorm 的数学直觉:先量梯度,再反向调节权重

2.1 平衡的对象是梯度范数,不是损失

GradNorm 的设定很简单。总损失写作:

L(t) = w_1(t)L_1(t) + w_2(t)L_2(t) + ... + w_T(t)L_T(t)

其中w_i(t)是第 i 个任务的权重,会随时间变化。对于一组共享参数 W,第 i 个任务的梯度贡献范数是:

G_W^i(t) = || ∇_W ( w_i(t) L_i(t) ) ||_2

这里 W 不是整个网络的参数,而是你指定的某一个共享层。之所以要指定层而不是对整个网络算,是因为整个网络参数太多,梯度范数会被最后一层和第一层的 scale 差异搅浑,失去“任务间对比”的意义。实践中一般选最后一个共享主干层,或者最后一个共享 Transformer block。

2.2 相对训练速度 r_i 和整体目标 G

光看梯度范数还不够,GradNorm 还引入了一个关键变量:相对训练速度。定义:

r_i(t) = [ L_i(t) / L_i(0) ] / [ mean_j ( L_j(t) / L_j(0) ) ]

分母是全部任务相对初始损失的均值。r_i < 1说明任务 i 比平均速度学得快;r_i > 1说明它学得慢。这里的聪明之处是:loss 掉了多少,天然反映了任务当前的学习进度。一个已经快收敛的任务,它的梯度范数即使不小,也不应该继续拿走太多更新能量。

于是 GradNorm 给每个任务设一个目标梯度范数:

target_i(t) = G_bar(t) * r_i(t)^α

其中G_bar(t)是所有任务当前梯度范数的平均值,α是平衡强度。整体 GradNorm 损失就是实际梯度和目标梯度的绝对差之和:

L_grad(t) = Σ_i | G_W^i(t) - target_i(t) |

这个L_grad只用来更新w_i和α,不直接更新网络参数。网络参数仍然由原始的加权损失Σ w_i L_i更新。α的作用很直观:α=0时退化成让所有任务梯度范数向平均值看齐;α越大,慢任务被抬得越狠、快任务被压得越狠。

2.3 两任务数值例子:看看 L_grad 怎么推动权重

我用一个两任务例子把公式落地。假设两个任务的初始损失分别是L_1(0)=2.0、L_2(0)=0.8,训练到某个 step 时,当前损失是L_1(t)=1.8、L_2(t)=0.2,共享层上的梯度范数都是1.0,α=1.5。

任务初始损失当前损失相对损失比r_i当前梯度范数目标梯度范数
任务12.01.80.90 / 0.575 ≈ 1.5651.5651.01.0 × 1.565^1.5 ≈ 1.96
任务20.80.20.25 / 0.575 ≈ 0.4350.4351.01.0 × 0.435^1.5 ≈ 0.29

任务 2 的 loss 已经掉了 75%,明显学得更快,所以 GradNorm 的目标是把它在共享层上的梯度范数从 1.0 压到约 0.29,同时把任务 1 的梯度范数抬到约 1.96。L_grad = |1.0-1.96| + |1.0-0.29| ≈ 1.67,优化这个值会让w_2减小、w_1增大。这就是 GradNorm 的刹车逻辑:给快任务踩刹车,给慢任务踩油门。

3. 从零落地:一套可跑的 PyTorch 简化实现

3.1 共享层怎么选

GradNorm 对共享层选择很敏感。我踩过最直接的一个坑是选了网络里一个 LayerNorm 层作为共享层,结果梯度范数一直在零点几浮动,更新权重完全像随机游走。原因是 LayerNorm 的参数是 σ、β 这类归一化参数,梯度规模不大,且和任务特征的耦合方式与卷积/线性层不同。

我的建议是选一个“真正做特征变换”的层:ResNet 里选最后一个 block 的 3x3 卷积,Transformer 里选最后一个 block 的注意力输出线性层,UNet 里选 bottleneck 的卷积层。这个层要满足两个条件:一是它下面的参数被所有任务共享,二是它越靠近各任务 head 越好,因为梯度在这里已经充分混合了任务差异。

3.2 训练循环里的 GradNorm 放在哪一步

下面这段是简化版的 PyTorch 核心逻辑,主要展示训练循环里 GradNorm 的更新位置。生产环境建议在它基础上加梯度裁剪、缓存管理和更严谨的 detach 策略。

# 示意实现:两个任务的 GradNorm w = torch.tensor([1.0, 1.0], requires_grad=True) # 任务权重 alpha = torch.tensor(1.5, requires_grad=True) # 平衡强度 initial_losses = None # 第一个 step 记录下来的初始 loss,之后固定不动 def grad_norm_loss(losses, initial_losses, shared_params, w, alpha): norms = [] for i, loss in enumerate(losses): # 计算 L_i 对共享层参数的梯度贡献,乘上 w[i],并保留计算图 grads = torch.autograd.grad( loss, shared_params, grad_outputs=w[i], create_graph=True, retain_graph=True, ) # 一个层里有多个参数,各自求 L2 范数后求和 layer_norm = sum(torch.norm(g * p.data, 2) for g, p in zip(grads, shared_params)) norms.append(layer_norm) norms = torch.stack(norms) G_bar = norms.mean() rel = losses.detach() / initial_losses rel = rel / rel.mean() target = G_bar * rel ** alpha return (norms - target).abs().sum(), norms # 主优化器和 GradNorm 优化器分开 main_opt = torch.optim.Adam(model.parameters(), lr=1e-4) grad_norm_opt = torch.optim.SGD([w, alpha], lr=1e-3) for step, (x, y1, y2) in enumerate(train_loader): out = model(x) losses = torch.stack([loss1(out[0], y1), loss2(out[1], y2)]) if initial_losses is None: initial_losses = losses.detach() # 第一步:更新 GradNorm 自己的参数 w 和 alpha main_opt.zero_grad() lg, grads_norm = grad_norm_loss(losses, initial_losses, shared_params, w, alpha) lg.backward() grad_norm_opt.step() # 归一化权重:保证 w 之和为任务数 with torch.no_grad(): w.data = w.data / w.data.sum() * len(losses) # 第二步:用当前权重更新网络本身 main_opt.zero_grad() final_loss = (w.detach() * losses).sum() final_loss.backward() main_opt.step()

这段逻辑里最容易被忽略的是w.detach()。更新网络参数时,权重必须被当作常数使用,否则final_loss的梯度会流回w,把 GradNorm 和主优化器的更新搅在一起。另外create_graph=True是必须的,因为w要通过L_grad拿到梯度;但如果共享层参数非常多,这一行会把显存需求拉高一大截,后面会在排障部分展开讲。

3.3 三个影响成败的超参

GradNorm 需要调的超参数不多,但每一个都直接影响成败。我按重要性排一下:

第一是α的初始值。论文里推荐从1.5左右开始,实际范围我建议控制在0.5~2.0。α越大越激进,越快给快任务刹车,但也越容易让权重出现震荡。如果你发现某个任务刚开始训练就学不进去,多半是α偏大。

第二是权重学习率lr_w。官方实现一般用1e-3,但当你发现w在几百个 step 内从1.0冲到5.0再掉回0.1,就是lr_w太高。我一般先用1e-4起,稳定后再试着放大到1e-3。

第三是 GradNorm 的更新频率。论文默认每个 batch 都更新,实操中我更推荐每4~8个 step 更新一次。稀疏更新能有效抑制单 batch 噪声对w的冲击,尤其当某个任务本身带标签噪声时,效果差异非常明显。

4. 和其他动态加权重做法放一起比,GradNorm 的位置在哪里

4.1 常见方案速览与对比表

多任务动态加权重不是只有 GradNorm 一家。我整理过几个常见方案,按“更新依据”和“典型问题”做了个对比:

方案更新依据优点典型问题
固定等权无,一开始就定死最简单、零成本完全无视任务量级和收敛速度差异
不确定性加权每个任务损失的概率方差理论优雅,适合带噪声标签的场景需要额外估计方差,训练不稳时方差本身剧烈波动
DWA损失下降速度几乎没有额外计算量只看 loss,不看梯度,可能被 loss 量级误导
GradNorm共享层梯度范数 + 损失下降速度直接针对真正影响更新的梯度;可解释性强需要额外显存和计算;对共享层选择敏感
PCGrad / 冲突消解类梯度方向能解决“梯度方向相反”的硬冲突只改方向,不改量级,常需配合量级调整使用

从这个表能看出,GradNorm 站在一个很特殊的位置:它改的是“每个任务梯度贡献的量级”,而不是方向。很多多任务冲突其实是方向冲突——两个任务在共享层上的梯度夹角接近 180 度,这时候你再怎么拉大缩小梯度范数都没用,因为最终梯度是向量相加,方向相反的部分会互相抵消。这种情况我更推荐 PCGrad 那类冲突消解方法。

4.2 我的选型习惯:什么时候开 GradNorm,什么时候关

我自己的选型习惯是三个判断条件,同时满足以上两条才会上 GradNorm:任务数量在 2 到 5 个之间,再多的话单组共享层参数很难承担全局面貌;任务 head 共用大部分主干;任务损失曲线整体平滑,没有严重的跳变。

如果任务之间天然存在强烈对抗,比如一个任务要求特征尽量离散、一个任务要求特征尽量连续,我会先用 PCGrad 处理方向冲突,再叠加 GradNorm 处理量级失衡。只开 GradNorm 而不管方向冲突,训练后期会出现一种很诡异的现象:两个任务单独看都在下降,但共享层的梯度范数始终在合格范围内波动,实际的共享特征却变得越来越“四不像”。

反过来,如果任务数量超过 5 个,我倾向于不用 GradNorm。因为L_grad对每个任务生成一个权重,任务一多,权重之间的相对关系会更敏感,调参成本会翻倍。这时候先把任务分组,每组内部用 GradNorm,组之间用固定权重,反而更容易控制。

5. 排障实录:GradNorm 训练发散时的排查链路

5.1 第一步:把该看的指标全部打到日志里

GradNorm 训练发散时最忌讳直接猜。我刚开始用的时候没打日志,发散了整整两个晚上都在怀疑优化器,后来才发现是w爆炸了。所以只要你决定用 GradNorm,日志里至少要有这几项:每个任务的当前损失、每个任务在共享层上的梯度范数、当前w_i、当前α、G_bar。每个 logging step 打一次,不需要每个 batch 都打,但前 2000 个 step 最好密集一点。

有了这些日志,排查链路就是固定的:先看w_i有没有出现脉冲,再看α有没有撞边界,最后看梯度范数是不是周期性爆炸。按这个顺序走,绝大多数发散都能定位到具体环节。

5.2 五个翻车点,每个都有对应修法

我在多个项目里反复踩过同样的坑,列出来给后来者省时间:

第一是权重w出现脉冲或负值。常见原因是lr_w太大、或 batch 太小导致单 batch 梯度噪声过大。修法很简单:降低lr_w到1e-4,把 GradNorm 更新频率改成每 4 个 step 一次,并且在每次更新后强制做权重归一化和 clip,保证w在合理区间。

第二是α长期撞边界。α是幂指数,原则上不能为负,我一般会 clip 到0.1~3.0。如果训练初期α就一路冲到上限,说明两个任务的初始学习速度差距极大,这时候不是调α能解决的,该检查某个 head 是不是初始化有问题、某些任务是不是带上了过多噪声标签。先修数据,再修平衡策略。

第三是显存暴涨。这是create_graph=True的代价,因为它保留了二阶图。如果你用整个 encoder 做shared_params,显存直接翻倍都不奇怪。修法是只把最后一个共享 block 的少量参数传入shared_params,其他层的梯度不参与 GradNorm 计算。不要贪心,GradNorm 需要的是一个“代表性观测点”。

第四是任务 A 先收敛后,GradNorm 继续压制它。这种情况的表现是任务 A 的 loss 已经平了,w_A还在缓慢下降,因为它学得最快,被判定为“应该让路”。但有时候任务 A 不是学完了,而是进入了一个非常平缓的平台,再压低它的权重会让它永远停在平台期。修法是在验证 loss 连续不降后对 GradNorm 做 freeze,只保留当前w不变,后面纯靠主优化器继续训。

第五是 BatchNorm 和 GradNorm 打架。共享层带 BatchNorm 时,梯度范数在训练模式下会跟着每个 batch 的统计量抖动,w也会跟着抖。我的做法是 GradNorm 计算梯度时用 eval 模式下的 BatchNorm 统计量,或者直接把 BN 排除在shared_params之外。这看起来是小事,实际上抖动幅度能差好几倍。

排障时还有一些符号上的细节值得注意。比如torch.autograd.grad返回的梯度是“多个参数各自一份”,要对它们分别求范数再聚合;不同参数维度不同,有的参数矩阵有上万个元素,有的偏置只有几十个,直接求和会让范数被大矩阵主导。我一般先对每个张量做 L2 范数,再在整个层上取平均,这样更接近论文里单层梯度范数的语义。

我个人最后的经验是:GradNorm 不是一个装了就能跑的工具,它更像是一个训练健康度观测窗口。你一旦把每步的梯度范数和权重都记录下来,就算最后决定撤回 GradNorm、改用固定权重,也能比以前更清楚地知道多任务训练到底卡在哪里。如果你现在正被两个任务互相拖后腿搞到头秃,建议先把梯度日志打出来,再决定要不要引入它,多数情况下你会得到比盲目调 loss 系数更稳定的结果。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询