網(wǎng)絡(luò):顯存優(yōu)化與輕量設(shè)計實戰(zhàn))
如果你手里只有一塊8GB顯存的普通顯卡又跟我一樣想訓練一個自己設(shè)計的神經(jīng)網(wǎng)絡(luò)那這篇東西大概率能幫你省下好幾個晚上的折騰。過去半年我一直在做一件事從零設(shè)計一個小規(guī)模神經(jīng)網(wǎng)絡(luò)并且把它開源出來。這個項目不求沖榜也不追求花哨結(jié)構(gòu)唯一的硬指標就是——普通顯卡能訓練能跑完能讓其他人照著代碼復現(xiàn)出來。項目代號叫NetLight屬于輕量圖像分類方向。文章里我會把結(jié)構(gòu)設(shè)計、顯存優(yōu)化、訓練配置和踩過的坑一次性講透適合正在用個人電腦做深度學習實驗的開發(fā)者。1. 顯存焦慮是起點普通顯卡到底被什么卡住了1.1 “模型小”和“顯存夠用”是兩碼事我在設(shè)計NetLight之前先在紙上估算過參數(shù)量模型大概3M參數(shù)心想這么小的網(wǎng)絡(luò)總該隨便跑了吧。結(jié)果第一次訓練就把8GB顯存撐爆了連第一個epoch都沒走完。后來我才意識到神經(jīng)網(wǎng)絡(luò)訓練時占用顯存的絕不只是模型參數(shù)而是四樣東西權(quán)重、優(yōu)化器狀態(tài)、梯度和中間激活值。其中中間激活值才是吃顯存的大戶。舉個直觀的例子假設(shè)某個特征圖是44×44×128batch size是32用單精度存儲光這一層就要占44×44×128×32×4約31MB。如果反向傳播時需要保存輸入用于梯度計算這個數(shù)字還會翻倍。而一個像樣的網(wǎng)絡(luò)里有幾十層這樣的特征圖層與層之間還要有臨時buffer顯存自然蹭蹭漲上去。所以個人開發(fā)者做自研網(wǎng)絡(luò)時第一課就是不能只盯著參數(shù)量必須算“激活顯存賬”。這也是NetLight設(shè)計時最重要的約束條件后面我會詳細說怎么算。1.2 現(xiàn)成模型倉庫為什么不能直接解決我的問題有人會問現(xiàn)成的輕量網(wǎng)絡(luò)已經(jīng)很多了直接拿預訓練權(quán)重微調(diào)不就行了嗎這話對一半。用預訓練模型做推理、做遷移學習確實方便但我想做的不只是拿模型填空而是想訓練一個自己能理解每層作用的自研網(wǎng)絡(luò)方便之后做結(jié)構(gòu)改動和消融實驗。另一方面很多開源倉庫的默認訓練配置非常“吃卡”。它們通常默認輸入分辨率224、batch size 128甚至256用多卡分布式訓練。普通8GB顯存顯卡跑第一輪就OOM你得先花半天去改配置、拆代碼才能讓它跑起來。與其在別人的設(shè)置里痛苦地找開關(guān)不如自己寫一個從頭到尾都可控的輕量訓練項目把“普通顯卡能訓練”作為第一優(yōu)先級寫進設(shè)計目標。2. 自研網(wǎng)絡(luò)的結(jié)構(gòu)設(shè)計先算顯存賬再定模塊2.1 用深度可分離卷積當骨架NetLight的骨干沒有用復雜的模塊而是把深度可分離卷積當作基礎(chǔ)單元。常規(guī)3×3卷積的計算量是輸入通道 × 輸出通道 × 9深度可分離卷積先把每個通道單獨做3×3卷積再用1×1卷積做通道融合參數(shù)和計算量都大幅下降。形象點說標準卷積像把所有衣服一起丟進洗衣機深度可分離卷積則是先按顏色分開洗再統(tǒng)一烘干效果接近但省水省電。具體結(jié)構(gòu)我做成了一張表方便你對照理解階段操作步長輸出尺寸通道數(shù)Stem3×3標準卷積288×8816Stage1深度可分離卷積殘差188×8832Stage2深度可分離卷積殘差244×4464Stage3深度可分離卷積殘差222×22128Stage4深度可分離卷積殘差211×11256Head全局平均池化全連接-1×1類別數(shù)整個模型參數(shù)量約2.8M輸入分辨率設(shè)為176而不是常用的224。很多人沒意識到224比176在面積上多了62%中間特征圖也跟著漲顯存自然壓不住。176×176對普通圖像分類任務(wù)來說信息量足夠但對顯存非常友好。2.2 給訓練過程算一筆顯存賬我建議所有想自研網(wǎng)絡(luò)的人都養(yǎng)成一個習慣拿到一個結(jié)構(gòu)先用公式粗算一下峰值顯存。簡化公式可以寫成激活顯存 ≈ batch size × 各層特征圖面積 × 通道數(shù) × 4字節(jié) × 反向傳播系數(shù)反向傳播系數(shù)通常取2因為前向特征圖存一份反向計算梯度時還要再訪問一次。NetLight第3階段的特征圖是22×22×128batch size取32時單層激活約31MB。雖然單層看起來不大但從Stage1到Stage4累加再加上Stem層的輸入輸出以及優(yōu)化器狀態(tài)等開銷整體峰值就進入GB級別了。這也是我把batch size默認設(shè)為32、輸入分辨率設(shè)為176的原因。這兩項直接卡住了顯存的大頭比換任何模型結(jié)構(gòu)都有效。加梯度累積之后等效batch size可以到64但顯存并不會翻倍后面會細說。2.3 沒有預訓練權(quán)重怎么保證訓練穩(wěn)定自研模型的另一個痛點是沒有公開預訓練權(quán)重必須從零開始訓練。很多人聽到“從零訓練”就慌其實只要結(jié)構(gòu)設(shè)計得收斂友好從零訓練完全可行。我在NetLight里做了三件事保證穩(wěn)定一是每個殘差分支都在卷積后、激活前加BatchNorm二是激活函數(shù)用ReLU6而不是普通ReLU輸出范圍有界配合BN更穩(wěn)定三是網(wǎng)絡(luò)深度只有4個Stage不會出現(xiàn)梯度消失。訓練時再用warmup和余弦學習率前面5輪學習率慢慢爬上去損失就不會亂跳。這樣設(shè)計的代價是表達能力不如大模型但換來的是個人開發(fā)者最需要的東西可訓練、可調(diào)試、可快速迭代。我做消融實驗時可以隨時刪除某個Stage或更換激活函數(shù)幾十分鐘后就能看到曲線變化這種掌控感是大模型倉庫給不了的。3. 訓練配方的“降顯存”打法不換卡也能跑更大的batch3.1 混合精度讓顯存占用降一個臺階NetLight默認開啟混合精度訓練。原理很簡單前向和反向計算時用半精度浮點數(shù)也就是fp16這樣中間激活值和梯度占用的顯存直接減半。與此同時模型主權(quán)重仍然用fp32保存避免精度損失。實際落地時要特別留意梯度下溢問題。fp16能表示的數(shù)值范圍比fp32小很多反向傳播時梯度稍微小一點就可能變成0所以主流實現(xiàn)會做loss scaling也就是把loss先放大若干倍等梯度算完再縮小回來。但自研網(wǎng)絡(luò)里最容易翻車的是BatchNorm。BN需要計算均值和方差在fp16下非常容易不穩(wěn)定。我在代碼里明確讓BN層的運算保持fp32這只增加了一點點計算量卻換來了穩(wěn)定收斂。訓練時如果發(fā)現(xiàn)loss在某個step突然變成NaN優(yōu)先檢查BN是不是被混進了fp16。3.2 梯度累積等效大batch真實小顯存梯度累積是我在小顯存顯卡上最常用的一招。它的思路是不一次性計算大batch的損失而是分成幾個小batch每個小batch正常算梯度暫存起來等累積到一定次數(shù)后再統(tǒng)一更新參數(shù)。核心偽代碼長這樣accum_steps 2 for i, batch in enumerate(loader): loss model(batch) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()注意這里必須把每個小batch的loss除以累積次數(shù)否則等效學習率會變大損失曲線容易出現(xiàn)尖刺。NetLight的默認配置是batch size32梯度累積2次等效batch size64但顯存峰值只按32算。有一點必須提醒如果網(wǎng)絡(luò)里有BatchNorm梯度累積的“等效batch增大”對BN并不成立。BN會在每個小batch內(nèi)部獨立更新統(tǒng)計量不會等累積后再算因此累積步數(shù)過大會讓BN的統(tǒng)計估計出現(xiàn)偏差。我實測下來累積2步影響很小但如果你的顯卡只有4GB顯存被迫把batch壓到8、累積4步BN最好改成同步BN或者換用不含BN的結(jié)構(gòu)。3.3 數(shù)據(jù)加載和訓練循環(huán)里的隱性顯存開銷除了模型本身訓練循環(huán)里還有幾個容易被忽略的顯存點。一是數(shù)據(jù)增強。把圖像縮放、翻轉(zhuǎn)、顏色抖動放到GPU上做雖然省CPU但會在顯存里新建大量臨時張量。NetLight的數(shù)據(jù)增強全部放在CPU端用多進程處理GPU只負責模型計算。這樣顯存干凈很多也不容易在第一個batch就飆到峰值。二是圖像解碼。很多人用普通顯卡訓練小數(shù)據(jù)集時發(fā)現(xiàn)GPU利用率很低一看CPU已經(jīng)拉滿原因是每個step都在重復解碼JPEG。我后來把訓練集圖像解碼成numpy數(shù)組緩存到內(nèi)存里顯存沒變但訓練速度快了將近40%。三是盡量開啟框架自帶的顯存優(yōu)化選項。如果是用動態(tài)圖框架通常有類似“不保存反向不需要的中間結(jié)果”的開關(guān)能進一步壓低峰值。我的經(jīng)驗是先關(guān)掉這些優(yōu)化跑通邏輯確認沒問題后再打開避免排查問題時多一層變量。4. 開源倉庫使用指南從克隆到跑通4.1 倉庫結(jié)構(gòu)一覽NetLight的代碼組織盡量保持簡單沒有做成臃腫的框架任何人下載下來都能看懂。目錄結(jié)構(gòu)大致是這樣netlight/ ├── configs/ │ └── default.yaml ├── models/ │ ├── __init__.py │ ├── blocks.py │ └── netlight.py ├── data/ ├── train.py ├── eval.py └── README.mdmodels目錄里只放模型定義blocks.py放深度可分離卷積和殘差模塊netlight.py負責組裝完整網(wǎng)絡(luò)。train.py包含訓練循環(huán)、混合精度、梯度累積和日志輸出。eval.py只做推理和精度統(tǒng)計不參與訓練所以哪怕你的顯卡再舊推理總是能跑的。環(huán)境要求也不復雜Python 3.8以上、一個主流動態(tài)圖深度學習框架、圖像處理庫和YAML解析庫。這些都是深度學習開發(fā)者的標配不需要額外裝奇怪的東西。4.2 準備數(shù)據(jù)和配置文件NetLight不挑數(shù)據(jù)格式最簡單的方式是準備一個train.txt每一行寫圖像路徑和對應(yīng)類別索引data/train/class0/sample_001.jpg 0 data/train/class0/sample_002.jpg 0 data/train/class1/sample_001.jpg 1如果你想直接按目錄結(jié)構(gòu)讀取代碼里也留了一個實現(xiàn)按類別子目錄自動生成標簽。默認配置文件長這樣input_size: 176 batch_size: 32 accum_steps: 2 epochs: 100 lr: 0.3 weight_decay: 4e-5 fp16: true num_workers: 4這里的lr: 0.3看著嚇人其實配合warmup和余弦退火完全沒有問題。如果把lr設(shè)成常見的0.001反而會因為優(yōu)化器狀態(tài)和梯度尺度不匹配在普通顯卡上訓練得很慢。我的建議是先按這個默認配置跑通再去調(diào)自己的學習率。4.3 一條命令啟動訓練和驗證訓練過程非常簡單打開終端進入倉庫根目錄執(zhí)行python train.py --config configs/default.yaml訓練結(jié)束后用python eval.py --checkpoint output/best.pth就能得到驗證集上的Top-1準確率。訓練過程中日志會實時顯示當前epoch、loss和驗證準確率我不習慣用花哨的可視化工具把曲線畫出來反而干擾判斷看數(shù)字就夠了。默認配置下100個epoch大約需要2小時出頭前提是你的顯卡和我一樣是8GB顯存的普通型號。如果你的顯卡顯存稍小也不需要改代碼直接改YAML里的batch_size即可但要注意學習率最好也按比例縮放。簡單經(jīng)驗是batch減半學習率也減半。4.4 如何判斷你的顯卡能不能“吃下”這個項目我收到過不少類似問題我的顯卡是XX能不能跑其實不用問別人跑一下就知道。在正式訓練之前把batch_size臨時改成8跑幾個step看顯存峰值然后再按比例往上加。NetLight在8GB顯存下跑到batch_size32完全沒有壓力峰值顯存大約5.4GB6GB顯存可以降到batch_size16或把input_size改為1604GB顯存則需要同時把input_size降到144、batch_size降到8并把梯度累積提高到4。我給了張配置參考表方便不同顯存的人直接抄作業(yè)顯存input_sizebatch_sizeaccum_stepsfp168GB176322開啟6GB176162開啟4GB16084開啟需要注意的是顯存和顯卡算力并不完全等價。8GB老卡可能算力弱訓練時間長一點但只要能跑通個人實驗的目的就達到了。NetLight設(shè)計的初衷就是讓這種“不夠頂級”的硬件也能完成從零訓練的完整閉環(huán)。5. 實測效果與翻車記錄我在這張普通顯卡上踩過的坑5.1 普通顯卡上的真實結(jié)果我在一個10類圖像分類數(shù)據(jù)集上做了完整測試訓練集大概5萬張圖從零訓練100個epoch最終Top-1準確率約91.7%。這個數(shù)字對輕量模型來說中規(guī)中矩但重要的是整個訓練過程在8GB普通顯卡上穩(wěn)定跑完單epoch約40秒總耗時不到2小時中間沒有一次OOM。顯存方面我做了對比如果不開啟混合精度batch_size32到了第10個epoch附近大概率會OOM開啟混合精度后峰值降到5.4GB再加上梯度累積整個訓練過程非常從容。我的建議是fp16永遠開著即使你的顯卡支持得不算好至少能讓顯存余量多出一截。5.2 三個讓代碼返工的坑第一個坑是混合精度下的BN溢出。最早我把整個網(wǎng)絡(luò)都切成fp16前幾個epoch一切正常到第30個epoch左右loss突然變成NaN。排查了一晚上最后發(fā)現(xiàn)問題出在BN的方差計算上。fp16一旦遇到某些分布比較極端的特征圖方差會溢出。解決辦法是把BN的輸入轉(zhuǎn)成fp32做統(tǒng)計再轉(zhuǎn)回fp16繼續(xù)后續(xù)計算這之后再也沒有出現(xiàn)過NaN。第二個坑是梯度累積導致BN統(tǒng)計值漂移。我試過在batch_size8、accum_steps8的情況下訓練損失曲線很漂亮但驗證準確率像過山車一樣抖。原因是BN在太小的小batch里估計統(tǒng)計量數(shù)據(jù)多樣性不夠。最后我選擇保留batch_size32、accum_steps2的組合用增大真實batch來保證BN穩(wěn)定而不是無限依賴累積。第三個坑和數(shù)據(jù)加載有關(guān)。剛開始訓練時GPU利用率只有30%多顯卡明顯沒吃飽CPU卻跑滿了。我以為是模型結(jié)構(gòu)太輕導致算力過剩后來才發(fā)現(xiàn)是數(shù)據(jù)增強和JPEG解碼都在主線程里跑成了瓶頸。把數(shù)據(jù)預處理移到獨立進程并做緩存之后GPU利用率才升到80%以上。對輕量網(wǎng)絡(luò)來說數(shù)據(jù)加載對總耗時的影響遠比想象中大得多。5.3 開源之后的一點體會做完這個項目之后我最大的感受是“可復現(xiàn)”三個字比“效果好”更重要。發(fā)布開源項目時我特意在README里寫了測試用的顯卡顯存、Python版本、隨機種子和處理后的數(shù)據(jù)格式。沒有這些信息別人下載代碼后復現(xiàn)不了第一反應(yīng)通常是懷疑你的代碼有問題實際上只是環(huán)境差異造成的。如果你也想開源自己的自研網(wǎng)絡(luò)我建議先跑一遍完整的“新手流程”用一個公開小數(shù)據(jù)集從零環(huán)境開始照著README操作看能不能一步步跑到最終結(jié)果。我平時習慣先開一個3個epoch的快速調(diào)試模式把耗時的部分全部縮小確認代碼沒有低級錯誤后再跑完整實驗。普通顯卡訓練自研神經(jīng)網(wǎng)絡(luò)的門檻其實沒有想象中那么高。只要把結(jié)構(gòu)做輕、把顯存賬算清、把訓練配置調(diào)對一張8GB顯卡足夠讓人完成從想法到開源的全過程。希望你的第一把訓練也能在普通顯卡上順利跑起來。