網(wǎng)絡(luò)參數(shù)初始化全解析:從梯度傳播原理到PyTorch實(shí)戰(zhàn))
訓(xùn)練神經(jīng)網(wǎng)絡(luò)這幾年我有一多半的“模型不收斂”最終都指向同一個(gè)元兇——不是網(wǎng)絡(luò)搭錯(cuò)了不是學(xué)習(xí)率沒調(diào)好也不是數(shù)據(jù)喂得不對而是參數(shù)初始化沒做好。很多人把PyTorch當(dāng)黑盒模型構(gòu)建完直接傳數(shù)據(jù)、算loss、backward等loss變成了水平線才想起來檢查權(quán)重。其實(shí)在一輪訓(xùn)練開始之前每一層權(quán)重分布就已經(jīng)悄悄決定了后面這條路是康莊大道還是萬丈深淵。這篇想把這個(gè)主題徹底聊透為什么參數(shù)初始化能在訓(xùn)練的第一步就決定成敗Xavier和Kaiming這些主流方案背后的數(shù)學(xué)動(dòng)機(jī)是什么怎樣在PyTorch框架下把初始化環(huán)節(jié)完全握在自己手里以及我在實(shí)際項(xiàng)目中踩過的初始化相關(guān)的坑。適合剛?cè)腴TPyTorch和神經(jīng)網(wǎng)絡(luò)基礎(chǔ)的朋友也適合那些訓(xùn)練過不少模型、卻從來沒主動(dòng)干預(yù)過初始化的同學(xué)。相信我看完你會(huì)忍不住去檢查自己那幾層網(wǎng)絡(luò)到底是怎么“起跑”的。1. 初始化是第一道關(guān)卡從梯度傳播看它為啥這么重要很多教程在講神經(jīng)網(wǎng)絡(luò)時(shí)把初始化當(dāng)成一個(gè)“選個(gè)隨機(jī)數(shù)就行”的步驟一筆帶過。但實(shí)際上初始化的質(zhì)量決定了網(wǎng)絡(luò)在一開始處于損失曲面上的哪個(gè)位置也決定了反向傳播的梯度信號能不能完整地傳回淺層。1.1 一個(gè)讓模型“學(xué)不動(dòng)”的真實(shí)場景先說一個(gè)我自己的案例。早年在本地跑一個(gè)淺層全連接網(wǎng)絡(luò)做回歸預(yù)測結(jié)構(gòu)很簡單三個(gè)隱藏層每層64個(gè)神經(jīng)元激活函數(shù)用ReLU輸出層不加激活直接回歸。數(shù)據(jù)歸一化做得干干凈凈學(xué)習(xí)率從1e-2一路試著降到1e-5loss就是卡在某個(gè)值附近一動(dòng)不動(dòng)連下降的苗頭都沒有。排查了很久最后把權(quán)重打印出來看分布發(fā)現(xiàn)是我在用PyTorch搭網(wǎng)絡(luò)時(shí)手動(dòng)把權(quán)重全部初始化成了均值0、方差特別大的正態(tài)分布隨機(jī)數(shù)。前向傳播時(shí)每一層的輸出都在指數(shù)級放大到最后一層已經(jīng)全是幾百幾千的量級loss直接爆炸式增長梯度也出現(xiàn)了大量NaN。把權(quán)重的初始方差降回合理范圍之后同一個(gè)模型、同一個(gè)學(xué)習(xí)率幾十個(gè)epoch就正常收斂了。這個(gè)案例給我留下一個(gè)很深的印象初始化問題往往不會(huì)在代碼層面爆出紅色報(bào)錯(cuò)而是以一種“l(fā)oss死活不降”的慢性病形式出現(xiàn)折騰你幾天幾夜。想通這一點(diǎn)就會(huì)明白我們?yōu)槭裁葱枰J(rèn)真對待每一層參數(shù)的初始值。1.2 梯度傳播中的連乘效應(yīng)信號如何消失或爆炸要理解初始化先要看一次完整的前向和反向傳播中發(fā)生了什么。以一個(gè)不帶偏置的線性層為例輸出滿足y Wx那一層的輸入 x 有 n 個(gè)維度輸出 y 有 m 個(gè)維度。如果 x 的每個(gè)分量方差是 Var(x)權(quán)重 W 中每個(gè)元素的方差是 Var(W)那么 y 中某個(gè)分量的方差假設(shè)各分量獨(dú)立大約是Var(y) ≈ n × Var(W) × Var(x)也就是說信號經(jīng)過一層之后方差被放大了大約 n×Var(W) 倍。為了讓信息在多層網(wǎng)絡(luò)中傳遞時(shí)既不衰減到零、也不膨脹到爆一個(gè)自然的目標(biāo)是讓 Var(y) ≈ Var(x)。于是就有了第一個(gè)直覺結(jié)論Var(W) ≈ 1 / nn 就是這一層的輸入維度在初始化理論里叫 fan_in扇入。反向傳播是同樣的邏輯只是信號變成了梯度。設(shè) loss 對 y 的梯度是 dy那么對 x 的梯度是dx W? dy此時(shí)經(jīng)過這一層反向傳播時(shí)梯度的方差由輸出維度 mfan_out扇出決定。要讓梯度反向傳播時(shí)保持穩(wěn)定需要Var(W) ≈ 1 / m前向希望按 fan_in 來定方差反向希望按 fan_out 來定方差兩個(gè)需求不一致怎么辦這就是不同初始化方法分道揚(yáng)鑣的地方。比如Xavier取兩者的調(diào)和折中Kaiming則根據(jù)激活函數(shù)特性調(diào)整系數(shù)。但不管哪種方法核心都是在控制連乘效應(yīng)的放大倍數(shù)讓它穩(wěn)定在1附近。1.3 對稱性陷阱所有神經(jīng)元變成同一個(gè)人除了梯度消失和爆炸初始化還藏著一個(gè)更隱蔽的陷阱——對稱性。如果同一層的所有權(quán)重初始化為相同的常數(shù)比如全零、全0.1那么這層所有神經(jīng)元的輸入分布完全相同。反向傳播計(jì)算出的梯度對每個(gè)神經(jīng)元也完全相同。于是無論怎么更新這些神經(jīng)元永遠(yuǎn)走一樣的路整個(gè)隱藏層實(shí)際上退化成了一個(gè)神經(jīng)元。這就是為什么“全零初始化”在理論上被明確否定它會(huì)讓多層網(wǎng)絡(luò)變成一層單神經(jīng)元的表達(dá)能力。實(shí)際工程中沒人蠢到全零但不少人會(huì)把權(quán)重設(shè)成全零而忘了偏置或者把線性層和卷積層的偏置習(xí)慣性設(shè)成全零——這本身沒問題只要權(quán)重本身是“各不相同”的隨機(jī)數(shù)對稱性就被打破了。偏置的初始化和權(quán)重邏輯不一樣。權(quán)重一旦全相同會(huì)造成對稱退化偏置全零卻無傷大雅因?yàn)槠玫妮斎雭碜郧耙粚拥姆菍ΨQ激活值。大多數(shù)框架里的默認(rèn)選項(xiàng)也是“偏置為0”或“很小的隨機(jī)數(shù)”這是合理的。提示隨機(jī)數(shù)只是表象核心是“破壞對稱”和“控制方差”這兩件事。任何初始化方案本質(zhì)上都是在回答這兩個(gè)問題每個(gè)參數(shù)的方差應(yīng)該多大在什么范圍內(nèi)取隨機(jī)數(shù)2. 主流初始化方法解析Xavier、Kaiming和其他選手的來龍去脈深度學(xué)習(xí)研究這么多年初始化方法也就那幾款主流選手在打天下。Mike的入門路線圖大概是先認(rèn)識(shí)Glorot即Xavier初始化再掌握專為ReLU而生的Kaiming初始化最后了解RNN場景下的正交初始化以及各類偏置、正則化層的處理。2.1 Xavier/Glorot初始化為對稱激活函數(shù)而生的理論基準(zhǔn)2010年Glorot和Bengio在《Understanding the difficulty of training deep feedforward neural networks》里提出了一個(gè)著名的方法PyTorch里叫xavier_uniform_和xavier_normal_。它針對的是tanh這類關(guān)于原點(diǎn)對稱、激活值在0附近的函數(shù)。前面說過方差的理想取值要兼顧前向的fan_in和反向的fan_out。Xavier的折中方案是Var(W) 2 / (fan_in fan_out)如果是均勻分布 U(-a, a)均勻分布的方差是 a2/3那么a2/3 2 / (fan_in fan_out)解得a √(6 / (fan_in fan_out))這樣xavier_uniform_的標(biāo)準(zhǔn)邊界就是W ~ U(-√(6/(fan_infan_out)), √(6/(fan_infan_out)))如果是正態(tài)分布直接取均值為0、方差為2/(fan_infan_out)。這套方案有個(gè)隱含的假設(shè)激活函數(shù)在零點(diǎn)附近的斜率接近1比如tanh在0點(diǎn)的導(dǎo)數(shù)是1sigmoid在0點(diǎn)的導(dǎo)數(shù)是0.25。所以你會(huì)發(fā)現(xiàn)PyTorch的xavier_系列文檔里明確寫著“推薦用于tanh和sigmoid類型激活”。如果把它硬套在ReLU上會(huì)因?yàn)镽eLU直接把一半信號砍成0而導(dǎo)致方差不匹配。2.2 Kaiming/He初始化專門給ReLU家族打的補(bǔ)丁2015年何愷明團(tuán)隊(duì)在《Delving Deep into Rectifiers: Surpassing Human-Level Performance on ImageNet Classification》中推出了針對ReLU的初始化方案也就是PyTorch里的kaiming_uniform_和kaiming_normal_。ReLU有個(gè)特性輸入負(fù)數(shù)時(shí)輸出恒為0。假設(shè)輸入x分布關(guān)于0對稱經(jīng)過ReLU后約有一半的信號變成了0剩余一半保持原樣整體方差大約只有輸入的一半。這相當(dāng)于信號每經(jīng)過一次ReLU就減半。為了補(bǔ)償這個(gè)衰減Kaiming初始化把權(quán)重方差調(diào)整為前向模式下Var(W) 2 / fan_in訓(xùn)練過程中前向傳播主要使用fan_in模式所以PyTorch的kaiming_normal_默認(rèn)modefan_in。如果設(shè)置modefan_out則方差為2/fan_out適合在需要保持反向梯度方差的場景下使用。對于均勻分布版本邊界相應(yīng)為W ~ U(-√(6/fan_in), √(6/fan_in))注意這里的分母里沒有fan_out。我之前見過有人把這個(gè)邊界和Xavier搞混結(jié)果模型淺層梯度衰減嚴(yán)重訓(xùn)練半天loss紋絲不動(dòng)。如果你用ReLU系激活請認(rèn)準(zhǔn)Kaiming。Leaky ReLU這類帶泄漏斜率的激活也適用Kaiming只是參數(shù)a要對應(yīng)設(shè)置。PyTorch的kaiming_uniform_支持通過a參數(shù)傳入負(fù)斜率。a√5時(shí)前面的計(jì)算已經(jīng)幫偏置提供了一個(gè)合理的默認(rèn)邊界這個(gè)在官方源碼的Linear初始化里也在用。2.3 偏置、BatchNorm和正交初始化容易被忽略的細(xì)節(jié)偏置初始化經(jīng)常被忽略但它也有自己的規(guī)范。全連接層和卷積層的偏置一般初始化為0即可因?yàn)闄?quán)重已經(jīng)打破了對稱性。PyTorch對Linear的默認(rèn)偏置初始化并非完全為0而是根據(jù)權(quán)重的fan_in算出一個(gè)邊界bound 1 / √(fan_in)然后偏置在U(-bound, bound)內(nèi)均勻采樣。這個(gè)設(shè)計(jì)其實(shí)很講究它讓偏置的初始量級和權(quán)重的輸出量級匹配避免網(wǎng)絡(luò)初期輸出分布發(fā)生大偏移。BatchNorm層的初始化就更特殊了權(quán)重縮放系數(shù)γ初始化為1偏置平移系數(shù)β初始化為0。這意味著BatchNorm在訓(xùn)練初期保持輸入分布的歸一化狀態(tài)先讓網(wǎng)絡(luò)“見見原貌”再逐步學(xué)習(xí)縮放和平移的必要性。如果你在加載別人代碼時(shí)看到把BatchNorm的γ亂設(shè)成大數(shù)模型通常很難收斂。RNN和LSTM這類循環(huán)結(jié)構(gòu)還經(jīng)常使用正交初始化。正交矩陣的列向量彼此垂直有一個(gè)很好的性質(zhì)矩陣乘法不會(huì)擴(kuò)大或壓縮向量長度。這能在長期依賴傳播中緩解梯度消失讓信息沿著時(shí)間步傳得更遠(yuǎn)。PyTorch的nn.init.orthogonal_就干這個(gè)活兒。不是說你非用不可但在訓(xùn)練較長序列時(shí)讓遺忘門偏置初始化為較大正值、讓輸入門偏置初始化為較小值再配合正交權(quán)重實(shí)際效果會(huì)明顯更穩(wěn)。提示第2節(jié)最后給個(gè)選擇邏輯——激活函數(shù)是tanh/sigmoid時(shí)優(yōu)先Xavier是ReLU/LeakyReLU時(shí)選Kaiming是RNN/LSTM時(shí)可在權(quán)重上加正交初始化。這三條基本覆蓋了絕大多數(shù)網(wǎng)絡(luò)。3. PyTorch初始化實(shí)操從默認(rèn)規(guī)則到完全掌控理論講再多最終要落到代碼。PyTorch給了我們幾層控制權(quán)默認(rèn)初始化、nn.init工具函數(shù)、apply批量處理、以及重寫reset_parameters。我從易到難逐個(gè)說。3.1 先搞清楚框架默認(rèn)做了什么很多人不知道PyTorch在創(chuàng)建模型時(shí)已經(jīng)做了初始化。比如nn.Linear和nn.Conv2d默認(rèn)權(quán)重使用kaiming_uniform_bias默認(rèn)使用U(-1/√fan_in, 1/√fan_in)。這對大多數(shù)ReLU網(wǎng)絡(luò)其實(shí)是夠用的。想確認(rèn)自己模型每一層的默認(rèn)初始化長什么樣可以打印weight的均值和標(biāo)準(zhǔn)差或者用m.weight.data.histc()看看分布。這里有個(gè)實(shí)用小技巧在第一次前向傳播前把模型各層參數(shù)的均值、標(biāo)準(zhǔn)差逐層打印出來一眼就能判斷有沒有層被初始化成了明顯不合理的量級。3.2 nn.init的API清單動(dòng)手改初始化PyTorch的torch.nn.init模塊提供了一套非常完整的函數(shù)。下面是常用的幾個(gè)我把它們的核心用途列出來xavier_uniform_(tensor, gain1)適合tanh/sigmoid均勻分布版本xavier_normal_(tensor, gain1)適合tanh/sigmoid正態(tài)分布版本kaiming_uniform_(tensor, a0, modefan_in, nonlinearityleaky_relu)適合ReLU/LeakyReLUkaiming_normal_(tensor, a0, modefan_in, nonlinearityleaky_relu)同上正態(tài)分布版本orthogonal_(tensor, gain1)適合RNN/LSTMones_(tensor)、zeros_(tensor)給偏置或BatchNorm的γ賦值constant_(tensor, val)固定常數(shù)eye_(tensor)單位矩陣初始化偶爾用于特定的注意力層使用方式也很直接拿到一個(gè)weight之后重新賦值即可import torch import torch.nn as nn linear nn.Linear(128, 64) nn.init.kaiming_normal_(linear.weight, modefan_in, nonlinearityrelu) nn.init.zeros_(linear.bias)需要注意kaiming系列有nonlinearity參數(shù)如果你漏傳了默認(rèn)是leaky_relua默認(rèn)0。如果實(shí)際激活是普通ReLU建議顯式寫nonlinearityrelu避免把自己繞暈。實(shí)際上對于a0leaky_relu和relu的數(shù)學(xué)計(jì)算完全等價(jià)但因?yàn)樵创a走的分支略有不同顯式寫relu更語義化、更安全。3.3 用model.apply批量接管整個(gè)模型的初始化一個(gè)模型幾十上百層不可能每層手動(dòng)去改。PyTorch提供了module.apply方法它會(huì)遞歸地把傳入的匿名函數(shù)作用在每個(gè)子模塊上。這是實(shí)戰(zhàn)中最常用的一招import torch.nn as nn def init_weights(module): acts (nn.Linear, nn.Conv1d, nn.Conv2d, nn.Conv3d) if isinstance(module, acts): nn.init.kaiming_normal_(module.weight, modefan_in, nonlinearityrelu) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.BatchNorm2d): nn.init.ones_(module.weight) nn.init.zeros_(module.bias) model nn.Sequential( nn.Conv2d(3, 16, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(16 * 14 * 14, 10) ) model.apply(init_weights)init_weights會(huì)被每個(gè)子模塊調(diào)用一次里面的isinstance判斷決定當(dāng)前子模塊走哪條初始化規(guī)則。BatchNorm2d的γ設(shè)為1、β設(shè)為0卷積權(quán)重用kaiming_normal_偏置置零。這個(gè)做法的好處是集中管理。改初始化策略時(shí)只需要改init_weights一個(gè)函數(shù)而不是去動(dòng)每個(gè)層定義。我習(xí)慣把init_weights放在模型文件最上方作為一個(gè)獨(dú)立工具函數(shù)修改時(shí)一目了然。3.4 覆蓋reset_parameters進(jìn)階玩法還有一種更“面向?qū)ο蟆钡姆绞骄褪侵貙懩愕淖远x模塊里的reset_parameters方法。PyTorch在每次創(chuàng)建模塊時(shí)都會(huì)調(diào)用這個(gè)方法。以自定義一個(gè)MLP塊為例import math import torch import torch.nn as nn class MyMLPBlock(nn.Module): def __init__(self, in_dim, out_dim): super().__init__() self.fc1 nn.Linear(in_dim, out_dim) self.relu nn.ReLU() def reset_parameters(self): nn.init.kaiming_uniform_(self.fc1.weight, amath.sqrt(5)) nn.init.zeros_(self.fc1.bias) def forward(self, x): return self.relu(self.fc1(x))執(zhí)行這個(gè)模塊的實(shí)例化時(shí)PyTorch會(huì)自動(dòng)調(diào)用reset_parameters。如果你什么都不寫默認(rèn)會(huì)調(diào)用父類nn.Module的reset_parameters對新創(chuàng)建的層執(zhí)行框架默認(rèn)初始化。重寫之后就完全由你說了算。我見過一些開源代碼直接在__init__里初始化權(quán)重不寫reset_parameters。這有個(gè)小隱患如果你后續(xù)想用model.apply重設(shè)整個(gè)模型的權(quán)重而那些層又沒有暴露apply能識(shí)別的結(jié)構(gòu)你的自定義層可能漏掉初始化。覆蓋reset_parameters配合apply是更規(guī)范的組合。4. 初始化不當(dāng)?shù)牡湫桶Y狀與排查路徑前面講了這么多原理下面說說實(shí)際運(yùn)行中最常遇到的問題。初始化的問題不會(huì)像語法錯(cuò)誤那樣直接拋異常它往往通過loss曲線和梯度統(tǒng)計(jì)來表達(dá)不滿。掌握這套“病理學(xué)”非常值錢。4.1 從loss曲線形態(tài)判斷初始化病情我見過最常見的情況有三種癥狀和“病因”對比如下癥狀可能病因解決方向loss初始值就是天文數(shù)字比如MSE回歸初始loss上萬輸出層或隱藏層權(quán)重方差過大前向輸出爆炸縮小初始標(biāo)準(zhǔn)差檢查是否有未初始化的自定義層loss從訓(xùn)練開始就完全不下降平穩(wěn)如直線梯度消失或權(quán)重對稱信號傳不到淺層切換為與激活函數(shù)匹配的初始化方法loss初期劇烈震蕩像心電圖一樣極端波動(dòng)初始方差偏大接近“混沌”狀態(tài)用Xavier/Kaiming的std或bound再減半訓(xùn)練跑了十幾個(gè)epoch后突然出現(xiàn)NaN初始化偏大疊加學(xué)習(xí)率偏高訓(xùn)練中期梯度爆炸降低學(xué)習(xí)率暫時(shí)縮小權(quán)重初始化方差其中“l(fā)oss初始值就是天文數(shù)字”這條特別容易騙到新手。有人看到初始loss幾萬第一反應(yīng)是改學(xué)習(xí)率或者改模型結(jié)構(gòu)其實(shí)只要把最后一層權(quán)重初始化方差調(diào)小一點(diǎn)loss立刻正常。4.2 逐層梯度檢查快速定位是哪層出了問題如果你懷疑初始化有問題最直接的辦法是打印每一層的梯度范數(shù)。PyTorch里可以用反向傳播后的grad屬性查看for name, param in model.named_parameters(): if param.grad is not None: grad_norm param.grad.norm().item() if grad_norm 1e-6: print(f{name}: grad norm too small, {grad_norm:.2e}) elif grad_norm 1e2: print(f{name}: grad norm too large, {grad_norm:.2e})把這段代碼塞進(jìn)訓(xùn)練循環(huán)里跑一個(gè)step之后看結(jié)果。如果靠近輸出層的參數(shù)梯度正常、靠近輸入層的梯度極小說明信號在反向傳播途中“熄滅”了典型的梯度消失大概率是激活函數(shù)和初始化不匹配。如果靠近輸入層的梯度極大、靠近輸出層的梯度正常說明梯度在傳播過程中被不斷放大這時(shí)應(yīng)優(yōu)先檢查是否有層被初始化的方差過大或者學(xué)習(xí)率是不是太高。還有一個(gè)應(yīng)該養(yǎng)成的好習(xí)慣在第一次backward之后順手打印一下全模型的梯度范數(shù)總和total_grad_norm sum(p.grad.norm().item() ** 2 for p in model.parameters() if p.grad is not None) ** 0.5 print(ftotal grad norm: {total_grad_norm:.3f})總梯度范數(shù)穩(wěn)定在一個(gè)合理范圍比如個(gè)位數(shù)到幾十訓(xùn)練通常健康。如果這個(gè)值在1e-4以下或1e4以上基本可以斷定初始化或?qū)W習(xí)率配置有問題。4.3 三個(gè)我踩過且容易復(fù)現(xiàn)的坑先說第一個(gè)坑在自定義模塊里創(chuàng)建層時(shí)忘了調(diào)用reset_parameters或做任何初始化。PyTorch對nn.Linear、nn.Conv這些內(nèi)置層會(huì)自動(dòng)初始化但如果你繼承autograd.Function或者手動(dòng)創(chuàng)建Parameter例如class MyLayer(nn.Module): def __init__(self, in_dim, out_dim): super().__init__() self.weight nn.Parameter(torch.empty(in_dim, out_dim))然后忘了給weight賦值這里面的內(nèi)存數(shù)據(jù)就會(huì)是“未定義”的隨機(jī)垃圾值可能是天大的正數(shù)也可能是NaN。正確的做法是立刻做初始化nn.init.kaiming_uniform_(self.weight, amath.sqrt(5))第二個(gè)坑用了Xavier初始化配ReLU激活。深層網(wǎng)絡(luò)里前向信號被ReLU不斷砍半后向梯度也跟著衰減幾十層之后幾乎傳不回去。典型表現(xiàn)就是深層CNN訓(xùn)練特別慢淺層權(quán)重梯度幾乎為零我實(shí)際項(xiàng)目中栽過跟頭換Kaiming之后收斂速度天壤之別。第三個(gè)坑Transformer場景下不看輸出方差直接用Kaiming初始化整個(gè)模型。Transformer里的注意力層涉及矩陣乘法和softmaxKaiming那套“為ReLU設(shè)計(jì)”的理論在這里并不完全適用。更合理的做法是配合更小的標(biāo)準(zhǔn)差比如0.02或者干脆依賴位置編碼和LayerNorm的配合。前兩年我調(diào)過一個(gè)小的Transformer用Xavier初始化注意力權(quán)重訓(xùn)練初期整個(gè)注意力分布全成了one-hotloss瘋狂震蕩改成0.02標(biāo)準(zhǔn)差后平穩(wěn)了許多。這說明初始化不能脫離具體結(jié)構(gòu)選型不要死板。提示排查初始化問題的正確順序是先看初始loss是否處于合理區(qū)間再看第一輪訓(xùn)練后的梯度范數(shù)是否健康最后才調(diào)整學(xué)習(xí)率和優(yōu)化器參數(shù)。別一上來就亂試超參數(shù)那會(huì)同時(shí)掩蓋多個(gè)問題。5. 不同網(wǎng)絡(luò)結(jié)構(gòu)的初始化選型與實(shí)戰(zhàn)心得到了這一節(jié)我想把實(shí)踐中的經(jīng)驗(yàn)匯總一下給出可以直接抄作業(yè)的選型建議和幾個(gè)心法。5.1 初始化選型速查表一個(gè)相對穩(wěn)妥的參考配置如下網(wǎng)絡(luò)結(jié)構(gòu)常用初始化偏置/特殊處理MLPReLUkaiming_normal_fan_inbias0MLPtanhxavier_normal_或xavier_uniform_bias0CNNReLUkaiming_normal_fan_inconv bias0BN γ1 β0LSTM/RNN正交初始化或xavier_uniform_遺忘門bias偏大如1.0Transformer各層標(biāo)準(zhǔn)差可取0.02或xavier_uniform_需配合LayerNorm和warmup遷移學(xué)習(xí)微調(diào)繼承預(yù)訓(xùn)練參數(shù)不對主干重置新增分類頭用較小隨機(jī)值這張表偏保守但能保證你在大多數(shù)任務(wù)里起步不翻車。這里多說一句遷移學(xué)習(xí)微調(diào)的場景。很多人從頭訓(xùn)練一個(gè)在預(yù)訓(xùn)練模型基礎(chǔ)上加分類頭的網(wǎng)絡(luò)時(shí)會(huì)習(xí)慣性地調(diào)用model.apply把所有層重置一遍結(jié)果把預(yù)訓(xùn)練權(quán)重全沖掉了。這是非常痛的教訓(xùn)。遷移學(xué)習(xí)的正確做法是預(yù)訓(xùn)練主干保持不動(dòng)只在新增層上做溫和的隨機(jī)初始化。因?yàn)轭A(yù)訓(xùn)練權(quán)重已經(jīng)蘊(yùn)含了良好的特征表示重新初始化等于毀掉這份資產(chǎn)。5.2 初始化與學(xué)習(xí)率、正則化的聯(lián)動(dòng)關(guān)系初始化從來不是一個(gè)孤立變量。我慢慢發(fā)現(xiàn)它和學(xué)習(xí)率之間存在明顯的“聯(lián)席效應(yīng)”如果初始化方差偏大即使學(xué)習(xí)率很小訓(xùn)練過程也可能震蕩如果初始化方差偏小又需要更大的學(xué)習(xí)率來彌補(bǔ)初期梯度過小的問題。所以調(diào)整初始化的時(shí)候要有意識(shí)地去配合學(xué)習(xí)率。我的習(xí)慣是先鎖定一種合理的初始化方法再把學(xué)習(xí)率放在一個(gè)中間值比如1e-3或3e-4然后只動(dòng)一個(gè)變量確認(rèn)效果后再動(dòng)另一個(gè)。很多人喜歡同時(shí)改一堆超參數(shù)到最后出了問題根本不知道是誰的鍋。初始化對正則化也有微妙的影響。權(quán)重初始方差越大相當(dāng)于模型一開始的“帶寬”越寬隱式的正則效果越強(qiáng)但也更容易過擬合或產(chǎn)生梯度問題。設(shè)置初始化時(shí)心里要有這根弦尤其在數(shù)據(jù)量不大的任務(wù)里更傾向使用偏小的方差。5.3 一些值得留意的細(xì)節(jié)習(xí)慣我在實(shí)際工作中養(yǎng)成了幾個(gè)和初始化相關(guān)的習(xí)慣分享給大家。建模型文件的時(shí)候我會(huì)在旁邊放一個(gè)自定義init_weights函數(shù)把所有初始化規(guī)則集中在一起。不管模型最后搭成什么樣apply一遍就到位。打印模型第一輪loss的時(shí)候我會(huì)順帶打印一下各層輸出的均值和標(biāo)準(zhǔn)差。如果某一層輸出的std突然比其他層大兩個(gè)數(shù)量級說明那一層初始化或結(jié)構(gòu)設(shè)計(jì)有問題趁早看比等loss曲線半小時(shí)后再后悔強(qiáng)多了。保存模型checkpoint的時(shí)候把模型結(jié)構(gòu)以及是否自定義過初始化一起寫在配置文件里省得幾個(gè)月后自己看著權(quán)重文件發(fā)呆不知道當(dāng)時(shí)的初始化策略是什么。最后如果遇到實(shí)在折騰不明白的不收斂問題不妨回到原點(diǎn)做一次“重新初始化zéro學(xué)習(xí)率測試”把學(xué)習(xí)率暫時(shí)設(shè)為0跑一步看看loss是否確定。如果學(xué)習(xí)率為0時(shí)loss都不穩(wěn)定那基本就是前向傳播或者初始化的問題而不是優(yōu)化器和反向傳播的問題。這個(gè)排查順序能幫你省下大量瞎猜的時(shí)間。結(jié)尾關(guān)于參數(shù)初始化這件事我現(xiàn)在的態(tài)度是它是整個(gè)訓(xùn)練流程里性價(jià)比最高的一環(huán)——改幾行代碼就能避免幾小時(shí)甚至幾天的無效訓(xùn)練。剛?cè)腴T的朋友一定要親手打印幾次權(quán)重分布看看不同初始化方案下數(shù)據(jù)的量級差異有經(jīng)驗(yàn)的朋友則可以花點(diǎn)時(shí)間讀一讀Glorot和He的兩篇經(jīng)典論文再回到PyTorch源碼里對照一下默認(rèn)實(shí)現(xiàn)那種“原來如此”的頓悟感是單純調(diào)參給不了的。如果這篇能幫你少走一次彎路那我就沒白寫。下一篇我打算聊聊學(xué)習(xí)率調(diào)度和優(yōu)化器選擇的聯(lián)動(dòng)問題那又是一個(gè)同樣容易被低估的坑。