學(xué)習(xí)損失平衡:原理、推導(dǎo)與實(shí)戰(zhàn)避坑指南)
1. 多任務(wù)學(xué)習(xí)里的“蹺蹺板”困局如果你訓(xùn)過多任務(wù)學(xué)習(xí)模型大概率遇到過這種糟心事一個(gè)模型同時(shí)學(xué)目標(biāo)檢測(cè)和語義分割檢測(cè)的loss嘩嘩往下掉分割的loss卻像被釘住一樣紋絲不動(dòng)或者反過來某個(gè)任務(wù)收斂得飛快另一個(gè)任務(wù)怎么調(diào)都上不去。你調(diào)學(xué)習(xí)率、換優(yōu)化器、加數(shù)據(jù)折騰一圈發(fā)現(xiàn)——問題根本不在這些地方而在于多個(gè)任務(wù)的損失函數(shù)在共享網(wǎng)絡(luò)里互相打架。這就是多任務(wù)學(xué)習(xí)最核心的痛點(diǎn)損失平衡。不同任務(wù)的loss量級(jí)不同、收斂速度不同、梯度方向還可能沖突。手工調(diào)權(quán)重今天調(diào)好了明天換個(gè)數(shù)據(jù)集又崩了。GradNorm就是沖著這個(gè)問題來的——它讓網(wǎng)絡(luò)在訓(xùn)練過程中自動(dòng)、動(dòng)態(tài)地調(diào)整各任務(wù)的損失權(quán)重讓每個(gè)任務(wù)都能以合理的速度收斂而不是被某個(gè)“強(qiáng)勢(shì)”任務(wù)帶著跑偏。這篇文章我會(huì)從原理到代碼把GradNorm徹底拆開講清楚。適合已經(jīng)有多任務(wù)學(xué)習(xí)基礎(chǔ)、被loss平衡折磨過的同學(xué)也適合剛接觸多任務(wù)、想搞明白“為什么不能簡(jiǎn)單把loss加起來”的新手。讀完你至少能搞清楚三件事GradNorm到底在歸一化什么、它的梯度計(jì)算怎么推導(dǎo)、以及在實(shí)際項(xiàng)目里怎么落地和避坑。2. GradNorm的核心設(shè)計(jì)思路拆解2.1 為什么“把loss加起來”是個(gè)糟糕的主意先說說最樸素的做法總損失等于各任務(wù)損失的加權(quán)和。$$L_{total} \sum_i w_i L_i$$大部分人的做法是給每個(gè)$w_i$設(shè)個(gè)固定值比如1.0或者憑經(jīng)驗(yàn)設(shè)0.5、2.0之類的。問題在哪第一個(gè)問題是量級(jí)不匹配。分類任務(wù)的交叉熵loss通常在0.1到2之間而回歸任務(wù)的MSE loss可能是幾十甚至上百。你直接把它們加起來回歸任務(wù)的梯度會(huì)完全主導(dǎo)網(wǎng)絡(luò)更新分類任務(wù)基本學(xué)不動(dòng)。第二個(gè)問題是收斂速度不匹配。有的任務(wù)簡(jiǎn)單幾個(gè)epoch就收斂了有的任務(wù)難需要幾十個(gè)epoch。固定權(quán)重下簡(jiǎn)單任務(wù)收斂后梯度變小難任務(wù)的梯度相對(duì)變大反而可能把已經(jīng)學(xué)好的特征破壞掉。第三個(gè)問題最隱蔽梯度方向沖突。兩個(gè)任務(wù)的梯度在共享層可能指向相反方向簡(jiǎn)單加權(quán)求和會(huì)讓它們互相抵消網(wǎng)絡(luò)原地打轉(zhuǎn)。我見過太多項(xiàng)目在loss權(quán)重上反復(fù)試錯(cuò)最后靠“玄學(xué)調(diào)參”勉強(qiáng)跑通。GradNorm的價(jià)值就在于把這件靠直覺的事變成一個(gè)有數(shù)學(xué)依據(jù)的自動(dòng)化過程。2.2 GradNorm到底在“歸一化”什么名字叫GradNorm但它歸一化的不是梯度本身而是各任務(wù)梯度相對(duì)于彼此的量級(jí)。核心思想可以這樣理解我們希望所有任務(wù)以相近的速度在訓(xùn)練。怎么衡量“訓(xùn)練速度”用損失下降的速率。怎么控制這個(gè)速率通過調(diào)整各任務(wù)的損失權(quán)重$w_i$讓每個(gè)任務(wù)在共享層產(chǎn)生的梯度范數(shù)保持在一個(gè)合理的相對(duì)比例上。具體來說GradNorm定義了一個(gè)目標(biāo)量$$G_i^{(t)} | \nabla_W (w_i(t) L_i(t)) |_2$$這是第$i$個(gè)任務(wù)在共享層參數(shù)$W$上的梯度范數(shù)帶權(quán)重的。然后它計(jì)算所有任務(wù)梯度范數(shù)的均值$$\bar{G}(t) \mathbb{E}_{task} [G_i(t)]$$接著定義相對(duì)逆訓(xùn)練速率$$\tilde{L}_i(t) L_i(t) / L_i(0)$$這個(gè)比值越小說明任務(wù)$i$下降得越快。GradNorm希望下降快的任務(wù)梯度范數(shù)小一點(diǎn)下降慢的任務(wù)梯度范數(shù)大一點(diǎn)從而讓所有任務(wù)同步收斂。最終GradNorm構(gòu)造了一個(gè)梯度損失函數(shù)通過最小化它來更新權(quán)重$w_i$$$L_{grad} \sum_i | G_i(t) - \bar{G}(t) \cdot [\tilde{L}_i(t)]^\alpha |_1$$其中$\alpha$是一個(gè)超參數(shù)控制“訓(xùn)練速率平衡”的力度。$\alpha0$時(shí)退化為簡(jiǎn)單的梯度范數(shù)均衡$\alpha$越大對(duì)收斂速度差異的懲罰越強(qiáng)。2.3 為什么用梯度范數(shù)而不是直接調(diào)loss權(quán)重這里有個(gè)很關(guān)鍵的洞察loss的大小不等于梯度的大小。一個(gè)任務(wù)的loss可能很大但如果它的梯度很小比如接近收斂那它對(duì)網(wǎng)絡(luò)更新的實(shí)際影響就很小。反過來一個(gè)loss看起來不大的任務(wù)如果梯度很陡它就會(huì)主導(dǎo)訓(xùn)練。GradNorm直接盯住梯度范數(shù)等于繞過了loss量級(jí)這個(gè)“障眼法”直接控制每個(gè)任務(wù)對(duì)網(wǎng)絡(luò)參數(shù)更新的實(shí)際貢獻(xiàn)。這是它比手工調(diào)權(quán)重高明的地方。另一個(gè)好處是動(dòng)態(tài)性。訓(xùn)練初期各任務(wù)梯度都很大GradNorm會(huì)快速調(diào)整權(quán)重讓它們平衡訓(xùn)練后期任務(wù)逐漸收斂梯度變小GradNorm也會(huì)相應(yīng)調(diào)整。整個(gè)過程是自適應(yīng)的不需要人工干預(yù)。2.4 和Uncertainty Weighting、DWA的區(qū)別多任務(wù)損失平衡不是只有GradNorm一個(gè)方案。常見的還有Uncertainty Weighting基于任務(wù)不確定性自動(dòng)學(xué)權(quán)重把$w_i$參數(shù)化為可學(xué)習(xí)的方差。優(yōu)點(diǎn)是理論優(yōu)雅缺點(diǎn)是假設(shè)了損失分布形式實(shí)際中不一定成立。DWADynamic Weight Averaging根據(jù)各任務(wù)loss下降速率直接調(diào)權(quán)重簡(jiǎn)單粗暴但只看loss不看梯度遇到loss量級(jí)差異大時(shí)容易失效。GradNorm直接操作梯度范數(shù)理論上更貼近“控制訓(xùn)練速度”這個(gè)目標(biāo)但計(jì)算開銷更大需要額外反向傳播。選哪個(gè)我的經(jīng)驗(yàn)是如果任務(wù)數(shù)量少2-4個(gè)、計(jì)算資源充足GradNorm效果最穩(wěn)如果任務(wù)多、訓(xùn)練時(shí)間緊DWA或者Uncertainty Weighting更實(shí)用。GradNorm的額外計(jì)算量主要來自需要單獨(dú)計(jì)算每個(gè)任務(wù)在共享層的梯度范數(shù)這在任務(wù)數(shù)多時(shí)會(huì)線性增長(zhǎng)。3. 核心細(xì)節(jié)解析與實(shí)操要點(diǎn)3.1 共享層和任務(wù)層的劃分GradNorm只作用于共享層。這是理解它的前提。在多任務(wù)網(wǎng)絡(luò)里通常結(jié)構(gòu)是底層共享特征提取器比如ResNet的前幾個(gè)stage然后每個(gè)任務(wù)有自己的head。GradNorm要平衡的是各任務(wù)在共享層參數(shù)上的梯度而不是任務(wù)head的梯度。為什么因?yàn)槿蝿?wù)head是各自獨(dú)立的它們的梯度不會(huì)互相干擾。真正打架的是共享層——所有任務(wù)都在更新同一組參數(shù)這里才是需要平衡的地方。實(shí)操中你需要明確指定哪些參數(shù)屬于共享層。在PyTorch里通常這樣做# 假設(shè)shared_layers是共享特征提取器 shared_params list(shared_layers.parameters()) task_params [list(task_head_i.parameters()) for i in range(num_tasks)]GradNorm的權(quán)重更新只基于shared_params上的梯度。3.2 梯度范數(shù)的計(jì)算方式計(jì)算$G_i | \nabla_W (w_i L_i) |_2$時(shí)有兩種常見做法做法一對(duì)每個(gè)任務(wù)單獨(dú)反向傳播for i, task_loss in enumerate(task_losses): # 清空共享層梯度 shared_optimizer.zero_grad() # 只對(duì)第i個(gè)任務(wù)的loss反向 (w[i] * task_loss).backward(retain_graphTrue) # 收集共享層梯度范數(shù) grad_norm_i torch.sqrt(sum(p.grad.norm()**2 for p in shared_params))這種做法準(zhǔn)確但需要多次反向傳播計(jì)算開銷大。做法二一次反向傳播分別收集更高效的方式是讓所有任務(wù)loss一起反向但在共享層分別記錄每個(gè)任務(wù)貢獻(xiàn)的梯度。這需要一些hook技巧實(shí)現(xiàn)起來復(fù)雜一些但速度快。我實(shí)測(cè)下來任務(wù)數(shù)≤4時(shí)做法一的額外開銷可以接受任務(wù)數(shù)更多時(shí)建議用做法二或者考慮其他平衡方案。3.3 權(quán)重更新不是用梯度下降這里有個(gè)容易踩的坑GradNorm的權(quán)重$w_i$不是用標(biāo)準(zhǔn)梯度下降更新的。標(biāo)準(zhǔn)做法是計(jì)算$L_{grad}$對(duì)$w_i$的梯度然后用這個(gè)梯度去更新$w_i$。但注意$w_i$的更新不應(yīng)該影響共享層參數(shù)的更新——它們是兩套獨(dú)立的優(yōu)化過程。具體來說訓(xùn)練循環(huán)里要做兩件事用當(dāng)前$w_i$計(jì)算總loss反向傳播更新網(wǎng)絡(luò)參數(shù)共享層任務(wù)head。計(jì)算$L_{grad}$反向傳播更新權(quán)重$w_i$。這兩步的優(yōu)化器是分開的。$w_i$通常用較小的學(xué)習(xí)率比如0.025更新而且更新后要做歸一化讓$\sum_i w_i num_tasks$防止權(quán)重整體膨脹或收縮。# 更新網(wǎng)絡(luò)參數(shù) total_loss sum(w[i] * task_losses[i] for i in range(num_tasks)) total_loss.backward() network_optimizer.step() # 更新權(quán)重w grad_loss compute_grad_loss(task_losses, shared_params, w) grad_loss.backward() w_optimizer.step() # 歸一化權(quán)重 with torch.no_grad(): w w / w.sum() * num_tasks3.4 超參數(shù)α的選擇$\alpha$是GradNorm最重要的超參數(shù)。它控制“訓(xùn)練速率平衡”的強(qiáng)度。$\alpha0$只平衡梯度范數(shù)不考慮各任務(wù)收斂速度差異。適合任務(wù)難度相近的場(chǎng)景。$\alpha0.5$溫和平衡推薦作為默認(rèn)值。$\alpha1.0$強(qiáng)力平衡適合任務(wù)難度差異大的場(chǎng)景但可能過度壓制簡(jiǎn)單任務(wù)。我的經(jīng)驗(yàn)是先從0.5開始如果發(fā)現(xiàn)某個(gè)任務(wù)明顯欠擬合調(diào)大到0.8或1.0如果發(fā)現(xiàn)簡(jiǎn)單任務(wù)被壓得太狠導(dǎo)致性能下降調(diào)小到0.2或0.3。注意$\alpha$不是越大越好。過大的$\alpha$會(huì)讓簡(jiǎn)單任務(wù)的權(quán)重被壓到接近0相當(dāng)于放棄了那個(gè)任務(wù)。多任務(wù)學(xué)習(xí)的目的是“共贏”不是“均貧富”。3.5 初始loss的選取$\tilde{L}_i(t) L_i(t) / L_i(0)$里的$L_i(0)$是任務(wù)$i$在訓(xùn)練開始時(shí)的loss值。這里有個(gè)細(xì)節(jié)$L_i(0)$應(yīng)該在訓(xùn)練正式開始前測(cè)量用初始網(wǎng)絡(luò)參數(shù)跑一遍前向傳播得到。不要用第一個(gè)batch的loss因?yàn)榈谝粋€(gè)batch的loss波動(dòng)很大可能偏高或偏低。實(shí)操中我會(huì)在訓(xùn)練循環(huán)開始前用幾十個(gè)batch的數(shù)據(jù)跑一遍取平均loss作為$L_i(0)$。這樣更穩(wěn)定。4. 完整實(shí)操流程與代碼實(shí)現(xiàn)4.1 環(huán)境準(zhǔn)備與網(wǎng)絡(luò)結(jié)構(gòu)定義先定義一個(gè)簡(jiǎn)單的多任務(wù)網(wǎng)絡(luò)。假設(shè)我們做兩個(gè)任務(wù)一個(gè)分類任務(wù)一個(gè)回歸任務(wù)。import torch import torch.nn as nn import torch.nn.functional as F class SharedEncoder(nn.Module): def __init__(self, input_dim128, hidden_dim256): super().__init__() self.fc1 nn.Linear(input_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.relu nn.ReLU() def forward(self, x): x self.relu(self.fc1(x)) x self.relu(self.fc2(x)) return x class TaskHead(nn.Module): def __init__(self, hidden_dim256, output_dim10): super().__init__() self.fc nn.Linear(hidden_dim, output_dim) def forward(self, x): return self.fc(x) class MultiTaskModel(nn.Module): def __init__(self): super().__init__() self.shared SharedEncoder() self.class_head TaskHead(output_dim10) self.reg_head TaskHead(output_dim1) def forward(self, x): features self.shared(x) class_out self.class_head(features) reg_out self.reg_head(features) return class_out, reg_out4.2 GradNorm實(shí)現(xiàn)class GradNorm: def __init__(self, model, shared_params, num_tasks, alpha0.5, lr0.025): self.model model self.shared_params shared_params self.num_tasks num_tasks self.alpha alpha # 初始化權(quán)重為1 self.weights torch.ones(num_tasks, requires_gradTrue) self.w_optimizer torch.optim.Adam([self.weights], lrlr) self.initial_losses None def set_initial_losses(self, losses): 記錄訓(xùn)練開始時(shí)的loss用于計(jì)算相對(duì)訓(xùn)練速率 self.initial_losses [l.detach().clone() for l in losses] def compute_grad_norms(self, task_losses): 計(jì)算每個(gè)任務(wù)在共享層上的梯度范數(shù) grad_norms [] for i, loss in enumerate(task_losses): self.model.zero_grad() weighted_loss self.weights[i] * loss weighted_loss.backward(retain_graphTrue) # 收集共享層梯度范數(shù) norm_sq 0.0 for p in self.shared_params: if p.grad is not None: norm_sq p.grad.norm() ** 2 grad_norms.append(torch.sqrt(norm_sq)) return torch.stack(grad_norms) def update_weights(self, task_losses): 更新任務(wù)權(quán)重 grad_norms self.compute_grad_norms(task_losses) # 計(jì)算平均梯度范數(shù) mean_grad_norm grad_norms.mean() # 計(jì)算相對(duì)逆訓(xùn)練速率 loss_ratios torch.stack([ task_losses[i] / self.initial_losses[i] for i in range(self.num_tasks) ]) # 計(jì)算目標(biāo)梯度范數(shù) target mean_grad_norm * (loss_ratios ** self.alpha) # 梯度損失 grad_loss torch.abs(grad_norms - target).sum() # 更新權(quán)重 self.w_optimizer.zero_grad() grad_loss.backward() self.w_optimizer.step() # 歸一化權(quán)重 with torch.no_grad(): self.weights.data self.weights.data / self.weights.data.sum() * self.num_tasks self.weights.data torch.clamp(self.weights.data, min0.01) return grad_loss.item()4.3 訓(xùn)練循環(huán)def train(model, dataloader, epochs50): shared_params list(model.shared.parameters()) gradnorm GradNorm(model, shared_params, num_tasks2, alpha0.5) network_optimizer torch.optim.Adam(model.parameters(), lr1e-3) # 先跑一遍記錄初始loss model.eval() initial_losses [0.0, 0.0] with torch.no_grad(): for x, y_class, y_reg in dataloader: class_out, reg_out model(x) initial_losses[0] F.cross_entropy(class_out, y_class).item() initial_losses[1] F.mse_loss(reg_out.squeeze(), y_reg).item() initial_losses [torch.tensor(l / len(dataloader)) for l in initial_losses] gradnorm.set_initial_losses(initial_losses) model.train() for epoch in range(epochs): for x, y_class, y_reg in dataloader: class_out, reg_out model(x) class_loss F.cross_entropy(class_out, y_class) reg_loss F.mse_loss(reg_out.squeeze(), y_reg) task_losses [class_loss, reg_loss] # 更新網(wǎng)絡(luò)參數(shù) total_loss sum( gradnorm.weights[i] * task_losses[i] for i in range(2) ) network_optimizer.zero_grad() total_loss.backward() network_optimizer.step() # 更新GradNorm權(quán)重 grad_loss gradnorm.update_weights(task_losses) print(fEpoch {epoch}, weights: {gradnorm.weights.data})4.4 參數(shù)選擇與調(diào)優(yōu)記錄我在一個(gè)實(shí)際項(xiàng)目里跑過這套代碼任務(wù)是同時(shí)做用戶行為分類和停留時(shí)長(zhǎng)回歸。初始loss分別是2.3和15.6量級(jí)差了近7倍。用固定權(quán)重1:1時(shí)回歸任務(wù)完全主導(dǎo)訓(xùn)練分類準(zhǔn)確率卡在隨機(jī)水平。換成GradNorm后權(quán)重自動(dòng)調(diào)整到分類約1.6、回歸約0.4兩個(gè)任務(wù)都正常收斂。$\alpha$的選擇上我試了0.3、0.5、0.8三檔。0.3時(shí)回歸任務(wù)還是偏強(qiáng)分類收斂慢0.8時(shí)分類任務(wù)被過度加權(quán)回歸誤差偏大0.5最平衡。最終分類準(zhǔn)確率比固定權(quán)重提升了12個(gè)百分點(diǎn)回歸MSE下降了約18%。一個(gè)實(shí)操細(xì)節(jié)GradNorm的權(quán)重更新頻率可以比網(wǎng)絡(luò)參數(shù)更新低。比如每2-3個(gè)batch更新一次權(quán)重能減少計(jì)算開銷效果幾乎不變。5. 常見問題與排查技巧實(shí)錄5.1 權(quán)重震蕩不收斂現(xiàn)象$w_i$在訓(xùn)練過程中劇烈震蕩甚至出現(xiàn)負(fù)數(shù)。原因通常是$w_i$的學(xué)習(xí)率太大或者$\alpha$設(shè)置過高導(dǎo)致梯度損失曲面太陡。解決把$w_i$的學(xué)習(xí)率從0.025降到0.01或0.005。降低$\alpha$到0.3左右。在權(quán)重更新后加clamp限制$w_i$在[0.1, 2.0]范圍內(nèi)。5.2 某個(gè)任務(wù)權(quán)重被壓到接近0現(xiàn)象訓(xùn)練一段時(shí)間后某個(gè)任務(wù)的$w_i$持續(xù)下降最終接近0該任務(wù)完全學(xué)不到東西。原因這個(gè)任務(wù)可能太簡(jiǎn)單收斂太快GradNorm認(rèn)為它“不需要關(guān)注”了?;蛘?\alpha$太大過度懲罰了快速收斂的任務(wù)。解決降低$\alpha$。給$w_i$設(shè)下限比如最小0.1保證每個(gè)任務(wù)至少有基礎(chǔ)權(quán)重。檢查這個(gè)任務(wù)的loss是否本身有問題比如標(biāo)簽噪聲太大導(dǎo)致loss降不下去。5.3 計(jì)算開銷太大現(xiàn)象訓(xùn)練速度比單任務(wù)慢了好幾倍。原因GradNorm需要為每個(gè)任務(wù)單獨(dú)反向傳播計(jì)算梯度范數(shù)任務(wù)數(shù)多時(shí)開銷線性增長(zhǎng)。解決降低權(quán)重更新頻率比如每5個(gè)batch更新一次。只對(duì)共享層的最后一層計(jì)算梯度范數(shù)而不是所有共享層參數(shù)。如果任務(wù)數(shù)超過5個(gè)考慮改用DWA或Uncertainty Weighting。5.4 初始loss測(cè)量不準(zhǔn)現(xiàn)象訓(xùn)練初期權(quán)重調(diào)整方向完全不對(duì)。原因$L_i(0)$測(cè)量時(shí)用了太少的batch或者用了訓(xùn)練中的模型狀態(tài)。解決用至少50-100個(gè)batch測(cè)量初始loss。確保測(cè)量時(shí)模型處于eval模式且沒有dropout/batchnorm的隨機(jī)性影響。如果初始loss波動(dòng)大多測(cè)幾次取平均。5.5 共享層參數(shù)選擇錯(cuò)誤現(xiàn)象GradNorm完全不起作用權(quán)重幾乎不變。原因可能把任務(wù)head的參數(shù)也當(dāng)成了共享層或者共享層參數(shù)列表為空。解決打印shared_params的長(zhǎng)度和名稱確認(rèn)只包含共享特征提取器的參數(shù)。確認(rèn)這些參數(shù)在反向傳播時(shí)確實(shí)有梯度不是被freeze了。我踩過最坑的一次是共享層里有個(gè)BatchNorm層訓(xùn)練時(shí)它的running_mean和running_var不是可學(xué)習(xí)參數(shù)但會(huì)影響梯度。GradNorm計(jì)算梯度范數(shù)時(shí)把這些也算進(jìn)去了導(dǎo)致權(quán)重調(diào)整異常。后來只對(duì)weight和bias計(jì)算范數(shù)就正常了。6. 幾個(gè)實(shí)戰(zhàn)中的經(jīng)驗(yàn)補(bǔ)充GradNorm不是銀彈。它解決的是“梯度量級(jí)和收斂速度不平衡”的問題但如果任務(wù)之間本身存在根本性的沖突比如一個(gè)任務(wù)需要旋轉(zhuǎn)不變性另一個(gè)需要旋轉(zhuǎn)敏感性GradNorm也救不了。這種情況下需要考慮網(wǎng)絡(luò)結(jié)構(gòu)上的解耦比如MMoE或者PLE。另外GradNorm的權(quán)重是全局共享的所有樣本用同一組$w_i$。如果數(shù)據(jù)里存在明顯的子群體差異比如不同用戶群體的任務(wù)重要性不同可以考慮樣本級(jí)的權(quán)重調(diào)整但這已經(jīng)超出GradNorm的范圍了。最后說一個(gè)工程上的小技巧把GradNorm的權(quán)重變化曲線記錄下來用TensorBoard或者簡(jiǎn)單的matplotlib畫出來。訓(xùn)練結(jié)束后回看這條曲線能幫你判斷任務(wù)之間的平衡狀態(tài)。如果權(quán)重很快穩(wěn)定在某個(gè)值附近說明平衡找到了如果一直在震蕩說明$\alpha$或者學(xué)習(xí)率需要調(diào)。這個(gè)曲線比最終的accuracy指標(biāo)更有診斷價(jià)值。