
1. 多任务学习里的“分赃不均”到底卡在哪做深度学习的人只要碰过多任务学习大概率都经历过这种憋屈模型结构设计得漂漂亮亮共享主干加多个任务头理论上应该互相促进、一起变强结果训练日志一拉出来某个任务loss唰唰往下掉另一个任务loss像被钉死了一样纹丝不动甚至越训越差。你调学习率、换优化器、加warmup折腾一圈发现——问题根本不在这些地方而是各个任务在反向传播时“抢梯度”抢得太难看。这就是GradNorm要解决的核心问题。它不是一个新的网络结构也不是什么玄学trick而是一种梯度归一化的动态损失平衡方法。说白了它干的事情就是在训练过程中实时观察每个任务的梯度状况然后动态调整每个任务loss的权重让所有任务以差不多的节奏往前跑而不是让某个“强势任务”把整个共享主干带偏。这篇文章适合谁看如果你正在做多任务学习手头有多个loss需要加权试过手动调权重但效果不稳定或者你用的是简单的loss相加、权重固定那这篇内容基本就是给你写的。我会把GradNorm的来龙去脉、核心原理、代码实现、参数选择、踩坑经验全部拆开讲尽量做到你读完就能往自己项目里搬。先给一个最直观的类比。想象一个团队里有两个员工一个能力强干活快一个能力弱干活慢。如果按固定比例给他们发奖金强的那个越来越有动力弱的那个越来越摆烂最后整个团队产出全靠强的那个人撑着。GradNorm的做法相当于每隔一段时间看一眼两个人的工作进度如果弱的那个落后太多就临时给他加钱、给强的那个减钱逼着两个人进度拉齐。最终目标是整个团队的总产出最大化而不是某一个人刷出漂亮KPI。2. GradNorm到底在算什么从梯度范数到动态权重2.1 为什么梯度范数能反映任务状态要理解GradNorm得先接受一个前提一个任务的梯度范数可以近似反映它当前的训练难度和收敛状态。梯度大说明这个任务离目标还远或者它的loss曲面还很陡梯度小说明它已经接近收敛或者它本身对参数的影响就很弱。在多任务学习里共享参数W上会收到来自各个任务的梯度。如果任务A的梯度范数远大于任务B那么更新时任务A的方向会主导W的变化任务B相当于被“淹没”了。GradNorm要做的就是给每个任务的loss乘上一个动态权重w_i(t)使得加权之后各个任务对共享参数的梯度范数趋于一个平衡状态。这里有个关键点GradNorm不是让所有任务的梯度范数完全相等而是让它们按照一个期望的相对速率去变化。具体来说它定义了一个目标每个任务的梯度范数应该正比于它的训练速率。训练速率快的任务梯度范数可以小一点训练速率慢的任务梯度范数应该大一点这样才能追上来。2.2 核心公式拆解别被数学符号吓到GradNorm的原始论文里公式看起来有点唬人但拆开看逻辑很清晰。对于每个任务i定义它在共享参数W上的梯度范数G_i(t) || ∇_W (w_i(t) * L_i(t)) ||注意这里w_i(t)是动态权重L_i(t)是任务i的原始loss。然后定义所有任务梯度范数的均值G_avg(t) mean(G_i(t))接着定义每个任务的相对训练速率。训练速率用loss下降的比例来衡量r_i(t) L_i(t) / L_i(0)这个比值越小说明这个任务训练得越快。然后定义期望的梯度范数G_target_i(t) G_avg(t) * [r_i(t)]^α其中α是一个超参数控制平衡的力度。α越大对训练慢的任务“照顾”越明显。最后GradNorm的损失函数是让实际梯度范数G_i(t)去逼近这个目标L_grad Σ_i | G_i(t) - G_target_i(t) |这个L_grad用来更新动态权重w_i(t)而原始的多任务loss用来更新网络参数。两者交替进行互不干扰。2.3 为什么用“训练速率”而不是直接看loss值这里有一个很容易踩的坑不同任务的loss量级可能差好几个数量级。比如一个回归任务MSE可能在10左右一个分类任务交叉熵可能在0.1左右。如果你直接比较loss绝对值那回归任务永远“看起来”更差权重会被无限放大。GradNorm用L_i(t)/L_i(0)这个比值相当于做了归一化把所有任务的起点都拉到1然后看谁下降得快。这样不同量级的任务之间就有了可比性。这个设计非常关键也是GradNorm比简单“按loss大小加权”更稳的原因。注意如果你的某个任务loss在训练初期就剧烈震荡r_i(t)会非常不稳定进而导致G_target_i剧烈波动。这种情况下建议先对loss做平滑或者等训练稳定几轮后再开启GradNorm。3. 手把手实现从零把GradNorm塞进你的训练循环3.1 准备工作确定共享参数和任务lossGradNorm只作用于共享参数。如果你的多任务模型里有些参数是任务独有的那些参数不参与GradNorm的权重计算。所以第一步是明确哪些层是共享的。通常来说底层的特征提取网络是共享的顶层的任务头是独立的。假设你有一个共享的encoder输出特征h然后两个任务头分别产生loss1和loss2。共享参数就是encoder里的所有可训练参数。你需要把这些参数收集起来后面计算梯度范数时只对这些参数求梯度。shared_params list(encoder.parameters()) task_losses [loss1, loss2]3.2 初始化动态权重和参考loss每个任务需要一个初始权重w_i通常初始化为1.0。同时记录每个任务在训练开始时的loss值L_i(0)作为归一化的基准。这个基准不需要每一步都更新一般用第一个batch或者前几个batch的平均值。w [torch.tensor(1.0, requires_gradTrue) for _ in task_losses] initial_losses [loss.detach().clone() for loss in task_losses]这里有个细节w必须是可学习的参数因为后面要用L_grad对w求梯度。但w不参与网络参数的更新它有自己的优化器。3.3 计算加权loss和梯度范数前向传播得到每个任务的原始loss后先计算加权lossweighted_losses [w[i] * task_losses[i] for i in range(len(task_losses))] total_loss sum(weighted_losses)然后对共享参数求梯度。这里要注意不能直接调用total_loss.backward()因为那样会累积梯度。正确做法是用torch.autograd.gradgrads torch.autograd.grad( total_loss, shared_params, retain_graphTrue, create_graphTrue )create_graphTrue是关键因为后面L_grad要对w求二阶梯度。retain_graphTrue是为了后面还能继续反向传播。计算每个任务的梯度范数时需要单独对每个加权loss求梯度G_i [] for i in range(len(task_losses)): grads_i torch.autograd.grad( weighted_losses[i], shared_params, retain_graphTrue, create_graphTrue ) norm_i torch.sqrt(sum(torch.sum(g**2) for g in grads_i)) G_i.append(norm_i)3.4 计算目标范数和GradNorm损失有了G_i之后计算G_avg和训练速率r_iG_avg torch.mean(torch.stack(G_i)) r_i [task_losses[i].detach() / initial_losses[i] for i in range(len(task_losses))] G_target [G_avg * (r_i[i] ** alpha) for i in range(len(task_losses))]然后L_grad就是L_grad sum(torch.abs(G_i[i] - G_target[i]) for i in range(len(task_losses)))3.5 更新动态权重和网络参数L_grad用来更新w网络参数用total_loss更新。两个优化器分开optimizer_w.zero_grad() L_grad.backward(retain_graphTrue) optimizer_w.step() optimizer_model.zero_grad() total_loss.backward() optimizer_model.step()注意顺序先更新w再更新网络参数。因为w的更新依赖于当前网络参数下的梯度范数如果先更新网络参数梯度范数就变了。实操心得w的优化器学习率不要设太大一般1e-3到1e-2之间。太大容易震荡太小则平衡效果出不来。我试过用Adam优化w学习率0.01大多数场景下比较稳。4. 参数选择与调优α、学习率和更新频率怎么定4.1 α的选择平衡力度旋钮α是GradNorm里最重要的超参数。α0时G_target就是G_avg所有任务的梯度范数被拉向同一个值平衡力度最强。α1时G_target正比于训练速率训练快的任务允许有更大的梯度范数平衡力度较弱。实际使用中α通常取0.5到1.5之间。我个人的经验是如果任务之间差异很大比如一个分类一个回归α取1.0左右比较合适如果任务比较相似α取0.5就够了。α太大比如超过2会导致训练慢的任务权重被过度放大反而引起震荡。4.2 动态权重的更新频率GradNorm不需要每一步都更新w。每一步都更新的话计算开销大而且w会抖动得很厉害。常见的做法是每N步更新一次wN取10到100之间。或者按epoch更新每个epoch结束时更新一次w。我一般用每50个iteration更新一次这样既能跟上训练节奏又不会太频繁。如果你的batch size很大可以适当减少更新间隔。4.3 初始loss的选取initial_losses决定了归一化的基准。如果只用一个batch的loss噪声会很大。建议用前10到50个batch的平均loss作为L_i(0)。另外如果某个任务在训练初期loss就下降得特别快r_i会迅速变小G_target也会变小这个任务的权重会被压低。这是符合预期的因为快任务不需要额外照顾。注意如果某个任务的loss在训练初期出现NaN或者Infinitial_losses会被污染后续所有计算都会出问题。建议在记录initial_losses之前先做一轮梯度裁剪和loss检查。4.4 和其他损失平衡方法的对比方法核心思路优点缺点固定权重人工设定w_i简单直接需要大量调参泛化差不确定性加权用可学习方差建模有概率解释方差容易坍缩GradNorm动态调整梯度范数自适应强理论清晰计算开销大超参敏感动态平均按loss下降速率调权实现简单对loss量级敏感GradNorm的优势在于它直接作用于梯度层面而不是loss层面。梯度才是真正决定参数更新的东西所以从梯度入手理论上更合理。代价就是每次更新w都要计算所有任务的梯度范数计算量大概增加30%到50%。5. 实战中遇到的坑与排查记录5.1 梯度范数计算报错“one of the variables needed for gradient computation has been modified”这个错误几乎每个第一次实现GradNorm的人都会遇到。原因是在计算G_i之后你又调用了total_loss.backward()而total_loss的计算图里包含了已经被修改过的变量。解决办法是确保所有autograd.grad调用都带retain_graphTrue并且在更新网络参数之前不要做任何in-place操作。如果还是报错检查一下你的模型里有没有BatchNorm或者Dropout。BatchNorm在训练模式下会更新running_mean和running_var这些是in-place操作会破坏计算图。建议在计算梯度范数时把模型设为eval模式或者用torch.no_grad()包住不需要梯度的部分。5.2 动态权重出现负数或者爆炸w_i在理论上是正数因为loss权重为负没有意义。但如果你用SGD更新w且L_grad的梯度方向不对w可能变成负数。解决办法有两个一是用Adam优化w它对符号变化更鲁棒二是每次更新后对w做clamp限制在0.01到100之间。w爆炸通常是因为α太大或者学习率太高。先检查α是不是超过了1.5再把w的学习率降到1e-3试试。5.3 某个任务权重一直涨其他任务权重一直跌这说明那个任务被判定为“训练太慢”GradNorm在拼命给它加权重。但如果它本身就是一个很难的任务权重涨到很大也追不上反而会拖累其他任务。这种情况下可以考虑给w设一个上限或者检查一下这个任务的loss是不是有问题比如标签噪声太大。我遇到过一次某个任务的标签里有大量错误标注导致它的loss根本降不下去GradNorm把它的权重加到了50多其他任务全被压死了。后来清洗了标签问题自然消失。所以GradNorm不是万能的数据质量永远是第一位的。5.4 训练前期GradNorm效果不明显GradNorm需要一定的训练步数才能积累出有意义的梯度范数统计。如果训练总步数只有几千步GradNorm可能还没来得及发挥作用就结束了。建议在训练步数超过1万步的场景下使用GradNorm否则用固定权重可能更划算。另外可以在训练前几百步先冻结w用初始权重跑一段时间等loss进入稳定下降阶段再开启GradNorm。这样能避免早期噪声对w的干扰。5.5 常见问题速查表现象可能原因排查方向解决办法计算图报错in-place操作检查BN/Dropout设eval模式或retain_graphw变负优化器不合适检查w更新方式换Adam加clampw爆炸α或lr太大打印w和G_i降α降lr某任务权重独大该任务太难或数据有问题看该任务loss曲线检查数据设w上限效果不如固定权重训练步数太少看总step数增加步数或延迟开启6. 一些进阶玩法和个人体会GradNorm本身是一个框架你可以根据具体场景做变体。比如你可以把G_target里的α做成动态的训练初期用大α强力平衡后期用小α让各任务自由收敛。也可以把梯度范数的计算从共享参数扩展到部分共享参数只对最底层的几层做平衡。还有一个很实用的技巧把GradNorm和梯度裁剪结合使用。先对每个任务的梯度做裁剪再计算范数这样能避免个别异常batch把w带偏。裁剪阈值一般设1.0到5.0之间看具体任务。我在实际项目里用GradNorm最大的感受是它确实能缓解任务之间的“霸凌”现象但前提是你的任务之间没有本质冲突。如果两个任务的标签互相矛盾比如同一个输入既要分类成A又要分类成B那GradNorm也救不了它只能平衡梯度不能解决标签冲突。这种情况下应该先做任务相关性分析把冲突任务拆开或者重新设计标签体系。另外GradNorm的计算开销在任务数量多的时候会很明显。如果你有10个以上的任务每个step都算所有任务的梯度范数训练速度可能直接减半。这时候可以考虑分组把相似的任务分到一组组内用固定权重组间用GradNorm。这样能在效果和效率之间取一个平衡。最后分享一个调试小技巧把每个任务的w_i和G_i都打到TensorBoard上观察它们随训练的变化曲线。如果w_i在某个点突然跳变往往对应着那个任务的loss出现了异常。这些曲线比单纯的total loss曲线信息量大得多能帮你快速定位问题。