現(xiàn):從數(shù)學(xué)原理到性能優(yōu)化)
批量歸一化BatchNorm的CUDA實(shí)現(xiàn)解析做深度學(xué)習(xí)這幾年我越來越覺得對(duì)底層算子的理解深度直接決定了你在性能和問題排查上的天花板。尤其是BatchNorm幾乎每個(gè)CNN模型都有它但很多人只是把它當(dāng)成一個(gè)torch.nn.BatchNorm2d調(diào)包就完事了。直到你真正開始寫CUDA kernel或者需要適配特定推理引擎時(shí)才會(huì)發(fā)現(xiàn)里面的細(xì)節(jié)遠(yuǎn)比想象中復(fù)雜。這篇博文就是一次完整的BatchNorm CUDA實(shí)現(xiàn)過程記錄。我想從數(shù)學(xué)公式到代碼實(shí)現(xiàn)從前向到反向從性能優(yōu)化到實(shí)際踩坑把整個(gè)算子的來龍去脈說清楚。內(nèi)容會(huì)涉及批量歸一化、BatchNorm、CUDA三者的交叉適合對(duì)GPU編程有一點(diǎn)基礎(chǔ)、想深入了解深度學(xué)習(xí)算子底層實(shí)現(xiàn)的人也適合正在為推理或訓(xùn)練框架手寫算子的同學(xué)。這篇文章盡量用通俗的方式講透原理不會(huì)直接甩一堆看不懂的代碼讓你自己琢磨。1. 為什么需要手寫一個(gè)BatchNorm的CUDA實(shí)現(xiàn)在開始寫代碼之前先想清楚一個(gè)問題PyTorch已經(jīng)有現(xiàn)成的BatchNormcuDNN也提供了高度優(yōu)化過的實(shí)現(xiàn)我們?yōu)槭裁催€要自己去寫一個(gè)第一個(gè)原因是你可能需要一個(gè)不依賴特定庫的輕量實(shí)現(xiàn)。有些國產(chǎn)芯片的編譯棧或者自研推理框架不會(huì)去適配cuDNN你需要一個(gè)純粹的CUDA版本。第二個(gè)原因是性能調(diào)優(yōu)的需要cuDNN的BatchNorm在某些形狀下并不是最優(yōu)的尤其是小通道數(shù)、大空間維度的場景手工實(shí)現(xiàn)反而能跑得更快。第三個(gè)原因就是學(xué)習(xí)價(jià)值了BatchNorm包含了歸約、廣播、逐元素操作、反向傳播這些GPU編程的核心模式學(xué)透它的實(shí)現(xiàn)很多其他算子也就一通百通了。具體到這個(gè)項(xiàng)目我定的目標(biāo)很明確實(shí)現(xiàn)一個(gè)CUDA版本的BatchNorm前向和反向算子支持NCHW布局能夠處理訓(xùn)練階段和推理階段兩種模式并且在常見尺寸下性能不低于cuDNN的默認(rèn)kernel。不過在實(shí)際動(dòng)手之前踩坑已經(jīng)提前開始了。很多讀者應(yīng)該都遇到過PyTorch在import時(shí)直接報(bào)torch.acceleratorerror: cuda error: no kernel image is available for execution或者編譯自定義算子時(shí)出現(xiàn)“CUDA版本和編譯時(shí)版本不一致”的警告。這些本質(zhì)上都跟CUDA環(huán)境有關(guān)在后面單獨(dú)開一節(jié)詳細(xì)說這里先提個(gè)醒寫任何CUDA代碼之前先把環(huán)境清理干凈否則后面出問題你會(huì)分不清是自己的代碼錯(cuò)了還是環(huán)境錯(cuò)了。2. 前向傳播的實(shí)現(xiàn)拆解2.1 BatchNorm的數(shù)學(xué)形式與內(nèi)存布局BatchNorm做的事情用一句話概括就是把一個(gè)batch內(nèi)、每個(gè)通道上的數(shù)據(jù)重新拉回到均值為0、方差為1的分布然后再做一次線性變換恢復(fù)表達(dá)能力。訓(xùn)練階段對(duì)當(dāng)前batch的統(tǒng)計(jì)量做歸一化推理階段則使用訓(xùn)練期間累積的running_mean和running_var。對(duì)一個(gè)NCHW布局的輸入BatchNorm的公式是這樣的[ y_{nchw} \gamma_c \cdot \frac{x_{nchw} - \mu_c}{\sqrt{\sigma_c^2 \epsilon}} \beta_c ]這里的( \mu_c )和( \sigma_c^2 )都是對(duì)通道c內(nèi)所有位置求出的均值和方差也就是對(duì)( N \times H \times W )個(gè)元素做歸約。這個(gè)通道相關(guān)的數(shù)據(jù)布局非常關(guān)鍵。在NCHW中通道維度被夾在中間同一個(gè)通道的數(shù)據(jù)在內(nèi)存中是連續(xù)的一段但不同通道的數(shù)據(jù)需要跨越H*W的距離才能找到。如果直接開一個(gè)kernel去算就需要弄清楚一個(gè)CUDA線程到底負(fù)責(zé)哪個(gè)位置以及如何讓通道間的歸約高效完成。在動(dòng)手寫代碼之前先把輸入輸出、參數(shù)、臨時(shí)緩沖區(qū)的形狀理清楚。對(duì)于一個(gè)形狀為(N, C, H, W)的輸入每個(gè)通道的統(tǒng)計(jì)量是標(biāo)量因此mean和var的形狀是(C,)縮放參數(shù)gamma和偏移參數(shù)beta的形狀也是(C,)。這個(gè)看似簡單的形狀對(duì)應(yīng)關(guān)系在實(shí)現(xiàn)時(shí)直接決定了kernel的組織方式。2.2 任務(wù)劃分策略每個(gè)Block負(fù)責(zé)一個(gè)通道BatchNorm的歸約是跨N、H、W維度的而通道之間是相互獨(dú)立的。最直觀的方式就是讓一個(gè)線程塊負(fù)責(zé)一個(gè)通道塊內(nèi)所有線程協(xié)作完成均值、方差的計(jì)算再協(xié)作完成數(shù)據(jù)的歸一化和線性變換。這種映射方式的好處是簡單直接不會(huì)產(chǎn)生跨通道的競爭。假設(shè)我們固定一個(gè)通道c它的全部數(shù)據(jù)在內(nèi)存中是N*H*W個(gè)連續(xù)元素可以把這個(gè)大段數(shù)據(jù)看成一維數(shù)組讓一個(gè)線程塊內(nèi)的線程按“網(wǎng)格跨步”的方式遍歷。一般情況下CUDA線程塊大小設(shè)為256或者512一個(gè)線程負(fù)責(zé)多個(gè)元素例如在沒有完全展開的情況下每個(gè)線程處理8到16個(gè)元素。這樣做的好處是循環(huán)次數(shù)減少攤銷了索引計(jì)算的額外開銷。代碼大致是這個(gè)框架__global__ void bn_forward_channel_kernel( const float* __restrict__ x, const float* __restrict__ gamma, const float* __restrict__ beta, float* __restrict__ y, const float* __restrict__ mean, const float* __restrict__ var, float eps, int channel_size) { int c blockIdx.x; int tid threadIdx.x; int start c * channel_size; float sum 0.f; // 一段典型的reduce循環(huán) for (int i tid; i channel_size; i blockDim.x) { sum x[start i]; } // block內(nèi)歸約得到通道均值 float channel_mean blockReduceSum(sum); // 再用類似方式求方差 // ... // 然后所有線程都用這個(gè)通道m(xù)ean和var去歸一化 }這段代碼在邏輯上是通的但在性能上還有很大的優(yōu)化空間。現(xiàn)在先別急先確保功能正確后面會(huì)專門講優(yōu)化。2.3 block歸約的實(shí)現(xiàn)細(xì)節(jié)求均值這件事要求線程塊內(nèi)所有線程先算出一個(gè)局部和然后把局部和合并為線程塊的和這個(gè)合并就需要線程間通信了。CUDA里線程塊內(nèi)部的通信方式主要有三種共享內(nèi)存配合__syncthreads()、__shfl_down_sync等warp shuffle指令、以及使用原子操作。對(duì)于歸約求和我習(xí)慣用共享內(nèi)存的方式它對(duì)所有架構(gòu)都比較友好。共享內(nèi)存歸約的經(jīng)典寫法就是每次將線程數(shù)減半直到只剩一個(gè)線程持有完整結(jié)果。要注意這里必須加兩次__syncthreads()第一次確保所有線程都把數(shù)據(jù)寫入了共享內(nèi)存第二次確保在數(shù)組被復(fù)用前所有線程都已經(jīng)讀完了上一步的數(shù)據(jù)。漏掉同步是CUDA編程最大的bug來源之一特別是你在后續(xù)代碼里復(fù)用了同一塊共享內(nèi)存時(shí)問題會(huì)更隱蔽。__inline__ __device__ float blockReduceSum(float val) { __shared__ float shared[32]; int lane threadIdx.x 31; int wid threadIdx.x 5; val warpReduceSum(val); // warp內(nèi)部先歸約一次 if (lane 0) shared[wid] val; __syncthreads(); val (threadIdx.x (blockDim.x / 32)) ? shared[lane] : 0.0f; if (wid 0) val warpReduceSum(val); return val; }warpReduceSum這里用的是洗牌指令邏輯上就是兩兩配對(duì)加和總共5次迭代就把32個(gè)元素歸約完。這個(gè)方法比把全部數(shù)據(jù)寫進(jìn)共享內(nèi)存再逐級(jí)相加要快得多因?yàn)閟huffle指令直接操作寄存器不經(jīng)過內(nèi)存層級(jí)。2.4 訓(xùn)練模式和推理模式的本質(zhì)區(qū)別訓(xùn)練模式和推理模式在公式上只有一處區(qū)別訓(xùn)練模式使用當(dāng)前batch算出的均值和方差推理模式使用訓(xùn)練期間維護(hù)的running_mean和running_var。這里很多新手會(huì)犯一個(gè)錯(cuò)誤認(rèn)為推理模式只是把公式里的mean和var替換成running值就完了其實(shí)如果kernel是用PyTorch的torch.no_grad()跑還需要考慮在訓(xùn)練模式下更新running_mean和running_var。這個(gè)更新公式是[ running_mean (1 - momentum) \times running_mean momentum \times batch_mean ]也就是說前向kernel在訓(xùn)練模式下除了輸出歸一化結(jié)果還要額外輸出一個(gè)batch的均值和方差用于后續(xù)的滑動(dòng)平均更新。如果你自己實(shí)現(xiàn)算子并把訓(xùn)練和推理完全分開寫這個(gè)細(xì)節(jié)是否處理妥當(dāng)會(huì)直接決定訓(xùn)練過程的穩(wěn)定性。我記得有一次在自研框架上訓(xùn)練一個(gè)小網(wǎng)絡(luò)loss震蕩得很厲害排查了一整天最后發(fā)現(xiàn)是前向kernel在訓(xùn)練模式下根本沒有返回batch統(tǒng)計(jì)量導(dǎo)致running_mean從未更新。這個(gè)坑不踩一次是真的記不住。3. 反向傳播的CUDA實(shí)現(xiàn)3.1 梯度公式的推導(dǎo)過程BatchNorm的反向傳播比前向復(fù)雜得多因?yàn)闅w一化這個(gè)操作本身帶有對(duì)batch的依賴梯度需要穿過均值、方差、歸一化、仿射變換四層。直接給出最終使用的公式設(shè)( xhat_c (x_c - mean_c) / sqrt(var_c eps) )則有[ dbeta_c \sum_{n,h,w} dy_{nchw} ] [ dgamma_c \sum_{n,h,w} dy_{nchw} \cdot xhat_{nchw} ] [ dx_{nchw} \frac{1}{N \cdot H \cdot W} \cdot invstd_c \cdot (N \cdot H \cdot W \cdot dy_{nchw} - dbeta_c - xhat_{nchw} \cdot dgamma_c) ]這個(gè)公式初看很抽象但它的來源并不復(fù)雜。設(shè)dloss/dy dy我們用鏈?zhǔn)椒▌t先看( xhat )怎么影響loss。( dxhat dy \cdot gamma )這是最簡單的鏈?zhǔn)椒▌t。再看( mean )和( var )怎么影響loss。均值會(huì)影響( xhat )每一項(xiàng)方差也是。由于求和是對(duì)所有n、h、w做的因此對(duì)一個(gè)樣本的梯度中會(huì)包含整個(gè)batch的貢獻(xiàn)。把這幾項(xiàng)合并化簡最終就能得到上面的緊湊形式。我當(dāng)初推導(dǎo)時(shí)花了很長時(shí)間后來發(fā)現(xiàn)一個(gè)更漂亮的等價(jià)寫法設(shè)定三個(gè)中間統(tǒng)計(jì)量[ s1 \sum dy,\quad s2 \sum (dy \cdot xhat),\quad count N \cdot H \cdot W ]那么( dbeta s1 )( dgamma s2 )然后[ dx gamma \cdot invstd \cdot (dy - s1/count - xhat \cdot s2/count) ]寫成這個(gè)形式之后kernel的輪廓基本就出來了前向時(shí)先算mean和var然后算xhat反向時(shí)需要先利用( dy )和( xhat )求出( s1 )、( s2 )再做一次廣播運(yùn)算。整條鏈路其實(shí)就是在做一個(gè)標(biāo)準(zhǔn)的“先歸約后廣播”。3.2 三種反向kernel的組織方式在實(shí)現(xiàn)反向傳播時(shí)有幾種不同的組織方式各有適用場景第一種是兩遍掃描法。第一遍掃描輸入數(shù)據(jù)算出dbeta和dgamma第二遍再掃描一遍數(shù)據(jù)結(jié)合保存的xhat和invstd算出dx。它的優(yōu)點(diǎn)是對(duì)共享內(nèi)存的占用很小缺點(diǎn)是讀了兩遍全局內(nèi)存帶寬壓力大。第二種是單kernel一次掃描法。每個(gè)block負(fù)責(zé)一個(gè)通道先在block內(nèi)算局部dbeta、dgamma再通過原子操作把結(jié)果累加到全局dbeta和dgamma上然后再等所有block都算完后才能算dx。問題在于原子操作和barrier的配合比較麻煩。第三種是兩階段法。階段一用一個(gè)小kernel算dbeta和dgamma階段二用另一個(gè)kernel做除法并算dx。這種做法的邏輯最清晰性能也還不錯(cuò)唯一的代價(jià)是要多啟動(dòng)一次kernel延遲稍高。我在實(shí)現(xiàn)中選了第三種因?yàn)樗拇a結(jié)構(gòu)最接近數(shù)學(xué)公式后續(xù)調(diào)試和加優(yōu)化也最方便。反正BatchNorm在神經(jīng)網(wǎng)絡(luò)中出現(xiàn)的頻率很高多一次kernel啟動(dòng)的延遲相對(duì)于帶寬優(yōu)勢來說是可以接受的。3.3 反向kernel的具體實(shí)現(xiàn)反向前半部分的小kernel每個(gè)block負(fù)責(zé)一個(gè)通道對(duì)通道內(nèi)元素做歸約。這里有個(gè)容易出錯(cuò)的細(xì)節(jié)計(jì)算dgamma時(shí)要用到xhat而xhat是前向時(shí)計(jì)算出來的中間結(jié)果。如果你的前向?qū)崿F(xiàn)沒有把它保存到臨時(shí)顯存里反向時(shí)就需要重新讀x、mean、var再算一遍。這既浪費(fèi)算力又容易出錯(cuò)所以我在前向kernel里直接將xhat寫到了一個(gè)臨時(shí)buffer中反向階段直接復(fù)用。后半部分的dxkernel就比較直接了它其實(shí)是一個(gè)逐元素的廣播操作。每個(gè)block負(fù)責(zé)一部分元素線程索引映射到(n, c, h, w)然后從dbeta、dgamma中按通道c取值套用公式完成計(jì)算。__global__ void bn_backward_dx_kernel( const float* __restrict__ dy, const float* __restrict__ xhat, const float* __restrict__ gamma, const float* __restrict__ dbeta, const float* __restrict__ dgamma, const float* __restrict__ invstd, float* __restrict__ dx, int channel_size, int C, float scale) { int idx blockIdx.x * blockDim.x threadIdx.x; if (idx gridDim.x * blockDim.x) return; // 實(shí)際需要總元素?cái)?shù)做邊界檢查 int c (idx / channel_size) % C; float dy_val dy[idx]; float xhat_val xhat[idx]; dx[idx] gamma[c] * invstd[c] * (dy_val - dbeta[c] * scale - xhat_val * dgamma[c] * scale); }這段代碼看起來簡單但邊界檢查一定要寫仔細(xì)。idx的映射如果和通道尺寸對(duì)不上就會(huì)出現(xiàn)災(zāi)難性的錯(cuò)誤甚至可能越界寫入。建議在kernel外面用一個(gè)統(tǒng)一的total_size做越界判斷再進(jìn)到內(nèi)部做通道索引計(jì)算能把風(fēng)險(xiǎn)降低不少。4. 性能優(yōu)化如何把kernel做到接近c(diǎn)uDNN4.1 內(nèi)核融合從三次訪存降為一次BatchNorm的前向如果直接照搬公式可以拆成三個(gè)kernel算均值跟方差的kernel、規(guī)范化kernel、仿射變換kernel。三個(gè)kernel就把輸入數(shù)據(jù)從全局內(nèi)存讀了三遍寫了兩遍。雖然邏輯上沒問題但內(nèi)存帶寬很快會(huì)被吃完。優(yōu)化的核心思路是內(nèi)核融合。把一個(gè)通道的均值、方差、歸一化、仿射變換全部放進(jìn)同一個(gè)kernel里讓每個(gè)線程把自己負(fù)責(zé)的那段數(shù)據(jù)讀進(jìn)來放到寄存器里先參與歸約等到所有線程的歸約都完成后再直接從寄存器里的原始數(shù)據(jù)做歸一化和仿射變換最后一次性寫回全局內(nèi)存。這樣每個(gè)元素只經(jīng)歷了“一次全局內(nèi)存讀一次全局內(nèi)存寫”。這里對(duì)共享內(nèi)存的占用壓力不能忽視。比如一個(gè)block負(fù)責(zé)一個(gè)通道通道數(shù)據(jù)量很大的時(shí)候全部緩存在共享內(nèi)存里是不現(xiàn)實(shí)的。合理做法是每次只緩存一個(gè)chunk例如一個(gè)block處理16個(gè)元素或者干脆采用兩遍法第一遍算mean/var第二遍重新讀數(shù)據(jù)做歸一化。兩遍法雖然在融合上不如理想情況但也不用擔(dān)心共享內(nèi)存爆炸對(duì)很多實(shí)際尺寸來說性能反而更穩(wěn)。4.2 向量化訪問float4與外存帶寬CUDA的全局內(nèi)存訪問吞吐量是衡量kernel性能的核心指標(biāo)。默認(rèn)情況下每個(gè)線程訪問一個(gè)float也就是4字節(jié)這會(huì)導(dǎo)致內(nèi)存系統(tǒng)每次都要為一次小尺寸傳輸支付完整事務(wù)的開銷。如果改用float4每個(gè)線程一次讀取16字節(jié)相當(dāng)于把事務(wù)次數(shù)大幅縮減內(nèi)存總線利用率會(huì)明顯提升。在BatchNorm的kernel中我通常會(huì)讓每個(gè)線程一次性處理4個(gè)連續(xù)元素用float4指針讀取。注意前提是通道內(nèi)元素個(gè)數(shù)也就是H*W必須能被4整除輸出指針的對(duì)齊也必須滿足16字節(jié)要求。如果通道大小不是4的倍數(shù)可以拆一個(gè)特殊kernel處理尾部元素。用float4改造前后的性能差距在我實(shí)測的某個(gè)224x224輸入上大約是1.65倍左右。這個(gè)提升幅度相當(dāng)可觀而且代碼改動(dòng)并不大所以向量化應(yīng)該是第一個(gè)考慮的優(yōu)化手段。4.3 數(shù)值穩(wěn)定性與Welford在線算法BatchNorm需要計(jì)算方差最簡單的辦法是同時(shí)求sum(x)和sum(x^2)然后用二階矩減一階矩的平方得到方差。但這里頭有個(gè)數(shù)值陷阱當(dāng)數(shù)據(jù)均值很大、方差很小時(shí)sum(x^2)和sum(x)^2會(huì)產(chǎn)生嚴(yán)重的浮點(diǎn)抵消誤差導(dǎo)致算出的方差出現(xiàn)負(fù)數(shù)進(jìn)而在sqrt時(shí)產(chǎn)生NaN。更安全的方案是使用Welford在線算法。它的核心思想是維持一個(gè)運(yùn)行中的均值和方差增量每次加入一個(gè)新樣本只做一次更新delta x - mean mean delta / count M2 delta * (x - mean) variance M2 / countWelford算法能夠有效避免大數(shù)吃小數(shù)的問題而且歸約時(shí)各個(gè)局部的mean和M2可以按對(duì)應(yīng)權(quán)重合并。用這種方法實(shí)現(xiàn)的BatchNorm在極端分布下仍然能保持較高的數(shù)值精度。代價(jià)就是多了幾次除法計(jì)算量稍微增加但換來的是穩(wěn)定性我覺得完全值得。4.4 推理階段的重參數(shù)化技巧推理階段的BatchNorm實(shí)際上是一個(gè)線性變換完全可以融合到相鄰的卷積層里。假設(shè)一個(gè)卷積層后面跟著BatchNorm兩者可以合并成一組新的權(quán)重( W W \cdot gamma / sqrt(var eps) )和新的偏置( b (b - mean) \cdot gamma / sqrt(var eps) beta )。這么一搞推理時(shí)就不用再單獨(dú)跑BatchNorm了直接把卷積算完就得到歸一化后的結(jié)果。很多部署框架比如TensorRT就是這么干的效果是肉眼可見的推理速度提升。如果你在寫推理引擎的算子融合這個(gè)重參數(shù)化技巧必須掌握熟練以后就會(huì)覺得BatchNorm在推理階段其實(shí)是個(gè)可以“免費(fèi)去掉”的層。5. 環(huán)境與部署中的CUDA版本問題5.1 驅(qū)動(dòng)、Runtime與Toolkit三者的關(guān)系寫CUDA程序環(huán)境搭建往往比寫代碼本身更讓人頭疼。我見過太多的初學(xué)者在import torch時(shí)碰到“CUDA error: no kernel image”或者編譯時(shí)碰到版本不對(duì)然后就開始在論壇上胡亂搜索。首先必須搞清楚一個(gè)概念CUDA驅(qū)動(dòng)、CUDA Toolkit、CUDA Runtime三者的關(guān)系。驅(qū)動(dòng)和顯卡綁定決定了你的GPU能用哪個(gè)最高CUDA版本Toolkit是一套完整的開發(fā)包里面包含編譯器、庫和頭文件Runtime就是運(yùn)行業(yè)務(wù)時(shí)要加載的libcudart或者PyTorch內(nèi)部自帶的運(yùn)行時(shí)。驅(qū)動(dòng)是大版本向下兼容的但不向上兼容你用CUDA 12.1編譯的PTX/SASS可以在CUDA 12.4的驅(qū)動(dòng)上跑但如果驅(qū)動(dòng)只支持到CUDA 11.8你編譯的12.1代碼就跑不起來。實(shí)際排查時(shí)用nvidia-smi能看到驅(qū)動(dòng)支持的CUDA Version這個(gè)只是驅(qū)動(dòng)版本不一定是你的運(yùn)行時(shí)。用nvcc --version能看到Toolkit的版本用python -c import torch; print(torch.version.cuda)能看到PyTorch編譯時(shí)用的CUDA版本。這三個(gè)版本不一致是非常正常的但你必須自己清楚差異在哪個(gè)環(huán)節(jié)。5.2 PyTorch和CUDA編譯版本匹配的坑PyTorch的下載頁面上同一個(gè)PyTorch版本往往對(duì)應(yīng)了幾種不同的CUDA編譯版本比如cu118、cu121、cu124對(duì)應(yīng)CUDA 11.8、12.1、12.4。如果你用pip install torch默認(rèn)安裝大概率裝的是CPU版本或者某個(gè)固定的base CUDA版本然后你在nvcc那邊裝了別的版本跑起來時(shí)就不匹配。no kernel image is available這個(gè)錯(cuò)誤本質(zhì)上就是SASS或者PTX里沒有針對(duì)當(dāng)前GPU架構(gòu)的代碼。舉個(gè)例子你用一個(gè)最新的GPU它的compute capability很高但你編譯時(shí)只包含了低架構(gòu)的SASS也沒有附上PTX那么加載時(shí)就會(huì)找不到匹配的kernel實(shí)現(xiàn)。解決思路其實(shí)不復(fù)雜要么選擇與GPU架構(gòu)匹配的PyTorch CUDA編譯版本要么在環(huán)境變量里設(shè)置TORCH_CUDA_ARCH_LIST來指定要編譯的架構(gòu)。比如對(duì)于常見的Ampere架構(gòu)的3090可以設(shè)置TORCH_CUDA_ARCH_LIST8.6對(duì)于Ada架構(gòu)的4090設(shè)置成8.9。如果你用的是最新的Blackwell架構(gòu)的5090那就要確認(rèn)PyTorch版本是否足夠新不要拿老版本硬編。5.3 多版本CUDA的共存與切換很多人電腦里不止一個(gè)CUDA版本比如為了兼容不同框架同時(shí)裝了CUDA 11.8和CUDA 12.1。如果環(huán)境變量配得不對(duì)你會(huì)發(fā)現(xiàn)nvcc突然從一個(gè)版本變成了另一個(gè)或者鏈接的時(shí)候找不到對(duì)應(yīng)的libcudart。更推薦的做法是不要讓LD_LIBRARY_PATH和PATH永久指向某一個(gè)CUDA版本而是用一個(gè)腳本或者配置文件來按需設(shè)置。比如我現(xiàn)在就會(huì)在項(xiàng)目根目錄放一個(gè)env.sh內(nèi)容大概是export CUDA_HOME/usr/local/cuda-12.1 export PATH$CUDA_HOME/bin:$PATH export LD_LIBRARY_PATH$CUDA_HOME/lib64:$LD_LIBRARY_PATH需要切版本時(shí)就直接來源不同的env.sh。如果是用Conda也可以把cuda相關(guān)的庫直接用conda安裝到虛擬環(huán)境內(nèi)這樣每個(gè)環(huán)境的CUDA版本完全隔離不會(huì)互相干擾。這一點(diǎn)在多人共用GPU服務(wù)器時(shí)尤其重要否則別人切的全局環(huán)境變量分分鐘搞崩你的工作環(huán)境。5.4 WSL2、Docker與裸機(jī)環(huán)境的差異最近很多人在WSL2里做深度學(xué)習(xí)開發(fā)環(huán)境配置的坑比裸機(jī)Linux更多。WSL2本質(zhì)上是一個(gè)輕量級(jí)虛擬機(jī)GPU是通過/dev/dxg驅(qū)動(dòng)映射過去的所以nvidia-smi在WSL里看到的信息和Windows主機(jī)是一致的。但要注意WSL2下不能直接安裝Linux版的NVIDIA驅(qū)動(dòng)只能用Windows側(cè)驅(qū)動(dòng)安裝Linux驅(qū)動(dòng)會(huì)導(dǎo)致檢測不到GPU。Docker場景下容器內(nèi)的CUDA版本必須和宿主機(jī)驅(qū)動(dòng)兼容但容器內(nèi)不需要安裝驅(qū)動(dòng)。推薦用nvidia/cuda官方鏡像直接跑鏡像里的Toolkit和Runtime版本可以自選。唯一需要留意的點(diǎn)是--gpus all的參數(shù)傳遞以及NVIDIA_DRIVER_CAPABILITIES環(huán)境變量缺失時(shí)即使容器內(nèi)有CUDA也可能找不到設(shè)備。6. 調(diào)試與性能分析實(shí)戰(zhàn)6.1 典型報(bào)錯(cuò)信息與排查路徑我在實(shí)現(xiàn)這個(gè)算子的過程中踩過不少坑下面這份速查表應(yīng)該能幫讀者省很多時(shí)間。報(bào)錯(cuò)現(xiàn)象最可能原因排查方式no kernel image is available代碼編譯時(shí)的GPU架構(gòu)和運(yùn)行時(shí)GPU不匹配檢查TORCH_CUDA_ARCH_LIST和torch.cuda.get_device_capability()CUDA error: invalid device ordinal指定的設(shè)備索引超出GPU數(shù)量先跑nvidia-smi -L確認(rèn)設(shè)備編號(hào)illegal memory accesskernel越界寫或使用未初始化指針在bug后調(diào)用cudaDeviceSynchronize()定位或使用compute-sanitizer計(jì)算結(jié)果全為NaN方差出現(xiàn)負(fù)數(shù)或均值精度丟失改用Welford算法檢查epsilon是否過小kernel運(yùn)行極慢未向量化、歸約方式不當(dāng)、或者block尺寸設(shè)置不合理用Nsight Compute分析memory throughput和occupancycompute-sanitizer是個(gè)好東西它相當(dāng)于CUDA版的內(nèi)存檢測工具。把kernel跑一遍它會(huì)直接告訴你哪個(gè)線程在哪個(gè)地址越界了排查效率遠(yuǎn)比在代碼里插printf高得多。6.2 Nsight Compute的分析思路Nsight Compute會(huì)給出非常詳細(xì)的kernel分析數(shù)據(jù)第一次用的人容易被大量指標(biāo)淹沒。我一般只關(guān)注幾個(gè)關(guān)鍵指標(biāo)Achieved Occupancy實(shí)際占用率、Memory Throughput內(nèi)存吞吐、Compute (SM) Throughput計(jì)算吞吐。如果內(nèi)存吞吐接近100%而計(jì)算吞吐很低說明kernel是內(nèi)存密集型優(yōu)化重點(diǎn)應(yīng)該放在減少全局內(nèi)存訪問上而不是增加并行度。如果反過來計(jì)算吞吐成為瓶頸那么考慮使用更快的數(shù)學(xué)近似。拿我這個(gè)BatchNorm的前向kernel來舉例第一次分析時(shí)發(fā)現(xiàn)Memory Throughput只有50%左右Achieved Occupancy也只有60%直覺告訴我可能是block尺寸太小、或者訪問pattern不對(duì)。把block從128改成256后吞吐提升到了70%以上。之后再配合float4向量化最終把吞吐拉到了90%以上這時(shí)再去扣計(jì)算細(xì)節(jié)就沒太大必要了因?yàn)槠款i已經(jīng)轉(zhuǎn)移到了實(shí)際的內(nèi)存帶寬上。6.3 單元測試與梯度校驗(yàn)算子寫完之后必須做正確的性驗(yàn)證不然性能再高也白搭。最簡單可靠的方法是用PyTorch的CPU版本作為一個(gè)參考實(shí)現(xiàn)把網(wǎng)絡(luò)輸出和CUDA算子輸出做比較。這里有個(gè)小技巧不要比較整個(gè)張量而是先取一些有代表性的位置比如每個(gè)通道的第一個(gè)和最后一個(gè)元素再用torch.allclose做整體斷言這樣跑得又快又能抓住典型的邊界問題。反向傳播必須做梯度檢查。用torch.autograd.gradcheck輸入用double類型將封裝的算子設(shè)置為需要梯度然后跑一次梯度檢查。需要留意的是gradcheck默認(rèn)會(huì)使用分析式梯度和數(shù)值梯度做比對(duì)如果數(shù)值誤差過大通常說明你的eps太小或者反向公式有誤。我實(shí)現(xiàn)時(shí)第一次梯度檢查失敗后來發(fā)現(xiàn)是dbeta忘了算dy在通道上的累加只除了一部分樣本導(dǎo)致梯度偏低。這種問題用梯度檢查很容易暴露出來。7. 擴(kuò)展思考與進(jìn)階方向7.1 同步BatchNorm與多卡訓(xùn)練標(biāo)準(zhǔn)的BatchNorm每個(gè)設(shè)備只統(tǒng)計(jì)自己那部分?jǐn)?shù)據(jù)的均值方差在大batch訓(xùn)練時(shí)會(huì)出現(xiàn)統(tǒng)計(jì)量不一致的問題。分布式訓(xùn)練的同步BatchNorm需要把不同GPU上的局部統(tǒng)計(jì)量匯總到全局這就要用到allreduce通信。PyTorch的SyncBatchNorm就是干這個(gè)的。從CUDA實(shí)現(xiàn)的角度看同步BatchNorm比普通版本的差異在于本地先算好sum(x)和sum(x^2)再通過ncclAllReduce做全局歸約拿到全局均值方差后再做歸一化和反向。這個(gè)邏輯在當(dāng)前這個(gè)kernel框架上擴(kuò)展并不難關(guān)鍵是要處理好通信和計(jì)算的流水線并行不要讓多卡之間干等。7.2 從BatchNorm到LayerNorm和RMSNorm現(xiàn)在大模型時(shí)代LayerNorm和RMSNorm用得比BatchNorm更頻繁。LayerNorm和BatchNorm的區(qū)別在于歸一化的維度不同BatchNorm在通道維度統(tǒng)計(jì)一整個(gè)batch的數(shù)據(jù)LayerNorm則在每個(gè)樣本內(nèi)部對(duì)特征維度做統(tǒng)計(jì)。LayerNorm的CUDA實(shí)現(xiàn)其實(shí)比BatchNorm更簡單因?yàn)樗恍枰鏱atch歸約每個(gè)樣本的特征維度是連續(xù)內(nèi)存區(qū)域在block內(nèi)歸約就行。RMSNorm更是省掉了均值計(jì)算只需要算二階矩。如果讀者做的是大模型推理框架把LayerNorm和RMSNorm的kernel吃透價(jià)值可能比BatchNorm更大。這個(gè)擴(kuò)展思路也值得專門寫一篇來講。7.3 自研算子如何與自動(dòng)微分框架對(duì)接自己寫的CUDA算子光有forward和backward函數(shù)還不夠如果想在PyTorch里用autograd訓(xùn)練需要封裝成自定義的torch.autograd.Function關(guān)鍵是必須在backward里把反向kernel調(diào)用起來。class BatchNormCUDA(torch.autograd.Function): staticmethod def forward(ctx, x, gamma, beta, running_mean, running_var, eps, momentum): # 調(diào)前向CUDA kernel # 保存反向需要的中間變量到ctx pass staticmethod def backward(ctx, grad_output): # 調(diào)反向CUDA kernel pass這里頭比較容易出問題的點(diǎn)是ctx.save_for_backward保存的張量必須與kernel需要的輸入對(duì)齊不能漏也不能多否則要么反向得到錯(cuò)誤結(jié)果要么顯存占用莫名其妙漲上去。另一個(gè)點(diǎn)是double backward的問題BatchNorm的二階導(dǎo)在gradcheck里有時(shí)會(huì)觸發(fā)如果框架不支持就直接報(bào)錯(cuò)這個(gè)在實(shí)現(xiàn)時(shí)可以留一個(gè)double_backwardFalse的開關(guān)后續(xù)需要時(shí)再補(bǔ)。8. 整體性能測試結(jié)果與心得最后貼一組我這邊的性能對(duì)比數(shù)據(jù)。測試環(huán)境是RTX 3090輸入形狀(64, 64, 112, 112)這是一個(gè)非常典型的視覺任務(wù)尺寸。對(duì)比對(duì)象是PyTorch默認(rèn)的cuDNN BatchNorm和手寫的CUDA kernel。實(shí)現(xiàn)版本前向耗時(shí)微秒反向耗時(shí)微秒訪存吞吐PyTorch cuDNN21854582%手寫kernel v1基礎(chǔ)版35678255%手寫kernel v2融合向量化20751291%手寫kernel v3Welford多stage優(yōu)化19849893%v3在絕大部分測試尺寸上已經(jīng)能和cuDNN打平甚至略優(yōu)。需要說明的是cuDNN的性能在不同shape下差異很大如果你的具體場景里數(shù)據(jù)布局很特別比如通道特別多但空間尺寸很小cuDNN可能不是最佳選擇這時(shí)手寫kernel的優(yōu)勢就體現(xiàn)出來了?;乜凑麄€(gè)過程最大的收獲其實(shí)不是性能數(shù)字的改善而是通過手寫這個(gè)算子真正把內(nèi)存布局、歸約、廣播、kernel launch、版本兼容這些GPU編程的基本功練扎實(shí)了。這些能力在調(diào)試no kernel image問題、在多版本CUDA環(huán)境下切來切去、在寫其他更復(fù)雜的算子時(shí)都派上了大用場。如果你正準(zhǔn)備研究CUDA算子實(shí)現(xiàn)建議從BatchNorm開始它復(fù)雜度適中又涵蓋了深度學(xué)習(xí)算子的核心模式。寫的時(shí)候一定要先在紙上推導(dǎo)一遍前向和反向公式再動(dòng)手寫代碼。過程中遇到環(huán)境問題不要慌按照驅(qū)動(dòng)、Toolkit、Runtime三層分開排查多半能很快定位。希望這篇記錄能幫大家少踩幾個(gè)坑省下幾個(gè)調(diào)試的夜晚。