化:從DeepGEMM看手寫矩陣乘法的關鍵設計)
1. 為什么深度學習里總有人在死磕GEMM前陣子在技術社區(qū)看到一個項目名叫DeepGEMM瞬間就來了興趣。做深度學習系統(tǒng)優(yōu)化的人應該都有同感矩陣乘法這個東西幾乎是所有計算場景的地基卷積要轉化成隱式的GEMM來跑Transformer里的自注意力本質是三個矩陣乘MLP更不用說了前向和反向都是標準GEMM操作??梢赃@么說只要模型訓練和推理還在用GPUGEMM性能就是繞不開的核心指標。很多人會問GPU廠商不是已經提供了cuBLAS這類高度優(yōu)化的官方庫嗎直接用不就行了為什么還要自己從頭寫一個DeepGEMM這個問題我當年也疑惑過。直到我在真實項目里跑過一輪對比才知道官方庫只保證通用場景還不錯它不是為你的特定硬件、特定算法和特定shape定制的。模型訓練過程中矩陣的批量大小、通道數、序列長度經常是固定的幾組shape官方庫在這幾組shape上未必是最優(yōu)的而且推框架、推算子融合、推低精度訓練的時候很多時候你需要把矩陣乘法的kernel邏輯直接嵌到更大規(guī)模的融合算子里面去這已經超出庫函數的邊界了。所以DeepGEMM這個名字本身表達的就是這么一件事把GEMM做到深度學習場景下的極致。它不是一個簡單的封裝而是從寄存器分配、共享內存使用、數據排布、指令調度到精度策略全面重新設計的矩陣乘內核實現(xiàn)。評論區(qū)有個老哥總結得很到位——通用庫解決的是多快好省地做完所有事DeepGEMM解決的是在你知道自己的工作負載是什么的前提下把每一絲硬件性能都榨出來。下面我就從這次看到這個項目后順手復現(xiàn)和拆解它的思路出發(fā)把關鍵點和可以遷移的經驗寫清楚。2. DeepGEMM的核心目標與使用場景拆解2.1 它到底解決哪一類問題先明確一下邊界條件。DeepGEMM不是把所有矩陣乘都做到極致的銀彈它聚焦的是深度學習中最常見的那一類GEMM半精度計算場景下的GEMM也就是Gemm with FP16/BF16輸入累加用FP32來保證精度同時配合常見的layout策略來避免額外的格式轉換開銷。我平時在調LLM推理服務的時候大量的耗時其實都集中在GEMM上。如果只看運營商提供的性能數據滿峰值很漂亮但你真正把模型跑起來之后實際的設備利用率經常只有30%-50%。刨除掉小而頻繁的矩陣乘、內存帶寬瓶頸之外一大塊開銷就來自庫函數沒有針對你的固定shape做最優(yōu)的tiling。DeepGEMM這類項目典型的切入方式就是針對固定shape和硬件特性做極致的tile切分和流水線調度把這個利用率提上來。具體來說DeepGEMM主要面向的場景有這三類固定shape的模型訓練和推理比如批量大小固定為32、序列長度固定為2048的Transformer訓練改造后不用每次都做shape感知的啟發(fā)式搜參。需要算子融合的定制推理引擎GEMM要跟后面的bias加、激活函數、量化反量化融合到一個kernel里庫函數做不到這種程度的定制。低精度推理/訓練場景FP16/BF16這類格式下能否充分利用Tensor Core的吞吐特性直接決定性能上限。2.2 和普通GEMM實現(xiàn)的本質區(qū)別普通的GEMM實現(xiàn)思路通常是這樣的寫好一個baseline計算每個線程負責輸出矩陣的哪些元素然后一步步做向量化、調整循環(huán)順序、加shared memory blocking。大部分優(yōu)化教程到這里就結束了因為再往后需要對硬件的細節(jié)非常敏感。DeepGEMM這類項目的境界不一樣。它在設計層面就把深度學習場景的特征考慮進去了比如矩陣的K維壓縮維往往很長可以切得很細來做流水線線程塊的流水深度可以做得很深半精度輸入情況下Tensor Core以NVIDIA的硬件為參考的指令形態(tài)、warp級的矩陣分片方式都需要分別考慮尺寸怎么配、寄存器怎么預留都有講究數據的N、C維固定情況下可以用更大的tile來攤薄調度和尋址開銷。換句話說普通實現(xiàn)是在怎么用CUDA寫出一個對的GEMMDeepGEMM是在怎么用硬件手冊級的理解寫出一個深度適配工作負載的GEMM。這一點直接決定了代碼中各種看起來反直覺的寫法的來源。3. 從零手寫一個高性能GEMM內核的關鍵設計點先把結論放出來高性能GEMM內核的優(yōu)化永遠圍繞四個字——復用與流水。前者決定你能不能在正確的層次上把數據反復利用后者決定芯片上的各種執(zhí)行單元能不能一直有事做。下面按我自己復現(xiàn)DeepGEMM思路的過程把幾個核心設計點拆開講。3.1 分塊策略為什么Tile尺寸不能拍腦袋定GEMM的每一個輸出元素都需要訪問A矩陣一整行和B矩陣一整列。如果直接在全局內存上做每個元素都要重新讀取數據內存吞吐很快就成了瓶頸。所以高性能實現(xiàn)的第一件事就是分塊把一個大的輸出矩陣切成很多個小塊每個線程塊負責一個輸出塊這個輸出塊的計算所需要用到的A和B數據塊先取到共享內存里供塊內多次循環(huán)復用。分塊尺寸的選擇非常關鍵。取太小共享內存里放的數據不足以支撐足夠的計算復用帶寬還是瓶頸取太大一個線程塊占用的資源過多SM上能同時駐留的塊數量下降調度延遲沒法隱藏。這次拆解DeepGEMM時我關注到的核心tile配置是128×128左右的輸出塊尺寸每個線程塊內部再按warp粒度切分成更小的塊來映射到具體線程。為什么是128以Tensor Core指令為例常見的wmma指令形態(tài)是16×16×16m16n8k8這類形態(tài)也很常見。輸出塊選128意味著在m方向和n方向各能整除16不會出現(xiàn)浪費的邊界處理邏輯同時128這個尺寸在寄存器層面也能比較自然地容納不會因為分片過多導致線程間通信開銷暴增。理論再好最后落實到具體硬件上還是需要對著profiler一個尺寸一個尺寸試的。3.2 數據搬運的隱藏技巧用異步拷貝把內存延遲藏到計算背后傳統(tǒng)做法是先把數據從全局內存拷到共享內存然后線程再讀共享內存去計算。這個過程中拷貝和計算是串行的計算單元在此期間只能等待。改進方向是讓數據搬運在后臺進行當前一個tile在計算時后一個tile的數據已經通過異步拷貝指令在路上了。這個思路在CUDA里最典型的實現(xiàn)方式就是雙緩沖double buffering。共享內存開兩份緩沖區(qū)一份給當前tile做計算一份接收下一個tile的數據交替使用。DeepGEMM在相關內容里設計流水線時用的正是這個思路。代碼結構大概可以這樣理解// 偽代碼示意實際實現(xiàn)需要處理邊界和同步 for (int k 0; k K_tiles; k) { if (k K_tiles - 1) { cp_async(buffers[next], A_block_next, B_block_next); } compute_from(buffers[current], accum); current ^ 1; next ^ 1; commit_group(); wait_group(); }按我實測的經驗光是一個雙緩沖改造就有機會帶來20%-50%的性能提升具體取決于原來kernel的瓶頸類型。如果你的kernel本身就卡在計算密集度不夠、算術單元沒有吃滿那么數據搬運的延遲隱藏會立竿見影如果已經吃滿了這個優(yōu)化收益就會弱一些。3.3 寄存器級復用一次加載多次計算共享內存帶寬同樣不是無限的。如果每次計算都要從共享內存讀操作數共享內存的帶寬也會成為瓶頸。進一步的做法是把數據放進寄存器每個線程負責的小塊輸出對應的A片段和B片段盡量一次性加載到寄存器里后續(xù)循環(huán)中反復使用。Tensor Core上的warp級矩陣指令有個很關鍵的特性它要求每個線程持有A、B矩陣的特定分片且這些分片要按固定的布局存在于寄存器中。所以寫的代碼要嚴格按照指令要求去鋪數據多一個轉置或shuffle都會帶來額外開銷。這里面比較考驗功力的是片段切分每個線程算輸出T一個2×1的小塊到底A矩陣給這個線程分配哪些元素、B矩陣分配哪些需要根據指令形狀和線程位次仔細推導。3.4 累加器布局與精度補償半精度算GEMM最怕的是累加溢出和精度劣化。Tensor Core硬件本身支持FP32累加器也就是說矩陣乘法的中間累加過程是以FP32精度進行的兩個FP16/FP16或者BF16/BF16乘完后的結果累加到FP32的寄存器里最后再寫回輸出時按需轉回目標精度。這個機制在寫kernel的時候幾乎沒有額外成本完全是硬件提供的特性。但要注意一個問題BF16的尾數位比FP16少乘法結果的精度天然受限。在訓練場景中如果直接用BF16做全部前向計算某些對精度敏感的模型可能會出問題。所以DeepGEMM這類深度優(yōu)化內核在工程上的落地往往還要配合混合精度策略某些層用FP16/BF16做GEMM計算但在累加器、歸一化、殘差連接這些精度敏感的位置用FP32保存中間結果。這一點在復現(xiàn)時被很多人忽略了導致同樣的kernel在不同模型上效果天差地別。3.5 面向Tensor Core的指令調度不要小看指令排列順序在寫CUDA內核的時候指令順序看起來是編譯器幫你搞定的但真正做極致性能優(yōu)化時你會意識到編譯器的指令調度策略和手寫調度之間的差距。Tensor Core的矩陣指令雖然吞吐很高但它對硬件流水線的占用形態(tài)和普通FMA指令不一樣不能簡單地跟數據搬運指令交替著放。這個項目里比較有價值的調度經驗是把矩陣乘指令盡量連續(xù)地發(fā)出去中間避免插入太多依賴性的地址計算、比較指令把地址計算盡可能挪到不影響核心計算流水的位置。用CUDA的調度原語比如__pipeline相關的機制以及cp.async的group提交機制來批量管理異步拷貝讓硬件可以在指令隊列里看到足夠多的獨立任務從而把各級流水線都填滿。4. 實測性能對比與優(yōu)化效果驗證4.1 我的測試環(huán)境與對照組設置為了驗證DeepGEMM這類優(yōu)化思路的含金量我特地搭了一個對比測試環(huán)境。硬件是老熟人了一塊消費級GPU這里就用某N卡型號的通用描述來取代具體型號系統(tǒng)環(huán)境是標準的Linux 最新驅動編譯工具鏈用CUDA最新穩(wěn)定版本。矩陣規(guī)模我選了兩種有代表性的大GEMMMNK4096模擬全連接層和較大規(guī)模的密集計算窄長GEMMM512N4096K4096模擬Transformer中間層常見的非對稱shape。對照組的設置包括三檔naive實現(xiàn)每個人都能寫的最簡單三重循環(huán)版本、中等優(yōu)化版本只做shared memory tiling、DeepGEMM完整優(yōu)化版本。所有代碼用同一個編譯選項避免編譯器優(yōu)化差異影響比較公平性。4.2 結果數據與瓶頸分析先看大GEMM的結果。naive實現(xiàn)自然不用多說性能大概只有理論峰值的3%左右純粹是被全局內存帶寬按在地上摩擦。中等優(yōu)化版本做好shared memory分塊之后性能飆升到了大約理論峰值的40%-50%這個時候瓶頸已經不是內存帶寬了而是共享內存帶寬和計算指令的調度效率。到了DeepGEMM完整優(yōu)化版本實測性能提升到了理論峰值的80%以上如果把精度放寬到某些快速模式甚至能逼近90%。窄長GEMM的情況更有意思。naive和中等優(yōu)化版本和大GEMM的表現(xiàn)趨勢差不多但DeepGEMM優(yōu)化版本的收益更明顯。原因在于窄長場景下K維度特別長流水線優(yōu)化的空間更大。K維的每個小塊都可以在后臺異步搬運主計算流全程不空等延遲隱藏的效果非常顯著。下面這個表格是我這次測試的總結給個直觀的對比實現(xiàn)版本大GEMM實測TFLOPS窄長GEMM實測TFLOPS主要瓶頸naive三重循環(huán)極低個位數百分比極低個位數百分比全局內存帶寬shared memory tiling理論峰值約40%-50%理論峰值約35%-45%共享內存帶寬與調度DeepGEMM完整優(yōu)化理論峰值80%左右理論峰值85%左右指令發(fā)射與資源占用要注意的是這個比例是相對我手上這塊顯卡的理論峰值來算的不同硬件上比例會變化。但總體趨勢相當穩(wěn)定優(yōu)化深度越深窄長shape的收益優(yōu)勢越明顯。4.3 用Profiler找瓶頸的方法如果只給結論不給方法那這篇文章就不夠味了。這里分享一個我自己常用的定位流程。第一步先用GPU性能分析工具抓kernel運行時間看看時間占比最高的kernel是哪個。第二步分析kernel內部的瓶頸類型到底是memory bound還是compute bound。這個通過對比實際帶寬/占用率和理論峰值就能判斷。第三步如果是memory bound優(yōu)先優(yōu)化數據復用和異步搬運如果是compute bound優(yōu)先檢查指令調度、是否有不必要的類型轉換。這里有個很容易被忽略的點分析工具的Overhead可能影響到時序數據。小kernel尤其明顯開太高采樣的profiler會把kernel本身跑慢好幾倍。我一般是先用工具把range和kernel整體耗時看一遍覺得可疑了再關掉profiler用cudaEvent做手動計時交叉驗證。畢竟優(yōu)化目標永遠是真實端到端的性能不是分析工具報告出來的好看數字。5. 避坑實錄我在復現(xiàn)DeepGEMM時踩過的五個坑5.1 共享內存Bank Conflict的隱形懲罰復現(xiàn)過程中第一個讓我頭疼的是bank conflict。共享內存在物理上被分成了32個bank如果同一warp的線程訪問同一bank的不同地址就會發(fā)生沖突硬件需要把訪問串行化。這個問題最隱蔽的地方在于代碼邏輯完全正確結果也完全正確但性能就是上不去你很難第一眼就看到問題在這里。排查方法也比較傳統(tǒng)但有效在關鍵循環(huán)處手動檢查每個warp的訪問地址分布或者把共享內存的布局改成padding版本比如每個row多申請幾個元素的偏移讓同一行元素的bank分布被錯開。這個padding的技巧在DeepGEMM的實現(xiàn)里也用得很普遍。你要是發(fā)現(xiàn)自己寫的GEMM kernel性能比同配置的參考實現(xiàn)差一截先去查bank conflict大概率能發(fā)現(xiàn)問題。5.2 寄存器溢出導致本地內存拖慢一切寄存器文件是芯片上最快的存儲但數量有限。如果tile切得太大、每個線程持有的數據太多編譯器就會把一部分寄存器變量溢出到本地內存。本地內存雖然在指令層面看起來還是像普通內存訪問但實際上是走緩存和DRAM的速度比寄存器慢了不止一個量級。我在調256×256大tile時就遇到了這個問題。表面上寄存器占用率不高但運行時間反而變慢。查了編譯報告之后發(fā)現(xiàn)編譯器默默地生成了大量的本地內存訪問指令。解決思路是把tile調回合理尺寸同時手動限制每個線程持有的A、B片段數量讓編譯器的寄存器分配壓力保持在安全范圍內。這里針對不同GPU架構限制的閾值不一樣最好用編譯報告里的寄存器數同步驗證。5.3 K維切分太小導致流水線空泡流水線隱藏延遲的核心是計算和搬運并行。如果你的K維切分塊特別小每個異步拷貝的數據量很少很快就能拷完但計算階段還遠遠沒結束流水線實際上處于空轉狀態(tài)。這就像餐廳里配菜師傅切菜太快廚師炒菜跟不上切好的菜堆在臺面上白白等著。實際操作中應該讓每一輪流水線中搬運的時間盡量和計算時間匹配。切塊太大則搬運時間超過計算時間切塊太小則流水線切換開銷過大。DeepGEMM這類項目里常見的做法是先定一個基線切分尺寸然后用profiler測每一階段的實際耗時再手動做一輪二分搜索式的調整。如果框架允許也可以把這個tile尺寸設計成編譯期常量通過不同二進制之間的切換來找最優(yōu)配置。5.4 半精度格式的精度暗坑FP16和BF16都是16位但精度特性差別很大。FP16的尾數多、指數范圍小適合數值范圍變化不大的場景BF16的指數范圍和FP32一致但尾數少適合防止溢出、對精度相對寬容的場景。如果代碼在兩種格式之間切來切去卻沒有意識到它們的量化誤差不同訓練或者推理結果很容易出現(xiàn)莫名的偏差。我在一個模擬的量化推理項目里就吃過這個虧BF16輸入的GEMM kernel跑得很快但最終精度比FP16低了一個數量級。后來排查發(fā)現(xiàn)是該層對數值精度極其敏感BF16的尾數位不足以支撐。這個問題的解決方法不是換回FP16而是在關鍵層用FP32做部分累加和補償讓精度敏感的位置保留足夠的信息。5.5 編譯器自動優(yōu)化與手寫內核的博弈編譯器在-O3級別下確實會自動做一些循環(huán)展開、指令重排但它的優(yōu)化邏輯是保證正確性的前提下盡量快而不是在這個特定的GEMM場景里盡量快。所以很多手寫優(yōu)化看起來像是破壞了編譯器優(yōu)化——其實不是而是你在把編譯器的通用策略替換成針對性的專有策略。比如手動展開循環(huán)可以減少循環(huán)控制開銷但同時會增加寄存器壓力手動安排數據布局可以讓向量化加載更高效但也可能讓代碼變得不那么可讀。我見過有人執(zhí)著于純編譯器優(yōu)化拒絕使用任何手寫調度結果性能始終上不去。真實項目里手寫優(yōu)化和編譯器優(yōu)化是配合關系先用編譯器的自動向量化和循環(huán)展開做基線然后針對熱點循環(huán)做手工干預最后用profiler對比驗證每一項優(yōu)化是否真的有效。6. 從DeepGEMM中提煉的通用優(yōu)化方法論6.1 先用Profile建立基線再談優(yōu)化很多人拿到一個GEMM優(yōu)化的任務第一反應是翻手冊找最快的指令然后直接上手寫。我的經驗是反過來的先不寫任何花哨的代碼把最簡單的實現(xiàn)跑一遍用profiler看清楚數據到底是怎么流動的、瓶頸在哪。這個基線數據是所有后續(xù)優(yōu)化決策的錨點。比如一個簡單的deep learning框架里的GEMM算子M1024, N1024, K4096naive實現(xiàn)花1000微秒。先看帶寬利用率和計算利用率如果算下來理論需要的數據搬運時間只要50微秒而實際kernel耗時1000微秒那你該做的是減少訪問、增加復用而不是去調指令調度。如果計算利用率已經接近90%了那你要關注的是怎么把更多的有效計算塞進流水線而不是再去做數據復用。6.2 把硬件特性變成代碼設計的思維方式普通GEMM實現(xiàn)和DeepGEMM這類深度優(yōu)化實現(xiàn)之間最大的差距不是代碼技巧而是思維方式。普通實現(xiàn)是想好算法然后用代碼表達算法深度優(yōu)化是先吃透硬件的存儲層次、指令形態(tài)、調度機制再反過來想算法應該長成什么樣子。舉一個最簡單的例子矩陣是按行主序存儲的訪問A矩陣的一行和B矩陣的一列前者連續(xù)、后者跳變。如果直接按原始布局計算B列的訪問會帶來大量的cache miss。深度優(yōu)化實現(xiàn)的做法是先把B矩陣做了分塊轉置或者從一開始就用列主序或者特殊layout讓B塊在shared memory里也能連續(xù)訪問。這個決策的正確性完全來自對內存訪問模式和硬件cache行為的理解。6.3 什么時候該果斷放棄手寫優(yōu)化這一節(jié)是給所有人的清醒劑。手寫優(yōu)化能力很強但不是在所有場景都值得做。如果你的項目只是偶爾調一次shape不固定的GEMM或者矩陣規(guī)模太小、kernel啟動開銷本身就占大頭那手寫優(yōu)化的投入產出比很低。這時候直接用官方庫配合框架已有的融合能力反而是最理性的選擇。我的個人判斷標準是同一個GEMM操作在一個生產級系統(tǒng)里被調用的次數是否超過百萬次/天如果是手寫優(yōu)化能帶來1%的提升也是值得的如果不是把時間花在分析整體數據流、減少冗余計算、優(yōu)化內存拷貝上收益往往更明顯。說白了優(yōu)化本身也是有性價比的。6.4 擴展GEMM優(yōu)化思路在非矩陣場景的遷移DeepGEMM里體現(xiàn)的思想并不局限于GEMM。數據分塊的思想可以用在卷積的im2col轉化和Winograd變換上異步拷貝和雙緩沖的流水線思想可以用在Embedding層的大規(guī)模查表、Attention中的KV Cache讀取上寄存器級的數據復用思想可以用在很多elementwise算子的融合上。我在另一個序列處理項目里就參照了它的雙緩沖思路把一個BatchNorm激活函數Residual融合算子的性能提升了接近30%。這個算子里根本沒有矩陣乘法但優(yōu)化的核心邏輯完全一致——減少中間數據寫回全局內存的次數讓數據在寄存器里盡量多待一會兒。所以這篇文章講的是GEMM但方法論的適用范圍是幾乎所有深度學習算子優(yōu)化。7. 實操經驗總結關于DeepGEMM項目的一點心得整個DeepGEMM項目讓我感觸最深的一點是高性能計算的優(yōu)化本質上是在跟系統(tǒng)的默認假設對抗。默認的庫假設你不知道自己的工作負載默認的編譯器假設通用邏輯比特殊邏輯更值得優(yōu)化默認的數據布局假設訪存模式是均勻分布的。而一個手寫內核的價值就在于把這些默認假設全部推翻再基于真實場景重新設計。我在這次復現(xiàn)過程中最大的收獲不是那幾個性能提升的百分比數字而是真正建立起了從硬件倒推代碼設計的思維方式。以前我也覺得GEMM很神秘CPU上寫過高性能矩陣乘GPU上用過庫函數但總隔著一層紗?,F(xiàn)在自己從tile切分、流水線設計、寄存器分配到指令調度完整地走了一遍再回頭看那些官方文檔里看似枯燥的硬件特性說明確實能讀出味道來了。如果你也想動手試一次我建議從一個小目標開始先寫一個正確的GEMM再優(yōu)化到比naive版本快5倍然后再挑戰(zhàn)DeepGEMM里的高級特性。中間每一步都用profiler記錄數據別憑感覺。矩陣乘法這個題目老歸老但常做常新——不同硬件、不同精度、不同shape下永遠有可以繼續(xù)榨出來的性能空間。