如果你训练过含多个任务的学习模型,大概率撞过这个怪现象:主任务精读涨得很漂亮,辅助任务却怎么拖都拖不动,或者反过来,辅助任务把主任务带偏了。我一开始的处理方式也很原始,手动调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], lr=0.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_graph=True) 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_graph=True) 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_graph=True) 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之后再进行clip,GradNorm的辅助更新和主网络更新顺序不能交叉。
第三个坑是跟学习率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的梯度日志打出来,看清你的多任务训练到底是谁在压制谁。