化實踐)
GEMM通用矩陣乘法是深度學習訓練和推理里最底層的重型計算單元從全連接層到注意力分數計算本質都是在搬矩陣。很多人以為直接調官方閉源數學庫就萬事大吉但真到了性能優(yōu)化、算子融合、低精度推理這些階段手邊能自己掌控的矩陣乘法內核反而不夠用。DeepGEMM 就是我做的一個面向深度學習場景的高性能矩陣乘法內核集合目的很直接在常見形狀上把算力吃透同時把 epilogue 融合、量化支持這些推理引擎真正需要的功能做進去。這篇內容適合正在做推理引擎、寫自定義算子、或者想搞懂矩陣乘法性能瓶頸的朋友我盡量把從分塊策略到 Tensor Core 指令、再到精度對齊的整個思路講清楚。1. 為什么深度學習場景需要一份專屬GEMM而不是直接調官方庫1.1 GEMM 是深度學習的發(fā)動機再快也不嫌快先說一句很多人可能不太在意的背景知識。深度學習模型跑一次推理計算量絕大部分落在卷積和矩陣乘法上而卷積在底層實現時也往往會轉換成矩陣乘法來處理。所以 GEMM 的性能直接決定了一次前向推理的快慢。GEMM 的數學公式非常簡單C alpha * A * B beta * C。A 是 M×K 矩陣B 是 K×N 矩陣C 是 M×N 輸出矩陣??雌饋砭褪侨匮h(huán)的事但問題在于數據量遠超芯片的片上存儲。一顆現代 GPU 有幾百 GB/s 甚至幾 TB/s 的顯存帶寬但一個 4096×4096 的 FP16 矩陣就有 32MB整個模型矩陣動輒幾十個這樣的塊不可能全部塞進片上緩存。所以 GEMM 優(yōu)化的核心從來不是怎么算而是怎么搬數據搬得少、算得密。我選擇自己寫 DeepGEMM不是因為官方閉源庫不行而是因為它解決不了我的三個問題。第一融合算子。推理引擎里 GEMM 后面基本都跟著 bias、激活函數、LayerNorm、量化縮放這些操作官方庫只負責輸出 C 矩陣剩下的操作要我重新起一個 kernel 讀寫一遍數據顯存帶寬全浪費了。第二形狀控制。官方庫對超大連續(xù)矩陣調校得很好但碰到 M1 的推理場景、或者 N 比較小的窄矩陣性能優(yōu)勢就沒那么明顯了。第三可控性。我想針對自己的模型形狀做定制想在看到性能瓶頸時能精確知道是哪一行代碼在等待閉源庫做不到。1.2 DeepGEMM 的定位不是重新造輪子是造一套能改的輪子DeepGEMM 的定位是一套面向深度學習推理場景的 GEMM 內核模板集不是要替代官方庫在所有場合的工作而是重點覆蓋推理引擎里最常用到的幾種情況常見批量大小下的大矩陣乘、低精度輸入FP16/BF16甚至量化后的 INT8/FP8、以及需要把額外算子焊進 epilogue 的融合需求。做這套東西之前我先明確了一個原則先讓內核在一個小形狀上跑出明顯的性能再考慮泛化而不是一開始就試圖處理所有邊界情況。我最初只實現了 M4096、N4096、K4096 的基礎版本跑通之后再逐步往外擴。這樣做的原因是 GEMM 優(yōu)化的變量實在太多tile 尺寸、寄存器布局、流水線深度、訪存模式任何一個改動都可能改變性能特征如果不固定形狀去調參很難知道到底是哪個改動起了作用。2. 分塊調度把大矩陣切成能塞進芯片的小方塊2.1 從數學公式到三層存儲層級矩陣乘法最直觀的實現就是三重循環(huán)累加但這樣每個 A 元素會被讀 N 次每個 B 元素會被讀 M 次數據搬運量驚人。分塊的目標是提高數據復用A 矩陣的某一行會被 C 矩陣同一行的所有列使用B 矩陣的某一列會被 C 矩陣同一列的所有行使用所以把計算切成 M_BLOCK × N_BLOCK 的小方塊后這個小方塊計算只需要加載對應的 M_BLOCK×K 的 A 分塊和 K×N_BLOCK 的 B 分塊數據復用率從 1 提升到了塊尺寸級別。在 GPU 上分塊要分層進行。第一層是把整個 C 矩陣分成若干 M_BLOCK × N_BLOCK 的塊每個 block線程塊負責一個輸出分塊第二層是 block 內部把 K 維再切段每次從顯存加載小塊 A、B 到共享內存第三層是每個線程從共享內存取數據用寄存器算對應的輸出小片累加結果最終寫回顯存。這個三層結構恰好對應 GPU 的三種存儲層級顯存、共享內存、寄存器。2.2 一個實例算清楚分塊參數怎么定舉一個具體例子。假設目標 GPU 有 132 個 SM目標形狀是 MNK4096。我最初選的 block 尺寸是 128×128那 grid 就是 32×32 1024 個 block平均每個 SM 要處理約 7.8 個 block負載基本均衡。K 維切片長度 BLOCK_K 的選擇更講究。BLOCK_K 越大單次加載的數據越多訪問顯存效率更高但共享內存占用也隨之增加。128×128 的 C 分塊用 FP32 累加寄存器需要 128×128/32 線程 每線程 64 個 FP32 寄存器加上操作數寄存器寄存器壓力已經不小。所以 BLOCK_K 我選了 32這樣 A 分塊是 128×32×2 字節(jié) 8KBB 分塊是 32×128×2 字節(jié) 8KB加上 C 分塊共享內存占用大概 20KB 左右在 128KB 的共享內存里可以輕松放下雙層緩沖。提示BLOCK_K 一旦超過 64共享內存占用會急劇上升留給雙層緩沖的空間就緊張了。實際測試下來在多數數據中心級 GPU 上BLOCK_K 在 32 到 64 之間是最穩(wěn)的選擇區(qū)間。選好這些參數之后主循環(huán)的控制流就非常機械了外層遍歷 K 維切片內層把切到的 A、B 分塊從顯存搬到共享內存然后所有線程執(zhí)行矩陣乘的小片計算。問題是這種搬一塊、算一塊的做法搬數據的時候計算單元是空閑的計算的時候搬運單元是空閑的性能只有理論峰值的一半左右。解法就是第 3 節(jié)要說的雙層緩沖。3. 真正提速的核心細節(jié)張量核指令、寄存器排布、流水線預取3.1 mma 指令與 FP32 累加為什么低精度輸入要用高精度加法現代數據中心級 GPU 架構都有一個專門做矩陣乘法的硬件單元通常稱為張量核心它通過底層的 mma 指令一次性完成一個小矩陣的乘累加。比如一條 mma 指令可以完成 16×8×16 這樣的操作也就是 A 是 16×16、B 是 16×8算出 16×8 的輸出小矩陣。向量單元要一個時鐘周期做幾次乘加張量核心一個周期能完成一個片段的乘累加吞吐量完全不在一個量級。DeepGEMM 里我采用的指令模式是加載 FP16/BF16 輸入累加器用 FP32。這不是我拍腦袋定的——直接用 FP16 累加會在多次累加后出現明顯的舍入誤差尤其在 K 很大時誤差會累積。FP32 累加寄存器雖然占的地方多一倍但換來的是結果精度明顯提高實測中與官方庫的 FP32 累加結果完全相同。PTX 層的指令看起來大概是這樣mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 {%0, %1, %2, %3}, // 輸出累加器 16x8 的四個片段 {%4, %5}, // A 矩陣片段 16x16 {%6, %7}, // B 矩陣片段 16x8 {%0, %1, %2, %3}; // 輸入累加器這條指令的難點在于操作數布局是固定的A、B、C 的每個片段分別放在哪些寄存器、哪些位都有嚴格規(guī)定。如果寄存器排布不符合指令預期編譯器會補一堆 mov 指令來回倒騰數據性能直接打折。所以我寫 DeepGEMM 時是先按指令要求的寄存器布局來分配數據而不是先分配好再期望編譯器去適配。3.2 共享內存 bank 沖突看不見的串行化陷阱共享內存按 bank 分塊硬件在同一周期可以同時服務多個不同 bank 的訪問。但如果多個線程訪問的是同一個 bank 的不同地址吞吐就會變成原來的幾分之一這就是所謂 bank conflict。在 GEMM 內核里,加載 B 矩陣分塊時最常踩這個坑。經典案例以 FP16 類型為例BLOCK_N 為 128B 分塊是 BLOCK_K × 128。如果每個線程連續(xù)取一行中的相鄰元素兩個線程的地址恰好落在同一個 bank 上整個 load 就會被串行化。解決方法是給 B 矩陣的訪問模式加一個偏置移位讓相鄰線程訪問的地址錯開到不同 bank。DeepGEMM 里我用了一個簡單的 swizzle 策略矩陣存儲時同一行的數據不再連續(xù)排列而是按 XOR 變換重新組織。這樣線程在列方向并行讀取時地址映射后的 bank 號自然錯開實測共享內存帶寬利用率從約 60% 提升到接近 90%。以 C 語言偽碼來說明一個 8 元素 swizzle 的思路// 每個線程要讀取的元素所在行 row、列 col // 傳統(tǒng)布局address row * row_pitch col // swizzle 布局address row * row_pitch \ // ((col 7) ^ ((row 1) * 7))關鍵是讓相鄰線程的列地址和行本身的偏移撞不到同一個 bank。3.3 雙層緩沖讓搬運和計算真正重疊軟件流水線的目標是讓顯存到共享內存的拷貝和矩陣乘計算重疊執(zhí)行。最樸素的做法是共享內存里放兩份 A、B 分塊一份給當前循環(huán)迭代使用另一份給下一次迭代預取。主循環(huán)寫成這樣而循環(huán)里的操作步驟是用 cp.async 指令異步發(fā)起下一次迭代的 A、B 數據拷貝用當前共享內存緩沖區(qū)里的數據執(zhí)行 mma 指令完成矩陣乘提交計算完成等待之前的異步拷貝完成交換兩個緩沖區(qū)的角色進入下一輪cp.async 是 GPU 上的一種異步拷貝指令數據從顯存搬進共享內存不占用線程計算時間硬件會自動完成。第一次實現時我直接把 cp.async 換成普通 load性能下降接近兩成原因就是每次主循環(huán)結束都得等數據搬運完才能繼續(xù)算。這個經驗在 DeepGEMM 的開發(fā)中反復被驗證只要共享內存裝得下雙緩沖幾乎是零成本提升。4. 精度對齊與排錯鏈路從錯誤結果到逐項排查4.1 第一次跑通不代表結果是錯的我第一版 DeepGEMM 跑出了結果乍一看和官方庫的輸出差不多但差分對不過最大絕對誤差在 1e-1 量級對推理模型來說這直接是錯誤輸出。當時我最懷疑的是張量核指令用錯了后來發(fā)現其實是邊界處理的問題。排查的第一原則是固定變量。我把輸入矩陣改成隨機數固定種子先用官方庫算出基準 C_ref再用 DeepGEMM 算出 C_test逐項對照。拿 M、N、K 都比較小的用例開始比如 64×64×64這樣任何一行代碼的行為都能手動推算。4.2 排查流程從索引、同步到邊界我當時整理的排查順序現在也推薦給你檢查項方法我遇到的問題索引公式在小矩陣上逐元素核對 A、B 分塊加載確認 mma 指令的 A 矩陣行主序轉列主序時搞混了一次同步位置檢查 __syncthreads() 是否覆蓋跨線程共享數據讀寫有一次預取數據已經發(fā)起當前計算還沒用完舊緩沖區(qū)就做了緩沖區(qū)交換邊界處理M、N 不能被 block 整除時越界位置是否置 0128 的 block 處理 4096 沒問題換到 1000 就出錯了累加精度與 FP32 基準對比誤差是否在可接受區(qū)間FP16 累加誤差在 K4096 時放大到不可接受4.3 兩個最容易踩的坑越界加載與緩沖區(qū)交換先說越界加載。BLOCK 尺寸通常是 32 或 64 的倍數但矩陣 M、N 不一定是整數倍。當尾部塊加載 A 分塊時最后幾行已經超出矩陣邊界讀到的顯存內容是什么不確定算出的結果自然不對。解法是加載時做邊界判斷越界位置置零。這個判斷放在主循環(huán)里太貴我把它放在每次加載的 if 分支里if (row M col K) { A_tile[local_row * BLOCK_K local_col] A[global_row * K global_col]; } else { A_tile[local_row * BLOCK_K local_col] 0.f; // 越界置零不影響累加結果 }置零的好處是無論該位置被計算多少次對最終累加結果都沒有貢獻尾塊照樣可以走和完整塊相同的 mma 指令。再說緩沖區(qū)交換。雙緩沖實現里有一個非常隱蔽的 bug我在異步拷貝還沒完成時就把緩沖區(qū)的指針換掉了導致下次計算用的還是舊數據。排查這個問題的代碼路徑是用__pipeline_commit和__pipeline_wait_prior這對異步接口時忘記在讀取共享內存前等待對應批次全部完成。注意多級流水線里等待條件不是上一批拷貝完成而是我這次計算需要的那一批完成。流水線越深這個關系越容易搞錯。我后來用了一個比較實用的檢驗方法把 K 維的切片數改成奇數跑一遍如果結果和偶數切片不一致基本就是流水線同步有 bug。因為切片數量奇偶變化會改變緩沖區(qū)和循環(huán)迭代的對應關系任何等待順序錯誤都會導致結果異常。5. 性能實測用 profiler 數據驅動不拍腦袋調參5.1 基線對比要分形狀不能只看一個數GEMM 性能受形狀影響極大。我拿 DeepGEMM 和官方閉源庫做了對比測試固定數據格式為 FP16 輸入 FP32 累加分別測了三種典型形狀結果很能說明問題形狀M × N × KDeepGEMM 達到的峰值算力占比與官方庫相對性能4096 × 4096 × 4096約 78%約 95%1024 × 1024 × 4096約 70%約 88%1 × 4096 × 4096約 15%約 20%前兩個結果說明 DeepGEMM 在常規(guī)大矩陣上已經有競爭力最后一個 M1 的形狀則暴露了問題block 里大量線程計算同一行數據復用不夠內存加載變成了瓶頸。這個結果提醒我DeepGEMM 目前的架構并不適合下沉到 M 很小的推理場景還需要針對這種情況做獨立優(yōu)化。5.2 用 profiler 找到瓶頸而不是靠猜性能出問題的時候我先用 GPU 的性能分析工具采集三個指標SM 忙碌率、共享內存帶寬利用率、請求停滯分布。第一次跑 4096 形狀SM 忙碌率只有 55%共享內存帶寬利用率也不是特別高說明問題不在計算而在流水線等待。逐個排查后發(fā)現主循環(huán)里每輪計算結束后有一個等待異步拷貝的操作按我的預期這里應該已提前預取完畢。但分析工具顯示 stall 主要發(fā)生在共享內存訪問階段原因是加載 B 矩陣分塊時的 swizzle 只做了部分偏移仍有少量 bank 沖突。我把 swizzle 粒度從 4 元素調整到 8 元素之后共享內存帶寬利用率從 72% 上升到了 88%整體算力占比也提升了大約 8 個百分點。5.3 編譯器參數同樣影響性能寫 GEMM 內核時編譯參數的影響經常被低估。我用到的兩個關鍵選項一個是指定 GPU 架構讓編譯器針對具體指令集優(yōu)化另一個是開啟更大的寄存器限制允許寄存器溢出到局部的優(yōu)化。最開始我用默認編譯參數寄存器分配比較保守性能差了約一成。把寄存器上限調高到 255 之后編譯器能把更多的中間結果留在寄存器里而不是反復寫回共享內存。需要注意的是這里有個平衡點寄存器占用太高會導致 SM 上同時運行的線程塊變少并行度下降。我在目標 GPU 上實際測試128 線程一 block、每個線程 168 個寄存器左右是可以同時跑兩個 block 的上限區(qū)間。6. 往融合算子走一步GEMM 的 epilogue 定制6.1 推理場景真正需要的是 GEMM 偏置 激活如果 DeepGEMM 只能輸出一個 C 矩陣那它對推理引擎的價值就打折了。實際上全連接層之后的模式非常固定先加 bias再過 ReLU/GELU最后才是輸出。把這三個步驟從獨立算子合并進 GEMM kernel 的末尾段epilogue可以減少一次完整的顯存讀寫。我在實現上把 epilogue 做成一個模板化的函數傳入輸出塊的共享內存指針和模型配置由每個線程在算完自己的輸出片段后執(zhí)行。模板參數里包含是否需要 bias、使用哪個激活函數、是否要做量化這些全都在編譯期確定運行時零分支開銷。6.2 量化場景把 scale 和 zero point 焊進融合流程低精度推理里GEMM 輸出是 FP32但下一層要求 INT8 輸入中間必然有一次量化操作y round(x * scale zero_point)。常規(guī)做法是把 C 矩陣寫回顯存再啟動一個量化 kernel 讀出來重新算一遍帶寬和時間成本都很高。在 DeepGEMM 的 epilogue 中我在每個線程輸出 FP32 片段之后直接做這個量化再把結果寫回顯存。如果量化是 per-token 的即每行一個 scale也只需要把 scale 向量提前加載到共享內存每個線程按自己的行索引取數即可。融合后一次 kernel 跑完 GEMM 和量化實測整體的顯存讀寫量下降約一半端到端時間縮短了約 30%。6.3 融合的邊界與后續(xù)擴展方向融合也不是越多越好。層歸一化雖然也在 GEMM 之后常用但它需要跨通道計算均值方差意味著要拿到一行的完整輸出而 GEMM 的輸出片段是按 block 分散在不同線程里的強行融合會導致復雜的跨線程規(guī)約。我在 DeepGEMM 里沒有把 LayerNorm 合進去而是讓它保持獨立 kernel這也是很多推理引擎的普遍做法。這套內核后續(xù)可以擴展的方向我目前最關心的是更小的 block 尺寸以適配推理批次較小的場景以及把 FP8 輸入的計算路徑補全因為在最新一代硬件上 FP8 的性能收益非常明顯。另一個思路是為結構化稀疏做配套優(yōu)化稀疏矩陣在帶寬上的節(jié)省潛力比純稠密更大。我實際寫 DeepGEMM 的體會是矩陣乘法的性能優(yōu)化沒有什么玄學所有瓶頸最后都能落到訪存模式、寄存器分配、流水線等待這幾個具體問題上。關鍵是不要一開始就追求大而全而是從一個小形狀出發(fā)把 profiler 給出的數據一項項磨平。你先跑通一個 block 都行把同步、邊界、精度都驗證對了再往上加復雜度和融合特性這個過程會穩(wěn)很多。