ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

GradNorm:动态平衡多任务学习梯度权重的实战指南

GradNorm:动态平衡多任务学习梯度权重的实战指南 如果你同时训练过两个任务多半遇到过这样的局面一个任务的损失稳定下降另一个却像没睡醒一样原地打转。这背后不是优化器的问题而是多个任务在共享网络里“抢梯度”谁的量级大谁就赢。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_gradTrue) # 任务权重 alpha torch.tensor(1.5, requires_gradTrue) # 平衡强度 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_outputsw[i], create_graphTrue, retain_graphTrue, ) # 一个层里有多个参数各自求 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(), lr1e-4) grad_norm_opt torch.optim.SGD([w, alpha], lr1e-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_graphTrue是必须的因为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_graphTrue的代价因为它保留了二阶图。如果你用整个 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 系数更稳定的结果。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进