戰(zhàn):從loss不收斂到顯存OOM的完整優(yōu)化方案)
先說一個(gè)讓人頭疼的場(chǎng)景我在訓(xùn)練一個(gè)中文生成模型batch size開到16就OOM開到8倒是穩(wěn)了但一個(gè)epoch要跑將近兩天。更氣人的是loss曲線前半段像心電圖后半段像便秘——要么不降要么突然跳水然后又彈回去。當(dāng)時(shí)我就明白一件事光把模型搭出來遠(yuǎn)遠(yuǎn)不夠真正決定項(xiàng)目能不能落地的是優(yōu)化這層功夫。后來我花了很長(zhǎng)時(shí)間折騰Model-Optimizer把訓(xùn)練過程中的梯度更新、學(xué)習(xí)率調(diào)度、顯存分配、精度策略這些環(huán)節(jié)一個(gè)一個(gè)拆開來看最終沉淀出一套可以復(fù)用的優(yōu)化框架。這篇文章不聊論文式的理論推導(dǎo)只聊我在實(shí)際項(xiàng)目中踩過的坑、驗(yàn)證過的配置、以及最終跑通的方案。如果你正在被loss不收斂、顯存不夠、訓(xùn)練太慢這些問題折磨這篇文章應(yīng)該能給你一些直接可抄的作業(yè)。1. 為什么我要自己折騰一個(gè)Model-Optimizer1.1 現(xiàn)成優(yōu)化器解決不了的實(shí)際問題先說結(jié)論P(yáng)yTorch自帶的AdamW、SGD這些優(yōu)化器本身沒做錯(cuò)什么但它們只是梯度更新規(guī)則這一個(gè)環(huán)節(jié)。實(shí)際訓(xùn)練一個(gè)稍微像樣的模型你會(huì)發(fā)現(xiàn)瓶頸根本不在這一個(gè)環(huán)節(jié)上。我當(dāng)時(shí)遇到的是三個(gè)具體問題優(yōu)化器狀態(tài)占用顯存過高。AdamW需要維護(hù)一階動(dòng)量和二階動(dòng)量這兩個(gè)緩存張量跟模型參數(shù)一樣大。一個(gè)7B參數(shù)的模型光優(yōu)化器狀態(tài)就得多吃將近56GB顯存按BF16算這還沒算梯度和激活值。學(xué)習(xí)率調(diào)度器和優(yōu)化器狀態(tài)脫節(jié)。模型中途恢復(fù)訓(xùn)練時(shí)如果只保存了模型權(quán)重而沒有保存優(yōu)化器的step計(jì)數(shù)和動(dòng)量狀態(tài)重新加載后學(xué)習(xí)率會(huì)跳回初始值訓(xùn)練節(jié)奏直接廢掉。loss曲線反復(fù)震蕩。batch size小的時(shí)候梯度噪聲大AdamW雖然能自適應(yīng)調(diào)整更新幅度但遇到個(gè)別梯度特別大的batch還是會(huì)出現(xiàn)loss尖刺。Model-Optimizer這個(gè)項(xiàng)目的出發(fā)點(diǎn)就是把優(yōu)化器從單一的參數(shù)更新算法擴(kuò)展成一套包含梯度處理、學(xué)習(xí)率策略、顯存優(yōu)化、狀態(tài)管理在內(nèi)的完整訓(xùn)練優(yōu)化解決方案。1.2 Model-Optimizer的設(shè)計(jì)目標(biāo)我在設(shè)計(jì)這個(gè)框架時(shí)給自己定了幾個(gè)原則模塊化組合每個(gè)優(yōu)化組件可以獨(dú)立開關(guān)和替換比如梯度裁剪可以單獨(dú)用也可以和混合精度配合??捎^測(cè)性優(yōu)先訓(xùn)練過程中的學(xué)習(xí)率、梯度范數(shù)、參數(shù)更新幅度等關(guān)鍵指標(biāo)必須能實(shí)時(shí)看到??床灰娋驼{(diào)不了。狀態(tài)可恢復(fù)任何時(shí)候中斷訓(xùn)練重新加載后應(yīng)該能精確恢復(fù)到中斷時(shí)的狀態(tài)包括學(xué)習(xí)率的step位置。很多人覺得模塊化這些是老生常談但真到自己寫訓(xùn)練腳本時(shí)全都圖省事直接調(diào)torch的默認(rèn)API結(jié)果出了問題只能干瞪眼。1.3 這個(gè)框架適合誰(shuí)如果你只是跑跑MNIST、CIFAR這種玩具數(shù)據(jù)集那完全不需要折騰這些。但如果你在訓(xùn)練GPT風(fēng)格的語(yǔ)言模型、Diffusion模型、或者任意超過1B參數(shù)的大模型你就會(huì)發(fā)現(xiàn)這里面的每個(gè)細(xì)節(jié)都在影響最終效果。我做這個(gè)項(xiàng)目時(shí)的基準(zhǔn)場(chǎng)景是單機(jī)多卡訓(xùn)練一個(gè)3B參數(shù)的對(duì)話模型這也是我認(rèn)為Model-Optimizer最適用的場(chǎng)景。2. 模型優(yōu)化器里的三大關(guān)鍵旋鈕2.1 梯度更新算法怎么選這不是三言兩語(yǔ)能說清的但我可以分享自己的選型經(jīng)驗(yàn)。首先是AdamW和SGD的對(duì)比。SGD配合momentum在CV領(lǐng)域一直表現(xiàn)穩(wěn)定尤其經(jīng)過長(zhǎng)時(shí)間訓(xùn)練后泛化性往往更好。但SGD對(duì)學(xué)習(xí)率太敏感需要在訓(xùn)練過程中頻繁調(diào)整而且自適應(yīng)能力弱遇到稀疏特征或者不同尺度的參數(shù)時(shí)表現(xiàn)不太穩(wěn)定。AdamW在NLP和生成模型里幾乎成了默認(rèn)選擇。它的優(yōu)勢(shì)是每個(gè)參數(shù)都有獨(dú)立的學(xué)習(xí)率縮放前期的訓(xùn)練速度明顯更快。但AdamW也有個(gè)很多人不知道的問題它對(duì)權(quán)重衰減的處理雖然比原始Adam規(guī)范但如果你把weight_decay設(shè)得太高比如0.1以上參數(shù)范數(shù)會(huì)被壓得特別小最終模型的表達(dá)能力會(huì)下降。我自己在Model-Optimizer里的默認(rèn)配置是優(yōu)化器: AdamW beta1: 0.9 beta2: 0.95 epsilon: 1e-8 weight_decay: 0.01這個(gè)配置在大多數(shù)語(yǔ)言模型任務(wù)上表現(xiàn)都比較穩(wěn)。beta2從默認(rèn)的0.999調(diào)低到0.95是我實(shí)測(cè)下來很管用的一個(gè)改動(dòng)——它讓二階動(dòng)量對(duì)梯度變化的響應(yīng)更快能明顯減少訓(xùn)練后期的loss尖刺問題。2.2 學(xué)習(xí)率調(diào)度的魔鬼細(xì)節(jié)學(xué)習(xí)率調(diào)度表面上看只是一個(gè)衰減曲線實(shí)際里面全是坑。我踩得最深的一個(gè)坑是warmup步數(shù)設(shè)置。剛開始訓(xùn)練3B模型時(shí)我沒加warmup直接用1e-4的學(xué)習(xí)率開跑。結(jié)果前500步loss不僅沒降反而從5.2漲到了5.8梯度范數(shù)一度飆到正常值的30倍。后來查資料才意識(shí)到模型參數(shù)剛初始化時(shí)分布不理想梯度統(tǒng)計(jì)量也不穩(wěn)定此時(shí)直接上大學(xué)習(xí)率會(huì)讓AdamW的動(dòng)量估計(jì)迅速偏移后面需要用很多步才能糾正回來。正確做法是加一個(gè)線性warmup讓學(xué)習(xí)率從0逐步升到目標(biāo)值。這個(gè)階段的主要作用是預(yù)熱優(yōu)化器的動(dòng)量狀態(tài)而不是真正學(xué)習(xí)。我在Model-Optimizer里的推薦配置學(xué)習(xí)率峰值: 3e-43B模型 warmup步數(shù): 總步數(shù)的1%約500步 衰減策略: cosine退火到峰值的1/10峰值學(xué)習(xí)率的選擇是另一個(gè)大頭。它和下一條直接相關(guān)——優(yōu)化器的更新幅度取決于學(xué)習(xí)率和梯度范數(shù)的乘積。我在實(shí)踐中的一個(gè)經(jīng)驗(yàn)是如果訓(xùn)練中出現(xiàn)loss前期不降不要盲目加大學(xué)習(xí)率先看看梯度范數(shù)。如果梯度范數(shù)本身在1e-2這個(gè)量級(jí)上下浮動(dòng)那3e-4的學(xué)習(xí)率是合理的如果梯度范數(shù)只有1e-3說明梯度太小考慮去掉梯度裁剪或者調(diào)整網(wǎng)絡(luò)初始化方式而不是去調(diào)學(xué)習(xí)率。2.3 混合精度和梯度累積的配合這兩兄弟配合好了能大幅提升訓(xùn)練效率配合不好會(huì)讓你懷疑人生。混合精度用的是PyTorch的torch.cuda.amp.autocast和GradScaler。核心邏輯是前向和反向計(jì)算用FP16加速但優(yōu)化器更新保留FP32的主權(quán)重副本同時(shí)用動(dòng)態(tài)loss scaling避免FP16下梯度過小被下溢吞掉。我在實(shí)際使用中踩過一次動(dòng)態(tài)scale失靈的問題。訓(xùn)練到第2000步時(shí)loss突然變成NaN程序卻沒報(bào)錯(cuò)。排查后發(fā)現(xiàn)是GradScaler的scale因子在反復(fù)迭到overflow后自動(dòng)變小但我的某個(gè)層在FP16下梯度一直下溢導(dǎo)致該層的權(quán)重長(zhǎng)時(shí)間不更新數(shù)值漂移越積越大最后整個(gè)模型崩了。解決辦法是給Model-Optimizer加了兩層保險(xiǎn)對(duì)特別容易下溢的層比如attention里的softmax和layer norm后的全連接層單獨(dú)走FP32計(jì)算不參與混合精度。作法是在模型前向里用with autocast(enabledFalse):包住這些層。實(shí)時(shí)檢測(cè)每個(gè)參數(shù)的梯度更新量如果連續(xù)100步某個(gè)參數(shù)組的梯度范數(shù)為0就告警提示。梯度累積的設(shè)置相對(duì)簡(jiǎn)單但有一個(gè)容易被忽略的點(diǎn)梯度累積會(huì)改變實(shí)際batch size進(jìn)而影響loss的尺度。如果你的累積步數(shù)是4那等效batch size就是單卡batch size乘4再乘卡數(shù)。模型更新一步時(shí)的梯度是所有微批次梯度的平均這時(shí)的學(xué)習(xí)率理論上也需要相應(yīng)放大。我實(shí)測(cè)的通用做法是梯度累積步數(shù)和學(xué)習(xí)率不聯(lián)動(dòng)保持學(xué)習(xí)率不變但warmup的步數(shù)可以適當(dāng)增多因?yàn)榈刃atch變大了梯度統(tǒng)計(jì)更穩(wěn)定warmup階段可以更平滑。3. 從loss曲線倒推優(yōu)化器配置問題的排查鏈路這是Model-Optimizer項(xiàng)目里最讓我覺得有價(jià)值的部分也是我踩坑最多的地方。調(diào)優(yōu)化器不能靠感覺要有系統(tǒng)的排查路徑。3.1 現(xiàn)象一loss前期紋絲不動(dòng)有次訓(xùn)練多模態(tài)模型前300步loss一直在5.4附近徘徊小數(shù)點(diǎn)后都看不出變化。我當(dāng)時(shí)的第一反應(yīng)是調(diào)大學(xué)習(xí)率但沒急著動(dòng)手。先查了三個(gè)東西第一個(gè)是梯度范數(shù)。打印出來發(fā)現(xiàn)是1.5e-4小得離譜這解釋了為什么參數(shù)更新幾乎為0。但梯度為什么這么小第二個(gè)是loss的絕對(duì)值。5.4對(duì)應(yīng)的是交叉熵還是MSE如果是交叉熵一個(gè)詞表大小為32k的模型隨機(jī)初始化的loss大約是log(32768)10.45.4已經(jīng)比隨機(jī)好很多了。這說明模型已經(jīng)學(xué)了一些東西只是速度慢。第三個(gè)是輸入輸出的數(shù)據(jù)分布。最后發(fā)現(xiàn)是某個(gè)預(yù)訓(xùn)練特征提取器把梯度傳到后面時(shí)幾乎衰減沒了——問題出在網(wǎng)絡(luò)連接方式而不是優(yōu)化器。這個(gè)排查給我留下的經(jīng)驗(yàn)是loss不降先看梯度再看loss的絕對(duì)水平最后看數(shù)據(jù)是否有效不要一上來就動(dòng)學(xué)習(xí)率。3.2 現(xiàn)象二訓(xùn)練中期loss尖刺訓(xùn)練到總進(jìn)度的40%左右loss曲線每隔幾百步就跳出個(gè)尖刺高了0.2-0.3。這種情況通常不是隨機(jī)噪聲而是某種可復(fù)現(xiàn)的系統(tǒng)性異常。我在Model-Optimizer里加了一個(gè)診斷工具當(dāng)單步loss超過該batch之前100步的平均值2倍以上時(shí)自動(dòng)記錄當(dāng)時(shí)的學(xué)習(xí)率、梯度范數(shù)、以及l(fā)oss最大的樣本對(duì)應(yīng)的樣本ID。后來發(fā)現(xiàn)尖刺往往集中在某些語(yǔ)義模糊的訓(xùn)練樣本上它們的特征是包含大量生僻詞或者超長(zhǎng)文本。處理方式有兩個(gè)層面對(duì)優(yōu)化器層面梯度裁剪是必須的。我用的max_grad_norm1.0把所有參數(shù)的梯度范數(shù)限制在這個(gè)值內(nèi)。注意這不等同于把每個(gè)參數(shù)clip到絕對(duì)值1.0效果差很多。對(duì)數(shù)據(jù)層面這種異常樣本即使梯度被裁剪仍然會(huì)污染模型的狀態(tài)最好是在數(shù)據(jù)預(yù)處理階段就把這類樣本單獨(dú)分桶或者降低其采樣權(quán)重。3.3 現(xiàn)象三訓(xùn)練集收斂但驗(yàn)證集差這是一個(gè)典型的泛化問題但我發(fā)現(xiàn)很多人把它歸咎為過擬合后就結(jié)束了沒有往優(yōu)化器配置上想。實(shí)際上優(yōu)化器的某些配置會(huì)明顯影響模型的泛化能力。我對(duì)比過同一模型在相同數(shù)據(jù)下的兩組實(shí)驗(yàn)配置項(xiàng)實(shí)驗(yàn)A實(shí)驗(yàn)Bweight_decay0.010.1最終loss2.12.3驗(yàn)證集準(zhǔn)確率68%72%梯度噪聲中等低實(shí)驗(yàn)B的訓(xùn)練loss更高但驗(yàn)證集表現(xiàn)更好。這不是巧合。weight_decay本質(zhì)上是給模型參數(shù)加了一個(gè)L2正則項(xiàng)約束參數(shù)范數(shù)不至于過大從而讓模型對(duì)訓(xùn)練集的特定噪聲不敏感。所以如果你發(fā)現(xiàn)驗(yàn)證集和訓(xùn)練集差距過大先別急著加dropout試著把weight_decay從0.01提上去也許效果更直接。3.4 一整套排查順序我把這段時(shí)間的排查經(jīng)驗(yàn)整理成一套固定順序現(xiàn)在調(diào)任何模型的優(yōu)化器配置都按照這個(gè)來看loss的初始值確認(rèn)它是否符合隨機(jī)初始化的理論預(yù)期。不符合先查數(shù)據(jù)管道和loss計(jì)算邏輯??辞?00步的梯度范數(shù)曲線。如果梯度范數(shù)趨近于0問題大概率在模型結(jié)構(gòu)或數(shù)據(jù)喂入方式而不是優(yōu)化器。看warmup結(jié)束后400步內(nèi)的loss變化趨勢(shì)。如果loss立刻上升考慮是通過降低峰值學(xué)習(xí)率或者增加warmup步數(shù)來緩解。如果訓(xùn)練后期出現(xiàn)尖刺檢查是否是特定batch導(dǎo)致的并考慮梯度裁剪和數(shù)據(jù)清洗。對(duì)比驗(yàn)證集指標(biāo)和訓(xùn)練集指標(biāo)的gap如果gap過大優(yōu)先調(diào)節(jié)weight_decay再考慮數(shù)據(jù)增強(qiáng)或dropout。這套排查鏈路不需要用到什么高級(jí)工具核心就是log好每一步的關(guān)鍵指標(biāo)。我在Model-Optimizer里默認(rèn)記錄了以下字段step、loss、lr、grad_norm、update_norm、loss_scale、顯存占用、當(dāng)前batch的樣本平均長(zhǎng)度。有了這些每次出問題都能對(duì)照歷史曲線快速定位。4. 顯存和吞吐的平衡術(shù)優(yōu)化器這個(gè)東西看起來只是更新參數(shù)的算法但實(shí)際上它站在顯存和吞吐的交叉點(diǎn)上。4.1 優(yōu)化器狀態(tài)本身就在吃顯存說個(gè)具體數(shù)字。我訓(xùn)練3B模型參數(shù)占用約6GBBF16如果不用任何顯存優(yōu)化完整訓(xùn)練狀態(tài)包括狀態(tài)項(xiàng)占用BF16/FP32模型參數(shù)約6GB梯度約6GBAdamW一階動(dòng)量約12GBFP32AdamW二階動(dòng)量約12GBFP32激活值動(dòng)態(tài)數(shù)GB到數(shù)十GB不等看到?jīng)]優(yōu)化器狀態(tài)占的顯存是模型參數(shù)本身的4倍。這也是為什么大模型訓(xùn)練框架里優(yōu)化器狀態(tài)往往是最先被優(yōu)化的對(duì)象。Model-Optimizer提供了兩種降顯存方案Adafactor替代AdamW。它的核心思路是只保存參數(shù)的逐行和逐列二階統(tǒng)計(jì)量而不是每個(gè)參數(shù)的完整二階動(dòng)量。顯存占用比AdamW減少約60%效果在部分任務(wù)上略有折扣但差距在可接受范圍內(nèi)。優(yōu)化器狀態(tài)offload到CPU。利用DeepSpeed的Zero-Offload類似思路把優(yōu)化器狀態(tài)放CPU內(nèi)存GPU只保留權(quán)重和梯度。我實(shí)測(cè)這個(gè)方法能在單卡上訓(xùn)練原本需要雙卡的模型代價(jià)是訓(xùn)練速度下降約30%。4.2 梯度檢查點(diǎn)和激活值重計(jì)算的取舍激活值的顯存占用經(jīng)常被忽略但實(shí)際上對(duì)3B模型來說激活值才是那個(gè)可能直接壓垮顯存的元兇。我遇到過一個(gè)典型情況batch size設(shè)為4時(shí)顯存剛好夠設(shè)5就OOM加的那一個(gè)batch把激活值的占用推到了極限。梯度檢查點(diǎn)的思路是前向傳播時(shí)不要保存所有激活值只保存關(guān)鍵的幾個(gè)錨點(diǎn)反向傳播時(shí)再?gòu)倪@些錨點(diǎn)重新計(jì)算需要的激活值。這個(gè)方案能把激活值顯存下降好幾倍但會(huì)讓訓(xùn)練時(shí)間增加約20%-30%。我個(gè)人的建議是先檢查你的模型是否能通過調(diào)整batch size和平行策略來規(guī)避OOM。如果實(shí)在繞不開再使用梯度檢查點(diǎn)而且只對(duì)最耗顯存的幾個(gè)模塊開啟不要全模型無(wú)腦開。4.3 batch size與吞吐的實(shí)測(cè)數(shù)據(jù)很多人以為batch size越大吞吐越高其實(shí)不完全是。我用3B模型做了幾組對(duì)比測(cè)試配置吞吐樣本/秒顯存峰值batch4無(wú)檢查點(diǎn)18.238GBbatch8無(wú)檢查點(diǎn)20.1溢出batch8梯度檢查點(diǎn)13.436GBbatch8激活重計(jì)算優(yōu)化16.839GBbatch4梯度累積2步17.938GB關(guān)鍵數(shù)據(jù)是最后一行用batch4加上梯度累積2步效果等效于batch8吞吐幾乎沒有下降顯存也沒有變大。原因很簡(jiǎn)單——梯度累積規(guī)避了激活值的峰值而吞吐瓶頸主要在計(jì)算單元不在batch size本身。結(jié)論是當(dāng)顯存有限時(shí)優(yōu)先用梯度累積不要優(yōu)先開梯度檢查點(diǎn)。5. 這套優(yōu)化策略在不同任務(wù)上的實(shí)測(cè)表現(xiàn)5.1 圖像分類任務(wù)ResNet-50ImageNet子集Model-Optimizer的第一個(gè)實(shí)測(cè)場(chǎng)景是一個(gè)經(jīng)典圖像分類任務(wù)。我拿它跟一個(gè)固定lr0.1的SGDmomentum配置做對(duì)比。兩組都訓(xùn)練90個(gè)epoch。Model-Optimizer用的是默認(rèn)的AdamWcosine退火學(xué)習(xí)率從0.001線性warmup到0.01再退火。結(jié)果很有趣SGD組在訓(xùn)練集上準(zhǔn)確率和AdamW組接近但驗(yàn)證集上SGD高了0.8個(gè)百分點(diǎn)。這和SGD本身的隱式正則有關(guān)。這個(gè)結(jié)果提醒我Model-Optimizer不能無(wú)腦用同一套配置。我在框架里加入了按任務(wù)類型切換優(yōu)化器預(yù)設(shè)的功能CV分類任務(wù)默認(rèn)切回SGDmomentum文本和生成任務(wù)才用AdamW。5.2 語(yǔ)言模型訓(xùn)練3B GPT風(fēng)格模型這是Model-Optimizer的主場(chǎng)。和原始配置AdamWlr3e-4無(wú)warmup相比加了線性warmup和cosine退火后的配置在15B token的訓(xùn)練數(shù)據(jù)上最終困惑度從16.8降到了15.2。關(guān)鍵是訓(xùn)練過程幾乎沒有出現(xiàn)過NaN和尖刺模型的恢復(fù)點(diǎn)都快了很多。在訓(xùn)練過程中我還嘗試過把beta2進(jìn)一步從0.95降到0.9困惑度略微提升到15.4但訓(xùn)練穩(wěn)定性更好了。如果你對(duì)loss曲線的心電圖感非常介意可以試試這個(gè)改動(dòng)。5.3 推薦排序模型DIN架構(gòu)淘寶風(fēng)格數(shù)據(jù)推薦模型的特點(diǎn)是特征極度稀疏embedding表巨大而且正負(fù)樣本比例懸殊。這種場(chǎng)景下AdamW的表現(xiàn)中規(guī)中矩但embedding層的梯度更新存在頻繁抖動(dòng)的問題。我給Model-Optimizer加了按層學(xué)習(xí)率的功能embedding層的更新步長(zhǎng)是主網(wǎng)絡(luò)學(xué)習(xí)率的0.3倍梯度裁剪只作用于主網(wǎng)絡(luò)不對(duì)embedding層做額外裁剪。這個(gè)配置讓AUC從0.78提升到0.79且訓(xùn)練速度提升10%因?yàn)閑mbedding層的更新變慢后顯存中的緩存命中率反而變高了。6. 踩過這么多坑之后寫下的經(jīng)驗(yàn)筆記最后這部分不按章節(jié)來了就寫幾條我在實(shí)際使用中最想告訴后來者的話。關(guān)于學(xué)習(xí)率的調(diào)節(jié)一次只動(dòng)一個(gè)變量。我最開始調(diào)優(yōu)化器配置時(shí)經(jīng)常同時(shí)改學(xué)習(xí)率、weight_decay、beta2結(jié)果loss變好了也不知道是哪個(gè)改動(dòng)起的作用變差了也不知道該回退到哪個(gè)配置。后來強(qiáng)制自己每次只動(dòng)一個(gè)變量記錄在案效果才能復(fù)現(xiàn)。Model-Optimizer的配置文件里我保留了每次實(shí)驗(yàn)的完整變更記錄這個(gè)方法幫我躲過了大量返工。關(guān)于日志越詳細(xì)越好但別讓日志本身拖垮訓(xùn)練速度。我的做法是把關(guān)鍵指標(biāo)每50步打印一次同時(shí)把更細(xì)粒度的數(shù)據(jù)每500步寫一次文件。這樣既能看到實(shí)時(shí)變化又不至于產(chǎn)生海量文件。關(guān)于恢復(fù)訓(xùn)練優(yōu)化器狀態(tài)必須和模型權(quán)重同步保存。我之前吃過虧覺得保存模型就足夠了結(jié)果恢復(fù)訓(xùn)練后warmup重新走了一遍前1000步完全在浪費(fèi)計(jì)算資源。Model-Optimizer里直接打包保存了model_state_dict、optimizer_state_dict、scheduler_state_dict、step任何時(shí)刻中斷都能無(wú)縫續(xù)跑。最后一個(gè)不起眼但很實(shí)用的小技巧訓(xùn)練開始前先跑一個(gè)50步的熱身測(cè)試。用極小的數(shù)據(jù)量、極短的時(shí)間把訓(xùn)練循環(huán)完整跑通一遍重點(diǎn)看有沒有NaN、顯存是否夠、日志是否正常。這個(gè)熱身測(cè)試能幫你避免在正式訓(xùn)練跑了幾小時(shí)后才發(fā)現(xiàn)配置錯(cuò)誤這種最痛苦的情況。我每次新接一個(gè)模型都會(huì)先用這個(gè)方式確認(rèn)環(huán)境再長(zhǎng)跑訓(xùn)練。