
跑过多任务模型的朋友大概率都有过这种体验分割、深度估计、法线估计三个任务拼在同一个backbone上loss直接相加。训练初期分割loss是个位数深度估计的loss却有十几法线估计只剩零点几——模型直接被深度任务带偏。手动调权重调了一星期好不容易跑出个能用的组合换个数据集又废了。这种多任务网络的Loss Balancing困境就是GradNorm要解决的问题。GradNorm全称Gradient Normalization梯度归一化是ICML 2018年提出的自适应损失平衡方法核心思路是把各任务的梯度范数拉回同一尺度同时兼顾各任务的学习速度差异。这篇文章不是论文复述而是站在实战角度把GradNorm的动机、算法拆解、PyTorch实现和我在几类任务上的实测经验完整讲一遍适合被多任务loss反复折腾、想系统解决任务权重问题的朋友。1. 先搞清楚多任务损失配平到底难在哪里1.1 损失直接相加并不是公平相加多任务学习理想很丰满一个模型多个头共享底层表征省参数省显存还能互相正则。但现实第一道坎就是损失合并。最简单也最常见的做法是L_total Σ L_i把所有任务loss直接加起来。问题是不同任务的loss天然不在一个量纲上——分类用交叉熵数值取决于类别数和模型置信度回归用MSE或L1数值取决于预测值与真值的量级有的任务用IoU loss有的用感知损失。直接相加时数值大的任务天然获得更大的反向传播梯度模型的学习重心就被它绑架了。我举个例子。之前做一个同时预测语义分割、深度、表面法线的模型backbone是ResNet。初始阶段深度任务的MSE在2到5之间波动分割的交叉熵在1附近法线任务用角度误差只有0.1到0.3。直接相加跑下来深度梯度的量级占绝对主导法线任务几乎分不到有效的优化信号。训了四十多个epoch分割和深度都像模像样了法线输出还是一团模糊。我一度以为是网络容量不够后来把三个loss各自归一化到同一量级再相加法线任务立刻活了过来。这就是loss balancing最粗暴的一层先解决尺度不均。1.2 梯度方向冲突这个隐藏问题就算你把三个loss的数值尺度调成差不多问题也还没完。任务之间还可能在梯度方向上打架。比如两个任务的梯度向量夹角超过90度更新共享参数时就互相抵消一部分甚至出现帮一个任务就必然损害另一个的跷跷板效应。任务数量超过三个以后这种冲突会快速累积让训练进入一种全局loss缓慢下降、个别任务始终拉垮的假收敛状态。这里我想特别区分两件事任务损失平衡和任务学习速度平衡。前者解决的是数值尺度让各任务在求和时贡献相近后者解决的是学习进度差——比如分类任务收敛很快回归任务收敛很慢。如果分类任务始终占着主导梯度等回归任务刚起步时共享特征已经被分类任务定形了想再调整就很难。这也是为什么直接给loss乘系数、固定权重的方法往往治标不治本。GradNorm的独到之处就是它把这两个维度同时纳入考虑而且不是预设一个静态权重而是在训练过程中动态调整。1.3 静态权重方案为什么会失效手动配平的方式五花八门网格搜索、先固定几个任务、按loss倒数的比例初始化权重、分阶段训练等等。这些方法在某些固定数据集上都能work但我用下来的感受是三个硬伤。第一搜索空间随任务数指数膨胀。三个任务还能靠蛮力试出来五个以上基本凭感觉毫无可解释性。第二最优权重本身是训练过程中的动态值不是一个常数。任务学到不同阶段它的梯度范数和loss下降速度都在变全周期用一个固定权重天然不可能最优。第三静态权重和网络初始化、batch构成、学习率调度强耦合A实验调好的值搬到B实验经常直接崩复现性很差。理解了这三点再看GradNorm的思路就会觉得顺理成章权重不该是人为拍脑袋的超参而应该是训练中随任务状态变化的可学习变量。2. GradNorm的核心机制盯着梯度范数做文章2.1 两个核心度量梯度范数与相对逆训练率GradNorm的关键洞察是与其盯着loss数值怎么缩放不如直接盯着梯度因为真正影响参数更新的是梯度。所有任务共享一个backbone每个任务对共享参数的影响力就可以用梯度范数量化。定义W为网络的共享层论文实验里通常取共享backbone的最后一层也就是各任务分支开始分裂之前的层。对任务i在训练时间t计算加权损失关于W的梯度范数G_W^(i)(t) ||∇_W (w_i(t) · L_i(t))||这里w_i(t)就是我们要学习并更新的任务权重。这个梯度范数越大说明任务i当前对共享层更新的影响越大。接下来需要两个统计量作为平衡标杆平均梯度范数Ḡ_W(t) (1/T) Σ_i G_W^(i)(t)代表所有任务的梯度影响均值。相对逆训练率r_i(t) (L_i(t) / L_i(0)) / ( (1/T) Σ_j (L_j(t) / L_j(0)) )。r_i(t)的含义值得仔细咂摸。分子是任务i当前loss相对它初始loss的比值表示我降到了初始的多少分母是所有任务这个比值的平均。如果任务i的loss降到了初始值的0.5倍而所有任务平均降到0.8倍那么r_i 0.5 / 0.8 0.625说明任务i学得比平均水平快。注意这里有个反直觉的点任务学得越快r_i反而越小因为它已经降得很低了。所以r_i小意味着我不那么缺梯度了可以把资源让给别人。2.2 α参数控制的是学习速度补偿强度有了上面两个量GradNorm为每个任务设定一个梯度范数目标G_W^(i)(t) → Ḡ_W(t) · [r_i(t)]^α当α 0时目标就是平均梯度范数所有任务的梯度范数被拉到同一水平论文里称为平衡梯度范数但不干预学习速度差异。当α 0时学过快的任务r_i 1目标值低于平均相当于主动压低它的梯度影响学得慢的任务r_i 1目标值高于平均相当于给它更强的梯度推动力去追赶。α本质上是一个恢复力强度参数控制我们多强硬地去补偿任务之间的学习速度差。α越大强迫所有任务学得一样快的力度越强但这也可能拖慢整体收敛速度因为所有任务都在互相迁就。α取0就是完全不干预速度取无穷大就是硬性让所有任务同步。论文默认值是1.5我后面实战部分会单独谈这个值的调法。这个α也是GradNorm论文里唯一需要手动设置的超参数其余全部自适应。2.3 一次前向两次反向的训练闭环GradNorm训练流程可以拆成两条并行的更新线。第一条线是正常的模型训练。用当前权重w_i对任务loss做加权求和L_total Σ w_i(t) · L_i(t)反向传播更新模型参数共享层和各任务分支。在这条线里w_i只是参与损失计算的前向系数模型参数的优化器不会更新它。第二条线专门更新权重w_i。构造一个GradNorm损失L_grad Σ_i |G_W^(i)(t) - Ḡ_W(t) · [r_i(t)]^α|这里的loss用L1距离对各任务的当前梯度范数与目标梯度范数之间的偏差求和。然后对w_i做一次梯度下降让当前梯度范数朝目标值靠拢。更新完w_i之后要注意还有一步全局归一化让所有权重的和等于任务数T即Σ_i w_i(t1) T。做法是w_i w_i / (Σ_j w_j) * T。这一步的目的是防止权重绝对尺度漂移保证总损失的量级不会因为权重整体放大缩小而剧烈变化。一个容易被忽略的要点GradNorm做梯度回传更新w时不能影响模型参数。实现上要么把w挂在独立的优化器上要么用torch.autograd.grad对w单独求梯度并stop住对模型的传播。如果你不加处理直接对gradnorm loss做backward()梯度会顺着计算图流进模型参数主优化器一更新就把模型训歪了。这个细节我后面实现部分会再强调。3. 从公式到PyTorch我落地的GradNorm实现3.1 共享层W的选法与梯度提取在PyTorch里复现GradNorm第一步也是最容易出岔子的一步是确定共享层W。我推荐取所有任务分支共享的最后一层参数。以常见的ResNet做多任务encoder为例W就是layer4之前的那部分或者干脆取backbone最后一个stage的参数集合。如果任务头是从某个特征层开始分叉的W就是最后一段公共路径。获取梯度有几种做法方案A对每个任务的加权loss单独调backward(retain_graphTrue)从param.grad里取梯度算完范数后清空梯度。方案B在W层注册反向传播hook在hook里直接拿到梯度避免多次backward导致的计算图问题。方案C通过torch.autograd.grad(weighted_loss_i, W_params)直接得到梯度不依赖.grad。我实测下来torch.autograd.grad最干净不容易造成梯度残留问题。算梯度范数时把W里所有参数的梯度展平后计算L2范数即可。注意不是只看单个参数的grad而是整个W参数集合的综合范数。3.2 权重更新、重归一化的正确顺序一份可以直接参考的核心代码逻辑如下# 假设 # model: 多任务网络T个任务头 # W: 共享backbone最后一层参数集合 # w: 形状为 (T,) 的可训练权重初始化为 [1.0, 1.0, ...] # alpha: 超参数默认1.5 # initial_losses: 训练开始时估计的各任务初始lossshape (T,) # optimizer: 主模型优化器 # w_optimizer: 只更新w的优化器实测常选SGD for step, (inputs, targets) in enumerate(train_loader): outputs model(inputs) losses torch.stack([ criterion[i](outputs[i], targets[i]) for i in range(T) ]) # shape (T,)注意在更新w之前计算并保留 # ----- GradNorm更新w ----- w_optimizer.zero_grad() grad_norms [] for i in range(T): weighted_loss_i w[i] * losses[i] # 直接对w_i * L_i 求关于W参数的梯度 grads_i torch.autograd.grad(weighted_loss_i, W_params, retain_graphTrue) grad_norm_i torch.sqrt(sum((g ** 2).sum() for g in grads_i)) grad_norms.append(grad_norm_i.detach()) grad_norms torch.stack(grad_norms) # shape (T,) G_bar grad_norms.mean() loss_ratio losses / initial_losses.detach() rate loss_ratio / loss_ratio.mean() target_norms (G_bar * (rate ** alpha)).detach() gradnorm_loss torch.abs(grad_norms - target_norms).sum() gnorm_grads torch.autograd.grad(gradnorm_loss, w)[0] w.grad gnorm_grads w_optimizer.step() # 重归一化所有权重之和等于任务数T with torch.no_grad(): w.data w.data / w.data.sum() * T # ----- 主模型更新 ----- optimizer.zero_grad() total_loss torch.sum(w.detach() * losses) total_loss.backward() optimizer.step()这段代码有个顺序是刻意的losses必须在更新w之前算好并保留因为更新w后重新forward会导致状态不同步目标梯度范数target_norms要detach()因为它是标杆不应该把梯度传给w以外的部分主模型更新时必须用w.detach()否则主模型优化器也会更新w造成双重更新。3.3 容易翻车的几个实现细节我在代码实现过程中踩过几个值得提醒的坑。第一个坑是初始loss的估计。L_i(0)如果直接用第一个batch的loss受随机初始化影响噪声很大特别是回归任务第一个batch的loss可能比稳定值高出几倍导致r_i在训练前几百步里严重失真。我后来改用训练开始阶段前几十个batch的均值或者用一个平滑的EMA替代权重的更新会稳定很多。第二个坑是w的优化器选择。论文用的SGD带小momentum学习率设置在0.025附近。我试过用Adam更新w并不是不行但Adam会加速w的漂移后期可能出现某些任务权重被压低到接近0的情况。用SGD加少量momentum权重变化更温和。第三个坑是梯度范数的计算粒度。不要只取最终任务头的最后全连接层梯度而是取共享参数W的梯度。如果W选得太靠近输出层GradNorm反映的基本就是分类头内部的事跟共享表征平衡关系不大了。 提示如果训练过程中某个任务权重出现“断崖式下跌”优先检查初始loss估计是否平滑、梯度范数是否加了数值下限epsilon以及重归一化是否放在了w更新之后。4. 实测记录三组任务下的表现与调参经验4.1 分类回归混合多任务收益最明显我最早测试GradNorm的场景是图像分类 回归的组合多任务网络。一个分支做场景分类另一个分支做深度估计。这种组合天然难配平因为交叉熵和MSE量级完全不在一个世界。手动调权重调了很久总是顾此失彼分类权重调大一点深度就开始飘深度权重调大一点分类精度就往下掉。上了GradNorm之后我没有手动设任何固定权重让w从1.0起步自动学。前几十个batch就能明显看到w的变化趋势分类任务学得快w被逐步压低深度任务学得慢w被抬高。最终稳定下来大约分类权重0.7、深度权重1.3左右两个任务的最终精度都超过了我在手动调参时找到的最优静态权重组合。这里的关键原因是GradNorm动态调节了不同训练阶段的任务权重而不是用一个固定比例迁就全程。4.2 目标检测多分支损失效果稳定但没有惊喜第二个场景是目标检测里常见的三损失结构分类损失、框回归损失、目标性损失。这类任务其实相对成熟很多训练框架里已经给了经验性权重比如YOLO系列固定用坐标权重高于分类权重。我试着把GradNorm套进去发现它能自动收敛到与经验权重接近的比例整个训练的收敛曲线比较平滑最终mAP和手工配平的基线差不多略高一点点。这个结果其实很有价值当任务间loss尺度差距不过分悬殊、已有比较合理的经验权重时GradNorm不会带来破坏性影响它能把人从三个权重怎么调里解放出来。有一点要注意目标检测的训练通常有很复杂的正负样本匹配策略不同阶段的loss统计特性变化大GradNorm的权重更新如果频率太高会跟着batch噪声波动。我给w的更新降低了频率每8个batch更新一次稳定性好很多。4.3 生成式任务多目标损失需要额外的梯度平滑第三个尝试是图像生成模型的多目标损失比如同时优化感知损失和对抗损失。这个场景有点特殊因为对抗损失的梯度统计特性极不稳定判别器强弱变化都会造成生成器梯度范数剧烈波动。直接上GradNormw会被噪声带着乱跳训练后期甚至出现两个任务权重周期性震荡的情况。我的处理办法是给梯度范数加EMA平滑再喂给GradNorm。计算公式变成smoothed_grad_norm 0.9 * smoothed_grad_norm 0.1 * current_grad_norm这样一来w的更新趋势稳定了震荡也消失了。另一个思路是降低α从默认1.5调到0.8降低速度补偿的强度也能让权重稳定下来。对于损失方差大的任务组合我建议优先调这两个旋钮。4.4 关于α和w初始化的经验规律α是我唯一需要手动扫的超参数扫过的范围是0到3。我的经验是α 0只做梯度范数拉齐不做学习速度补偿。适合任务本身收敛速度接近的场景w比较稳定。α 1.0 到 2.0适合分类和回归混合、任务学习速度差异明显的场景。我遇到的多数任务组合在1.2到1.8之间效果最好。α 2.5补偿力度过强慢任务被过度扶持整体收敛速度明显变慢多数情况下不划算。w的初始化我一直用的是全1.0没有踩到明显问题。如果任务loss初始量级差异特别大也可以按1 / L_i(0)归一化初始化能缩短w预热的周期。但注意GradNorm的更新第一步就会矫正权重所以初始化的影响在训练充分时几乎可以忽略。另外一个调参技巧w的更新频率和学习率要有联动。如果w更新频率调低比如每N个batch一次w的学习率可以相应调大一点。我自己通常把w的SGD学习率设在0.01到0.05之间再用余弦annealing衰减效果比较稳。 提示当任务数量较多5个以上时梯度范数宜采用对数坐标再计算避免单个任务的梯度范数过大把平均梯度范数整体抬高导致其他任务目标值失真。5. GradNorm的边界、开销与值得跟进的变体5.1 计算开销与梯度噪声GradNorm的额外计算开销主要来自对共享层梯度的多次提取。任务数为T时需要对W参数重复求T次梯度计算量大约是T倍的反向传播开销。T为2到3时几乎可以忽略但到了10个以上就值得权衡了。一些工程化做法是每N个batch才更新一次w把GradNorm的计算摊薄只要N不是太大w依然能跟上训练趋势。梯度噪声是另一个要注意的点。Batch越小、任务loss方差越大单次梯度范数的估计越不稳定w更新就越抖。建议batch size不低于32如果任务头本身loss就波动剧烈就必须用EMA平滑梯度范数。还有一个实操细节不同任务如果用不同的损失函数它们的梯度分布形状差异可能很大有的梯度集中在少数参数上有的分布很均匀。这时L2范数只能刻画总量无法体现分布差异。如果发现GradNorm效果不明显可以尝试改用梯度向量直接拼起来算范数或者对每个参数单独做归一化再聚合。5.2 不适合用GradNorm的场景识别不是所有多任务问题都适合GradNorm我用下来觉得有几类场景要谨慎。第一任务之间本身没有共享参数或共享层很少。如果两个任务网络结构上基本独立GradNorm在共享层上计算的梯度范数意义不大。第二某个任务的loss函数在训练中会发生非平滑变化比如任务头结构动态变化、样本权重动态变化梯度范数会出现跳变w更新易失稳。第三当任务权重本身的含义非常重要、必须符合业务解释时自适应权重难以审计这时候宁可手动配平。另一个被很多文章忽略的点是GradNorm解决的是共享层梯度平衡它不会解决任务分支内部头部的优化问题。如果某个任务头自己就训不动GradNorm会不断抬高它的权重但可能只是把它的梯度噪声放大对最终任务精度没有帮助。遇到这种情况先排查任务头结构和损失函数设计是否合理再考虑用GradNorm。5.3 与不确定性加权、DWA等方法的适用性对比做loss balancing的方法不止GradNorm一种我在实际项目中横向对比过几个主流方案这里整理成一张表方便参考方法核心思路超参数优点不足Uniform各任务loss直接相加无简单对loss量级敏感任务间冲突时偏差大GradNorm拉齐共享层梯度范数 补偿学习速度差α原理清晰适配任务进度差异额外计算量对梯度噪声敏感Uncertainty Weighting用同方差不确定性可学习噪声加权各任务噪声初始值有贝叶斯解释常用于分类回归混合权重量级容易漂移需要正则DWA按loss下降率的指数平均分配权重T温度参数实现极简无需额外梯度只关注loss变化率不含梯度方向信息MGDA多目标最优化求帕累托方向无理论上最优解计算复杂度随任务数上升快工程落地少对比下来GradNorm的优势在于它直接作用于梯度而非只作用于loss数值能感知到任务对共享参数的真实影响。Uncertainty Weighting实现简单、在分类回归组合中表现也好但它基于高斯噪声假设对非平滑任务不友好。DWA和GradNorm一样是动态权重不过DWA只看loss下降率不校正梯度尺度当两个任务loss量级差异大时依然会偏向前者。MGDA理论上最漂亮但大规模网络里求多任务梯度组合的帕累托方向计算量和稳定性都不容易控制。5.4 我后续会尝试的几个改进方向GradNorm还有不少扩展玩法我列几个我觉得值得跟踪的方向。一是把相对逆训练率里的初始loss改成滑动基线比如用过去K个epoch的平均loss来定义学习速度这样r_i对训练中后期更敏感也减少对初始状态的依赖。二是把单层W的梯度范数扩展成多层共享特征的加权组合避免某一层梯度小但其他层梯度大的情况被淹没。三是把w加上约束范围比如限制在[0.2, 5]之间防止某个任务权重被压到接近0后彻底失联。这些改动都不复杂工程上很容易做效果在某些场景下比原版更稳。还有一点我想提醒GradNorm和很多训练技巧不是天然兼容的。它会动态改变各任务的梯度尺度所以对学习率warmup和梯度裁剪比较敏感。我自己习惯在GradNorm启用前先跑十几步纯Uniform热热身让batch norm统计量稳定后再让w参与更新整体训练会更省心。如果你发现上了GradNorm反而比基线差先不要急着调α检查一下是不是和warmup、梯度裁剪冲突了。多任务loss balancing没有银弹。GradNorm是我目前愿意默认首选的自适应方案因为它把权重是静态超参这个思维定式打破了让网络自己在训练中协调任务关系。但实际落地时还是要根据任务组合、loss特性、工程约束去做适配。上面这套经验是我在几次项目里踩坑换来的希望对你少走弯路有帮助。