
如果你训练过含多个任务的学习模型大概率撞过这个怪现象主任务精读涨得很漂亮辅助任务却怎么拖都拖不动或者反过来辅助任务把主任务带偏了。我一开始的处理方式也很原始手动调loss权重主任务不行就调高主任务辅助任务拉了就把辅助任务拉起来折腾几天效果随缘。直到重新读了GradNorm那篇工作才算把这个问题捋清楚——它解决的不是“权重设多少”而是“训练过程中权重怎么动态变”并且依据的是梯度范数这个更本质的信号。这篇文章我会从原理讲到PyTorch实现再把我在真实实验里踩过的坑一并写出来希望对正在处理多任务失衡的人有帮助。1. 多任务训练失衡固定loss权重为什么经常翻车1.1 同样的权重不同任务的“发言权”完全不对等多任务学习在工程里很常见比如检测里同时做回归框和分类模型既要判断目标类别又要精确回归坐标再比如推荐系统里同时优化点击率和转化率。做法一般是在最终loss处把各个任务的loss相加乘上各自的权重系数L_total w1 * L1 w2 * L2 ... wN * LN权重w通常是手动给的最常见的就是所有任务权重设为1或者按经验拍脑袋给几个值。问题就出在这里不同任务的loss在数值尺度上可能差好几个数量级。一个均方差回归loss可能动辄几十一个交叉熵分类loss可能只有零点几。你把两者都乘上权重1实际上这个模型在训练时95%的更新方向都被回归loss牵着走分类任务相当于一直在“陪跑”。这其实是很多初级玩家容易忽略的点权重是小但最终影响网络更新方向的是梯度范数不是loss数值大小。任务A的loss数值大并不意味着它“重要”只能说明它的梯度过大在大权重下直接淹没了其他任务。正确的姿势应该是看每个任务对共享参数产生的梯度有多大也就是梯度范数。多任务梯度相乘再相加谁梯度范数大谁就在抢模型容量。1.2 看绝对loss不可靠要用相对学习速率有人会说那我不用固定的我把两个任务的loss都归一化到同一量级再相加不行吗可以但这又是另一条弯路。loss绝对值小不代表任务学得差不多了。举个极端例子一个任务初始loss就是0.1现在已经降到0.09看起来数值小且稳定另一个任务初始loss是50现在降到30看起来还在大起大落。但你从相对变化去看第一个任务只下降了10%第二个任务却能下降40%。显然第二个任务还在快速学习阶段你更应该给它的梯度留更多“发言权”。所以GradNorm的思路是动态调整权重w核心参考信号不是任务loss的绝对大小而是各任务相对于自身的初始loss下降了多少。这个相对下降比例就是学习速率。如果一个任务学得快说明它进步明显需要适当降低它的梯度贡献如果一个任务学得慢说明它被压制了需要放大它的权重把梯度范数拉上来。看似很直觉但实现起来需要一套规范的定义和更新策略这就是GradNorm的核心价值。2. GradNorm核心机制用梯度范数衡量每个任务的学习状态2.1 算法想做的事情和三个关键量GradNorm要做的就是一件事在训练过程中动态调整各任务loss权重w让所有任务以“接近一致的速度”学习。这里的“速度”不是指loss步长收敛速度而是指每个任务对共享参数产生的梯度范数大小让它跟上整体的平均水平。先定义共享参数集合W通常是骨干网络的参数而不是任务各自的decoder头。训练步t上GradNorm会计算三个量第一个是总梯度范数记为G_W(t)它把当前所有加权loss的总和对W求梯度再取范数G_W(t) || ∇_W Σ_i ( w_i(t) * L_i(t) ) ||这就是网络当前整体更新力度的量度。第二个是每个任务单独的梯度范数记为G_W^(i)(t)把第i个任务单独产生的加权梯度在W上取范数G_W^(i)(t) || ∇_W ( w_i(t) * L_i(t) ) ||这个量告诉我们每个任务在总梯度里的实际“出力大小”。第三个是相对学习速率记为r_i(t)衡量第i个任务相对自己的初始loss下降了多少r_i(t) L_i(t) / L_i(0)初始步L_i(0)作为分母r_i越小说明这个任务降幅越大、学得越快r_i接近1说明它几乎没有进步、学得很慢。有了这三个量目标就非常清楚了希望每个任务的单任务梯度范数G_W^(i)都朝一个动态的目标值靠近这个目标不是固定值而是与全局总梯度范数和该任务的相对学习速率有关核心公式为target_i(t) G_W(t) * ( r_i(t) ^ α )然后构造一个GradNorm的辅助loss拉近目标与实际单任务梯度范数的距离L_grad(t) Σ_i | G_W^(i)(t) - target_i(t) |GradNorm的做法就是用这个辅助loss的梯度去更新w_i而非直接利用训练loss的梯度。这也意味着w的更新方向和网络参数更新方向是分开的两套逻辑。2.2 权重更新公式与归一化细节实现时需要注意一个之前容易忽略的点w更新完之后如果不加约束所有w可能在训练过程中漂移到某些极端值导致整体loss数值失控。所以GradNorm在每步更新w后还会做一次归一化保持所有任务权重的总和近似等于任务数N也就是说初始权重全部等于1时总权重和始终约为N。更新的逻辑大致如下前向算出每个任务的loss并取一个初始的L_i(0)用于计算r_i。对每个任务单独反向传播获取该任务在共享参数W上的梯度范数G_W^(i)。对总加权loss反向传播获取总梯度范数G_W。根据公式算出每个任务的target_i构造L_grad对w_i求梯度。用单独优化器更新w_i再把w_i做归一化保持权重之和稳定。有一个关键点反复踩到更新w时G_W^(i)和G_W都要从计算图中脱离当成常数处理。也就是说梯度的梯度不反向传导到网络参数上。GradNorm更新w只改变训练loss的加权比例它不应该改变网络参数的梯度方向以外的任何东西。如果这里不做detach操作loss_weights和模型参数就会被耦合进了奇怪的二阶梯度训练一多就会出现诡异震荡甚至崩溃。2.3 α超参数控制“拉一把”的力度公式里的α是一个需要自己调的平衡参数。α取0时target_i等于G_W(t)相当于让所有任务的单任务梯度范数向同一个全局平均对齐α取正数时学习慢的任务r_i大会被给到一个更大的target_i从而进一步放大权重、加速收敛。α越大对学习慢的任务的“鼓励”越强理论上收敛速度越快但盲目设大容易让模型在困难任务上过拟或者引起训练波动。论文里的实验常用α在0.5到1.5之间我自己实践时也基本落在这个区间。如果你手头的任务差距特别悬殊比如一个任务已经拟合得很好另一个任务还在挣扎可以尝试把α设得大一些如果两个任务本身体量差不多α0到0.5就够用。用太大的α配合不合适的初始学习率会在训练早期产生很大的loss抖动因为权重更新幅度过猛等于把一个本来就不稳定的任务用更大的权重继续放大。3. 从零复现GradNorm的PyTorch实践3.1 实验场景一个人为制造的难收敛任务纸上谈兵没有说服力我搭了一个很简单的多任务回归模型来做对比实验。共享网络是一个两层的MLP输入维度128隐藏层维度64输出层分裂成两个独立分支对应任务1和任务2。为了模拟失衡我故意让两个任务的噪声差异很大任务1的噪声很小很容易学任务2的噪声是任务1的5倍而且特征主要依赖输入的后半段学习难度明显更高。如果按固定等权去训练任务2通常会被任务1“带偏”收敛极慢。每个人都可以在同类合成数据上验证不用准备复杂数据集重点是看GradNorm的权重演化逻辑。3.2 第一步获取共享层梯度范数PyTorch实现里最核心的一个函数是“计算指定参数集合的梯度范数”。这里不要直接写动手算所有参数我封装了一个函数遍历backbone中所有共享参数把每个参数的梯度的平方累加起来再开根号def grad_norm_of_params(params): 计算一组参数的梯度L2范数 params: 需要统计的共享参数列表 total_square 0.0 for p in params: if p.grad is not None: g p.grad.detach() total_square g.pow(2).sum().item() return total_square ** 0.5这里有个容易忽略的细节不是所有层的梯度都参与计算。我强烈建议只统计共享参数W的梯度也就是backbone部分的参数不要统计各任务独享的decoder头。原因很好理解我们想平衡的是“共享表示层”上的任务竞争任务各自的head本来就是各算各的不存在竞争关系。你把head梯度也算进去会把任务自身头部大小的变化带进来干扰GradNorm的判断。3.3 第二步把GradNorm接入训练循环完整训练循环的核心结构如下我简化成只保留关键逻辑。首先定义共享参数列表和初始loss记录import torch import torch.nn as nn # 假设模型有两个任务分支task1_head, task2_head共享部分记为 backbone model TwoTaskModel() shared_params list(model.backbone.parameters()) # 两个任务的初始loss用于计算相对学习速率 initial_loss [None, None] # 权重初始化为1使用Parameter形式便于参与GradNorm的loss反传 w nn.Parameter(torch.ones(2)) # 给w单独开一个优化器不要塞进主优化器 w_optimizer torch.optim.SGD([w], lr0.01) alpha 1.5然后每一轮的训练逻辑我将整个流程拆成四步。第一步是记录初始loss第二步是分别反传拿每个任务的单任务梯度范数第三步是算总梯度范数第四步才是正常更新网络def train_step(x, labels): optimizer.zero_grad() # 1. 前向 pred1, pred2 model(x) loss1 mse_loss(pred1, labels[0]) loss2 mse_loss(pred2, labels[1]) # 记录初始loss只记录一次即可 if initial_loss[0] is None: initial_loss[0] loss1.item() initial_loss[1] loss2.item() # 2. 获取每个任务单独的梯度范数 task_grad_norms [] for i, task_loss in enumerate((loss1, loss2)): model.zero_grad() (w[i].detach() * task_loss).backward(retain_graphTrue) task_grad_norms.append(grad_norm_of_params(shared_params)) # 3. 获取总梯度范数 model.zero_grad() total_loss w[0] * loss1 w[1] * loss2 total_loss.backward(retain_graphTrue) total_grad_norm grad_norm_of_params(shared_params) # 4. 更新GradNorm权重w w_optimizer.zero_grad() grad_loss 0.0 for i, task_grad_norm in enumerate(task_grad_norms): r loss_i.item() / initial_loss[i] target total_grad_norm * (r ** alpha) grad_loss abs(task_grad_norm - target) grad_loss.backward() w_optimizer.step() # 归一化保持权重总和等于任务数 with torch.no_grad(): w.data.mul_(2.0 / w.sum().item()) # 5. 最后再用总loss反传一次更新网络参数 optimizer.step()等等这里写串了参数更新逻辑我需要调整一下上面步骤里已经对total_loss调用了backward在第5步我们调用optimizer.step()进行网络参数更新即可不能再调用total_loss.backward()。上面的写法目的只是展示计算顺序不要照抄错误版。下面是整理过没有注释歧义的关键版本# 计算梯度范数含单任务与总任务 task_grad_norms [] for i, task_loss in enumerate((loss1, loss2)): model.zero_grad() (w[i].detach() * task_loss).backward(retain_graphTrue) task_grad_norms.append(grad_norm_of_params(shared_params)) model.zero_grad() total_loss w[0] * loss1 w[1] * loss2 total_loss.backward() # 更新GradNorm的w w_optimizer.zero_grad() grad_loss torch.tensor(0.0) for i, task_loss in enumerate((loss1, loss2)): r task_loss.item() / initial_loss[i] target total_grad_norm * (r ** alpha) grad_loss torch.abs(task_grad_norms[i] - target) grad_loss.backward() w_optimizer.step() with torch.no_grad(): w.data w.data * 2.0 / w.data.sum() # 更新网络参数利用的是第2步里total_loss.backward()累积好的梯度 optimizer.step()注意total_grad_norm是在total_loss.backward()之后、还没执行optimizer.step()之前计算好的时机需要准确。3.4 两个容易翻车的实现细节我实际写这段代码时翻过两个车。第一个是单任务梯度范数收集完以后一定要把梯度彻底清零再用总loss做backward。否则单任务backward的梯度会残留在共享参数上和总loss的梯度混在一起算出的总梯度范数直接就错了。第二个是w不参与网络参数优化器的更新。如果图省事把w挂进model的parameters里再跟网络一起调用optimizer.step()w会同时被网络参数的优化器规则比如weight decay衰减更新很快w会退化成0或很小GradNorm彻底失效。所以w要单独立优化器而且要小心weight decay的干扰。如果你的主优化器用了weight decay那建议w优化器干脆就用SGD且不带weight decay保持w天然可控。另一个细节是初始学习率问题。GradNorm的w更新虽然用的是单独loss但它前面乘的是普通梯度下降w的更新跟模型当前loss大小强相关。如果模型初始loss特别大相应地单任务梯度范数也会很大导致grad_loss很大w更新步长变得很猛。可以考虑对grad_loss做一定缩放或者将w优化器学习率调低到0.001至0.01量级。我遇到一次w在第一个epoch就从[1,1]干到[0.3,1.7]的情况后来把w的学习率降到0.005才稳定下来。4. 跑完实验之后的对比结果4.1 收敛速度对比我在这个合成多任务场景上对比了三种方案固定等权、手动调优后的固定权重、GradNorm动态权重。固定等权不用说了任务2收敛极其缓慢到第100轮才勉强到合理水平手动调优的固定权重好一些但需要提前尝试多组权重组合成本高且只在固定场景有效。GradNorm的方案明显更稳。因为权重是动态的任务1一旦学得差不多了它的w会自动下降失去“压制力”给任务2让出一条路。从loss下降曲线来看同等epoch下任务2的相对loss在GradNorm下平均低20%到30%而且任务1最终精度也没有明显损失。这个结果符合GradNorm的设定逻辑——它不牺牲任何任务来强行平衡而是让慢任务追上来。4.2 loss权重随时间的变化曲线打印w的变化曲线非常有意思。初始两个任务权重都是1前几十步里任务1loss快速下降r_1不断变小而任务2还处在高亏损状态r_2更大因此target_2要比target_1大w_2被不断上调w_1逐渐下降到0.6左右。等到训练后半段任务2也开始明显下降r_2逐渐向r_1靠近w_1回升w_2回落两者趋向于某个稳定比例。这个曲线是GradNorm最直观的“体检报告”如果训练结束w没有明显变化说明两个任务本身就比较均衡没必要上GradNorm如果w长期集中在某个任务上则说明任务之间竞争确实剧烈GradNorm的调节是有意义的。我建议做这类实验尽量把w输出到日志观察它的动态非常有价值。4.3 几个让我意外的现象第一个意外GradNorm并不一定让总loss最低它让的是“坏任务”的loss显著下降而好任务只是轻微上升甚至不变。很多人看总loss反而被误导觉得GradNorm“没效果”。正确的评价方式是分别看每个任务的指标多任务场景本来就没有单一的总准确率指标。第二个意外α过大反而导致早停收益变小。我试过α1.5和α3前者在验证集上表现更稳后者训练后期任务2虽然吞掉了更多权重但泛化并不一定更好甚至在某些轮次出现过拟合。原因可能是训练后期慢任务的梯度已经被强行拉到很高模型在继续过拟合噪声而不是学特征。这也是为什么论文和社区实践都建议α不要过大的原因。第三个意外w的震荡比预想的频繁。如果你的w初始化不是1而是某个偏离较大的值比如1和0.1前几十步w会出现明显的上下波动因为grad_loss没有考虑当前r_i的平滑性。后来的经验是尽量让w初始都等于1利用GradNorm自己调节别想“给个先验值”。5. GradNorm的局限、踩坑记录和可用变体5.1 我实际踩过的坑第一个坑是共享参数选择的范围。一开始我把所有模型参数都丢进shared_params里包括两个任务的head参数结果梯度范数被某一层尺寸较大的参数主导GradNorm的行为变得很怪。后来把共享参数严格限定为backbone参数效果立刻恢复。共享参数集合的选择会直接影响GradNorm的平衡对象一定确保它是所有任务真正共享的那部分。第二个坑是梯度裁剪。很多训练流程里会做gradient clip如果你在计算GradNorm梯度范数之后、更新w之前执行clip实际上w的更新用的是裁剪前的数据梯度范数是真实值但网络更新是裁剪后的。这会引入一种微妙的不一致个别loss巨震的任务会被GradNorm放大而网络参数却没同步接受那么大更新。我的解决方案是在计算完grad_norm之后再进行clipGradNorm的辅助更新和主网络更新顺序不能交叉。第三个坑是跟学习率scheduler的冲突。如果你的主网络用了cosine或者step型scheduler来降低学习率GradNorm的w优化器最好也配一个同频率的scheduler否则w还在不停调整网络却已经进入“微调期”两头节奏不匹配训练会出现锁死现象表现为epoch后期w更新很大但loss纹丝不动。5.2 和不确定性权重法的关系与取舍很多人问GradNorm跟Kendall那篇不确定性加权有什么异同。不确定性加权同样做动态加权但它是从概率分布出发把每个任务的语言模型视为一个带噪声的观测权重由噪声参数softmax得到。它更适合那些任务loss服从一定概率分布回归问题而且天然自带噪声结构。GradNorm不依赖任何分布假设它完全从梯度的几何指标出发无论你是分类还是回归都能用。实验里我的感觉是分类任务组合上GradNorm的稳定性更好回归任务组合上两者都还可以但GradNorm不用调不确定性权重的初始参数相对省心。一个实际方案是把两者结合用GradNorm设置w的初始化再用不确定性加权在线微调这个思路在几个公开多任务项目里效果不错。5.3 我在项目中使用的几个实用变体GradNorm最消耗额外时间的地方是计算每个任务单梯度范数的过程每个任务都要单独backward。实际任务数量多的时候这个额外计算变成了不小的开销。我的做法是第一只在交替步中使用GradNorm比如每两步网络更新其中一步做GradNorm的权重刷新另一步纯网络更新效果损失不大第二用指数移动平均平滑当前r_i降低单步的随机噪声避免w在step到step之间跳变太剧烈。另一个变体是不在全网络层上均匀统计梯度范数而是只统计共享网络后端最后一两层即最接近任务分支的共享表示层。这样更贴合“每个任务抢共享表示”的直觉计算量也更小。如果你在某个任务上改过网络结构建议重新评估一下选用哪些层作为shared_params的观测点。还有一点值得注意GradNorm对梯度的计算依赖当前batch的loss波动如果batchsize很小单任务梯度范数噪声很大。工程上至少保证batchsize在64以上再直接使用否则建议配合EMA平滑一起用。这几个变体归纳起来就是GradNorm是主干逻辑平滑和稀疏更新是工程加固两件事不冲突。我自己跑完GradNorm之后最大的体会是多任务失衡本质上不是一个调参问题而是一个观测问题。你只要把“每个任务对共享层的梯度出力到底有多大”量化出来平衡策略就会自动浮出水面。与其继续在手工权重里挣扎不如先把GradNorm的梯度日志打出来看清你的多任务训练到底是谁在压制谁。