
做过多任务学习的朋友应该都有被损失权重支配的恐惧。一个网络同时学检测、分割、深度估计三个任务三个loss的量级天差地别训练起来总是某个任务一边倒。GradNormGradient Normalization梯度归一化就是专门解决这个问题的它通过自适应地调整每个任务的损失权重让多任务网络的梯度保持平衡。这篇文章我会从原理到实现把这篇论文的细节完整过一遍包括PyTorch复现的代码和踩过的坑适合正在做深度多任务网络、推荐系统多目标建模或者任何多loss训练任务的读者。1. 为什么多任务训练总是一边倒1.1 一个典型的三任务训练现场先看一个最实际的例子。假设你在做一个自动驾驶感知模型一个共享的ResNet编码器后面接三个头2D检测、语义分割、单目深度估计。三个任务各自的loss长这样检测分支用的是Smooth L1加分类交叉熵数值通常在1到10之间分割分支是像素级交叉熵稳定期大概在0.5到1.5深度估计用的是尺度不变的回归loss初始阶段轻松破百训练中期也可能还在几十附近徘徊。如果你把三个loss直接相加会发生什么深度估计任务的梯度会占据绝对主导共享编码器的特征全部朝着“拟合深度”的方向更新检测和分割的任务头只能在夹缝中勉强学习。这是我实际调过模型之后的体会共享层学到的特征基本是深度任务形状的检测和分割的性能烂到没法看。1.2 手工调权是条死路有人会说给每个任务设个权重不就行了比如总loss等于 λ1 * L_det λ2 * L_seg λ3 * L_depth三个λ慢慢调。问题在于你调的是“同一组固定的λ”但训练过程中每个任务的难度和收敛速度是动态变化的。训练初期深度任务loss下降很快这时候它需要的梯度比例其实可以小一些训练后期深度任务开始收敛变慢反而需要更强的梯度推力。固定权重根本无法适配这种动态变化。更麻烦的是三个任务还算能靠经验调如果是推荐系统里的多目标建模CTR、CVR、点击时长、完播率、互动率五六个目标叠加每个目标还对应不同的业务权重靠手调基本不现实。我在做多目标推荐模型时光是确认一组相对可用的权重组合就要来回跑三轮AB实验两周时间就耗在这上面了。1.3 症结在于我们盯错了对象手工调权痛苦的根源在于我们始终盯着loss的数值在调权重。但真正决定共享网络参数往哪个方向更新、更新多大幅度的是损失函数对共享层权重的梯度而不是loss本身。举个例子一个任务的loss是0.5但它对某个共享参数的梯度可能非常大另一个任务的loss是50梯度却可能很小。如果只看loss数值你会觉得该给第一个任务加大权重实际上它的梯度已经把共享特征推向极端了。GradNorm这篇2018年的论文标题是 Gradient Normalization for Adaptive Loss Balancing in Deep Multitask Networks就是把这个矛盾直接摆到台面上与其在loss空间里调权不如到梯度空间里调权。它用一组可学习的任务权重让每个任务在共享层上的梯度范数去匹配一个自适应目标从而实现在整个训练过程中动态平衡。这也是“梯度归一化”这个名字的由来。2. GradNorm的核心机制从损失权重到梯度范数2.1 任务梯度范数的定义为了讲清楚原理我们先把符号约定好。设共享网络层的参数为 W任务 i 的损失为 L_i任务权重为 w_i。所谓任务梯度范数就是任务 i 的加权损失对共享层参数的梯度的L2范数G_W^(i) ||∇_W (w_i · L_i)|| w_i · ||∇_W L_i||因为 w_i 是一个正标量它可以被提到范数外面实测中要防止 w_i 为负数一律做截断或归一化。这个 G_W^(i) 描述了一件事在当前这一步任务 i 到底想以多大的“力度”去推动共享层参数。如果某个任务的这个范数长期远大于其他任务共享层就会被它单向侵占。2.2 目标梯度范数不是拉平而是按难度分配很多人第一次看GradNorm会误以为它要让所有任务的梯度范数相等。其实不是。如果所有任务梯度范数强行拉平反而会掩盖任务难度的差异。有些任务天生难需要网络投入更多容量有些任务已经收敛得差不多就不该再占据太多梯度资源。GradNorm的做法是先求所有任务当前梯度范数的平均值 G_W(t)再根据每个任务的“相对训练速度”对平均值做一个缩放得到该任务的目标梯度范数。训练速度越慢的任务目标范数会被放大迫使网络给它更多梯度资源训练速度快的任务目标范数缩小让出资源。2.3 相对训练速度 r_i(t) 怎么算任务的快慢不能直接看loss的绝对值得看它相对自己初始状态的下降比例。定义 loss ratioL̃_i(t) L_i(t) / L_i(0)也就是任务i当前loss与初始loss的比值。初始时这个比值是1训练一段时间后变成0.5就说明loss降了一半。但每个任务收敛速度不一样还要把这个比值和所有任务的平均比值做对比r_i(t) L̃_i(t) / ( (1/T) · Σ_j L̃_j(t) )这里的 r_i(t) 就是任务i的相对训练速度。如果 r_i(t) 1说明任务i比所有任务的平均水平下降得慢需要加大它的目标梯度范数如果 r_i(t) 1说明任务i已经降得很快了可以适当减小它的目标梯度资源。2.4 不对称参数α的作用直接拿 r_i(t) 乘到平均梯度范数上存在一个问题慢任务的“补偿”会被放得过大而快任务被压得过狠导致训练来回震荡。论文引入了不对称参数 α默认取1.5目标梯度范数写作G_W(t) · r_i(t)^αα 0 时退化为所有任务的目标梯度范数都等于平均值也就是完全拉平α 0 时才会对慢任务做额外补偿α越大补偿越激进。这个参数相当于一个“难度响应强度”旋钮实践里我不是每次都按默认1.5来的后面会讲到怎么调。2.5 GradNorm的优化目标把上面几块拼起来GradNorm对任务权重 w_i 的优化目标就是一个简洁的L1 lossL_grad Σ_i | G_W^(i)(t) - G_W(t) · r_i(t)^α |对 w_i 求梯度并更新就能让每个任务的梯度范数逐步靠近自己的目标值。注意这个 L_grad 只更新任务权重 w_i不直接更新网络参数 W。网络参数还是由原始的总loss来更新。3. 完整算法流程与PyTorch复现3.1 论文的算法步骤逐条解读论文Algorithm 1给出的流程可以拆成以下几个动作初始化网络参数 W任务权重 w_i 全部设为1记录每个任务第一个batch的初始loss L_i(0)。前向计算所有任务的loss L_i(t)。构造当前总loss L_total Σ w_i(t) · L_i(t)。分别计算每个任务对共享层参数的梯度范数 G_W^(i)(t)。计算平均梯度范数 G_W(t) 和相对训练速度 r_i(t)。计算GradNorm loss对 w_i 求梯度并更新 w_i。对 w_i 做归一化保持所有权重之和等于任务数T。用 L_total 的梯度去更新网络参数 W。这个顺序里有一个容易误解的地方第4步算梯度范数和第8步更新网络参数都是对共享层参数求梯度但两者目的不同。梯度范数只被用来算GradNorm loss真正更新网络的是L_total的梯度。两者是解耦的。3.2 为什么任务权重更新和网络参数更新必须分开这是我在第一次实现时踩过的大坑。如果直接把 w_i 塞进网络参数里一起做反向传播会发生什么总loss对 w_i 的梯度正好就是 loss_i 本身那么优化器为了让总loss最小会把所有 w_i 无限往0压任务权重很快退化到没有意义。GradNorm真正需要的 w_i 梯度不是来自总loss而是来自“梯度范数与目标值的差距”这个目标。所以 w_i 必须用独立的一套更新逻辑和网络参数完全隔离。更严谨地讲L_grad 对 w_i 的梯度依赖的是加权梯度范数 G_W^(i) 中 w_i 的系数作用而不是总loss对w_i的导数。这个区分如果没做对整个算法就变味了。3.3 PyTorch实现一个可直接改的模板下面是我在实际项目里使用过的GradNorm核心训练循环模板做了简化但保留了关键细节。场景假设是你已经有一个多任务模型shared_params是共享编码器的参数列表task_losses是各任务loss的列表。import torch import torch.nn as nn class MultitaskModel(nn.Module): def __init__(self, shared_encoder, task_heads): super().__init__() self.shared shared_encoder self.heads nn.ModuleList(task_heads) def shared_parameters(self): return list(self.shared.parameters()) # 训练超参数 num_tasks 3 alpha 1.5 # GradNorm 不对称参数 grad_norm_lr 0.025 # GradNorm 的权重更新学习率 base_lr 1e-4 # 网络参数学习率 # 任务权重初始化为1作为Parameter方便原地更新 task_weights nn.Parameter(torch.ones(num_tasks)) model MultitaskModel(shared_encoder, task_heads).cuda() optimizer torch.optim.Adam(model.parameters(), lrbase_lr) # 记录初始loss用于计算loss ratio init_losses None for step, batch in enumerate(train_loader): x batch[input].cuda() targets batch[targets] # 各任务标签 outputs model(x) task_losses [] for i in range(num_tasks): loss_i loss_fns[i](outputs[i], targets[i]) task_losses.append(loss_i) # 第一次迭代记录初始loss if init_losses is None: init_losses [l.item() for l in task_losses] # 1. 计算每个任务原始loss对共享层参数的梯度范数 raw_grad_norms [] for loss_i in task_losses: grads torch.autograd.grad( loss_i, model.shared_parameters(), retain_graphTrue, create_graphFalse ) flat torch.cat([g.reshape(-1) for g in grads]) raw_grad_norms.append(torch.norm(flat)) raw_grad_norms torch.stack(raw_grad_norms) # 2. 构造加权梯度范数对w_i可导 weighted_grad_norms task_weights * raw_grad_norms G_W weighted_grad_norms.mean() # 3. 计算相对训练速度 r_i current_losses torch.tensor([l.item() for l in task_losses]) loss_ratio current_losses / torch.tensor(init_losses) r_i loss_ratio / loss_ratio.mean() # 4. 计算GradNorm loss target_norms G_W.detach() * (r_i ** alpha) grad_norm_loss torch.mean(torch.abs(weighted_grad_norms - target_norms)) # 5. 更新任务权重用GradNorm的梯度沿独立学习率 grad_w torch.autograd.grad(grad_norm_loss, task_weights)[0] with torch.no_grad(): task_weights.data.sub_(grad_norm_lr * grad_w) # 归一化保持权重之和 num_tasks task_weights.data.mul_(num_tasks / task_weights.sum()) # 6. 更新网络参数注意用detach()的权重构造总loss total_loss sum( w.detach() * l for w, l in zip(task_weights, task_losses) ) optimizer.zero_grad() total_loss.backward() optimizer.step()这段代码里有几个细节我想单独说明一下。3.4 为什么网络权重更新时要用detach()total_loss里的task_weights.detach()非常重要。如果不加detachtotal_loss.backward()会顺带给task_weights填充梯度梯度值正好是各个任务的loss。而第5步我们已经用GradNorm的梯度更新了一次task_weights这个历史梯度如果不清掉到下一轮会被错误累积权重更新就乱套了。用detach()构造总loss让网络反向传播时不把任务权重当作叶子节点任务权重的梯度只由GradNorm loss产生两条更新通道彻底隔离。另外用w.detach() * l构造总loss网络参数的梯度仍然是w_i * ∇L_i。因为detach只截断了w的梯度通道w的数值仍然作为系数参与了链式求导。这一步我实测过结果和论文里的更新方式一致。3.5 共享层怎么选代码里model.shared_parameters()返回的是共享编码器全部参数。但实务中不一定要用全部共享层。论文建议取“最靠近任务头的那个共享子层”的梯度比如ResNet的最后一个stage。道理很简单共享层越深它的梯度越直接决定了任务头拿到的特征质量如果取整个ResNet所有参数前面浅层的梯度会被大量卷积的参数量稀释范数计算反而失真。我实际使用时是把共享编码器拆成早期层和晚期层用晚期层也就是和任务头接壤的那一段的梯度做GradNorm。比如对ResNet-50取 stage4对Transformer取最后一层encoder block。这个选择对效果影响不小后面常见问题里我会再展开。3.6 每个任务单独求梯度的性能代价代码里torch.autograd.grad对每个任务分别求了一次共享层梯度这意味着每步要多做 T 次反向传播retain_graphTrue 会在内存里保留计算图。任务数量一旦超过5个训练速度会明显下降。这是GradNorm绕不过去的代价。如果任务太多可以考虑每隔固定步数比如每20步才更新一次任务权重中间步数继续用当前权重训练网络性能损失可控效果也不会差太多。4. 关键超参数与调参经验4.1 GradNorm学习率不是越大越快GradNorm的权重更新学习率grad_norm_lr我建议默认从0.025开始。这个值和网络主学习率是独立的网络用1e-4或者1e-3都行但GradNorm的步子不宜迈得太大。任务权重w_i的数值通常在0到2之间波动如果grad_norm_lr给到0.1以上权重的震荡会直接反射到网络训练上一个任务可能在某几步突然拿到特别大的权重loss曲线出现尖刺。我测试过几个取值0.01偏慢权重调整跟不上任务难度的变化节奏0.025到0.05是比较稳的区域0.1以上基本都会出现训练震荡。如果你的网络本来就用了大学习率比如1e-3的Adam建议取0.01到0.025的下沿。4.2 不对称参数α任务差异越大α越大论文默认 α1.5。它控制的其实是“对慢任务的补偿强度”。当任务间难度差异很大时比如一个任务收敛到0.1另一个还在50r_i 的差距会非常大此时如果α也大慢任务的目标梯度范数会被顶得很高有可能把快任务直接饿死。我在一个分割深度估计的联合训练任务里把α从1.5改成0.8分割任务的mIoU反而提升了1.2个点。反过来如果任务间相对均衡α可以调大一点让算法对任务难度的响应更敏锐。实践建议是先跑50个step打印每个任务的加权梯度范数和目标范数如果看到慢任务的梯度范数被持续放大到离谱程度比如超过均值5倍以上就把α往下降如果各任务梯度范数几乎没有区别就把α往上提。4.3 初始权重的选择不要因为GradNorm能自适应就轻视初始权重。虽然论文里统一把w_i初始化为1但在任务loss量级差异极大的场景下前几步网络会被大loss任务带着猛冲GradNorm要花不少步数才能把权重拉回来。稳妥的做法是根据经验先设一个粗略的初始权重比如深度任务初始给0.01分类任务给1然后让GradNorm在这个基础上做动态修正。实测这会缩短训练前期的不稳定阶段对最终效果也有正面帮助。4.4 初始loss的记录时机代码里init_losses用的是第一个batch的loss。这个选择其实有点赌运气因为第一个batch的loss受输入数据分布影响很大。最好用前若干个batch的均值做初始化比如前50步的滑动平均。如果某个任务第一个batch恰好遇到一个异常样本loss异常偏高那么整个训练阶段的 loss ratio 都会被这个错误基线带偏。我在一个任务上踩过这个坑检测loss第一个batch爆炸到300初始化值被污染GradNorm把检测权重一路压到接近0花了几千步才恢复。改用前50步滑动平均后问题自然消失。5. 实验设计与效果观察5.1 论文里的验证场景GradNorm原论文在三个多任务场景上做了验证MultiMNIST多位数识别、CelebA属性分类、CityScapes语义分割深度估计。对比的baseline包括固定等权、手工调权、不确定性加权Uncertainty Weighting。结论是在多数任务组合下GradNorm能提升整体指标特别是CityScapes这种跨类型任务分类回归的场景提升幅度最明显。体验上MultiMNIST属于相对简单的多任务固定等权已经能取得不错的成绩GradNorm的收益主要体现在训练稳定性和收敛速度上CityScapes这种分类和回归混合的任务loss量级差异大GradNorm带来的收益就非常可观。5.2 我自己的复现结果我曾在一个人脸属性多任务项目上复现过GradNorm共享骨干加7个属性分类头。当时对抗baseline有两个等权相加和不确定性加权。固定等权训练下7个属性中4个能正常收敛另外3个发色、眼镜、表情因为样本量和难度差异指标明显落后不确定性加权会让个别难度高的属性被压得更弱。GradNorm训练后我最直观的感受是每个属性的梯度范数分布从“参差不齐”变成“有一定梯度差但整体可控”困难属性表情、发色的权重自动升高最终7个属性的宏观F1从等权的87.2提升到89.6。注意这个过程我完全没手动调过任务权重只是初始化给了大致范围。5.3 一个值得注意的现象在一个回归分类混合的任务里我观察到GradNorm会把回归任务的权重压得比预期低很多。原因是回归任务loss曲线的绝对数值在后期依然不小但它的梯度范数已经很小了——因为回归目标逐渐被拟合得很好梯度自然趋向于0。GradNorm看到的是“这个任务梯度已经很小”于是继续压低权重这其实没问题因为此时网络确实不需要再为回归任务付出太多更新力度。如果你发现某个回归任务指标在后期出现回退通常不是你权重给太低了而是网络容量问题这时候应该考虑加任务头复杂度而不是调权重。6. 实际使用中的常见问题与避坑6.1 任务权重直接崩到0这是GradNorm使用中最常遇到的问题之一。可能的原因有三个一是grad_norm_lr过大权重更新步子太大直接越过合理区间二是某任务的loss_ratio长期偏低r_i远小于1目标梯度范数被压到接近0权重也一路被压向0三是初始loss基线记录异常导致r_i计算失真。我推荐的排查顺序是先打印前50步的r_i和target_norms确认是不是某个任务的r_i长期低于0.5如果是考虑把α调小同时给task_weights做一个下界截断比如最小不能低于0.05实战中能避免很多次训练事故。注意截断后归一化时要把截断后的值纳入计算保证权重和仍然等于任务数。6.2 加载检查点后GradNorm状态丢失任务权重w_i、init_losses这些状态如果只在内存里而没有存到checkpoint中断续训时就会出问题init_losses为None会重新记录但这不是从零开始训练记录的“初始loss”其实是中断时刻的loss整个loss ratio的计算就错位了训练行为会突然变化。解决办法是在保存checkpoint时把task_weights.data和init_losses一起存进去加载时恢复。这个坑我遇到过不止一次每次都是恢复训练后指标莫名下降最后发现是GradNorm状态没跟上。6.3 共享层参数有共享Embedding的场景如果是推荐系统的多目标模型共享层往往包含Embedding。Embedding参数动辄几百万把它们全算进梯度范数会有问题Embedding中大量稀疏特征对应的参数梯度为零或极小会稀释整个共享层的梯度范数值导致GradNorm对“网络更新力度”的判断失真。我的建议是计算梯度范数时只取共享层的稠密部分比如MLP的权重不考虑Embedding。或者单独把Shared Bottom最后两层的参数作为GradNorm的观察对象。这样计算出的梯度范数更聚焦训练也更稳定。6.4 权重更新频率与网络更新频率不匹配GradNorm本身是每步更新但很多实践者包括我会改成每K步更新一次权重网络参数还是每步更新。这个做法对性能几乎没有影响同时能省掉每步多次反向传播的开销。需要注意的是K不宜过大我建议K在5到20之间超过50步的话权重反应太慢会出现任务失衡持续很久的情况。6.5 和其他多任务技巧共用GradNorm不是万能的它只管梯度范数平衡不管梯度方向冲突。任务A和任务B的梯度如果方向相反就算范数拉平了共享层依然会被两个任务扯来扯去。所以在实际项目中我经常把GradNorm和梯度冲突处理方法结合起来用PCGrad、CAGrad这类方法负责处理方向冲突GradNorm负责处理力度失衡两者互补。这里有个经验先用GradNorm让各任务梯度范数对齐再叠加梯度冲突消除手段比单纯用任何一种都稳。我在一个检测分割的模型上做过对照实验GradNorm单独提升1.8个点PCGrad单独提升1.2个点两者结合提升2.6个点确实是正向叠加的。7. 与其他自适应权重方法的横向对比7.1 不确定性加权Uncertainty WeightingKendall在2018年提出的方法用同方差不确定性σ来构造权重loss Σ_i (1 / 2σ_i²) L_i log σ_i。这个方法的优点是很轻量不需要算梯度范数缺点是不一定和梯度对齐。σ_i是通过优化出来的优化目标是最小化似然负对数而不是显式平衡梯度。在loss尺度方差大的场景下不确定性加权有时会给小loss任务分配过大的权重因为1/2σ²可能变得非常大导致训练不稳定。GradNorm从梯度范数出发物理意义更直接。7.2 DWADynamic Weight AverageDWA来自ICCV 2019思路比GradNorm简单直接用任务loss的变化率决定权重变化快的任务权重降低变化慢的任务权重升高。它不需要计算任何梯度开销几乎为零但效果通常弱于GradNorm因为它只看到了loss变化的相对速度看不到梯度本身的绝对力度。GradNorm同时考虑了loss下降速度通过r_i和当前梯度范数信息量更大。7.3 梯度冲突处理方法PCGrad/CAGrad注意区分这类方法和GradNorm解决的不是同一个问题。GradNorm解决的是“力的分配”PCGrad解决的是“力的方向”。多任务训练里两种问题往往同时存在一个任务梯度大但方向和其他任务冲突另一个任务梯度小但方向一致。只用PCGrad不做范数平衡可能还是被大梯度任务主导只用GradNorm不处理方向冲突共享参数还是会被撕裂。我个人的组合建议是先跑一个baseline打印每个任务梯度范数和两两夹角。如果范数差异大优先上GradNorm如果夹角超过90度的情况频繁出现再叠加PCGrad或CAGrad。两个都上的时候注意要先做梯度裁剪再交给优化器避免异常梯度带来的训练发散。7.4 什么时候别用GradNorm不是所有多任务场景都适合GradNorm。如果你的任务数量超过20个每个任务每步都要单独算梯度范数这个开销可能不可接受。另外如果任务之间共享层很浅各任务主要在独立分支上学习GradNorm能调整的空间就很小效果也有限。还有如果各任务loss本身已经做过严格归一化比如都限制在0到1之间且权重由业务规则硬性规定那就没必要引入GradNorm直接用固定权重更简单可控。最后分享一个我的使用习惯在把这套方法用了两年之后我现在每次设计多任务训练流程都会在一开始就把GradNorm当成默认组件放进去但做两处自定义一是把初始权重改成根据任务业务重要性和loss量级先验确定的粗略值二是每10步更新一次任务权重同时开启一个小的weight下限截断。这两个小改动让GradNorm在我的项目里稳定落地很少出现论文复现时常见的发散问题。另外我会坚持把每个step的w_i、梯度范数、r_i都打进日志里训练到一半时翻出来看往往能快速定位到是哪个任务在抢占资源。多任务训练的平衡问题没有银弹但GradNorm绝对是一个值得放进工具箱的通用解法。