計實踐)
1. 現(xiàn)狀CATLASS算子模板庫到底解決了什么問題這幾年做高性能計算的同行應(yīng)該都有同感算子開發(fā)已經(jīng)從“能跑就行”卷到了“必須榨干每一絲算力”。我們團(tuán)隊維護(hù)的這套自研算子模板庫內(nèi)部代號CATLASS說白了就是一套面向CUDA/GPU環(huán)境的C模板化算子框架專門用來快速生成高性能算子尤其是矩陣乘、卷積、歸約這類計算密集型和訪存密集型內(nèi)核。它借鑒了CUTLASS的思路但在調(diào)度策略、數(shù)據(jù)流編排和代碼生成層面做了一定程度的定制適配我們內(nèi)部的計算平臺和業(yè)務(wù)場景。CATLASS這個詞拆開看就是CUDATemplateLibrarySystem的合成但實際定位不只是“又一個GEMM庫”而是一套面向算子復(fù)用的基礎(chǔ)設(shè)施。過去寫一個高性能算子基本流程是通讀架構(gòu)手冊手工排布線程束和共享內(nèi)存調(diào)優(yōu)異步拷貝、流水線階段數(shù)、寄存器緩存策略測一遍性能然后換一個shape一切重來。有了CATLASS之后核心算子的計算主循環(huán)、數(shù)據(jù)搬移、切分調(diào)度被參數(shù)化成模板通過組合模板參數(shù)就能派生出不同規(guī)格的實現(xiàn)性能和手工調(diào)優(yōu)版本基本持平甚至在某些形狀下能超過。這套庫目前在我們團(tuán)隊內(nèi)部已經(jīng)覆蓋了三大類算子矩陣乘GEMM及其變體、卷積前向和反向、以及若干融合算子比如GELUGEMM、LayerNormGEMM。訓(xùn)練和推理側(cè)都有落地。整體代碼規(guī)模在五萬行左右核心模板頭文件大約二十多個配合一套構(gòu)建腳本和性能基線測試形成了從模板定義到benchmark回歸的完整閉環(huán)。很多剛接觸這套庫的同事會問一個問題現(xiàn)在cuBLAS、CUTLASS都開源了而且生態(tài)成熟、適配充分為什么還要自己搞一套這個問題其實正中要害。我的回答通常分兩層第一自研模板庫的核心價值不是“避免用第三方庫”而是“面對黑盒算子和高復(fù)雜度代碼生成框架之間提供一個中間層級”。cuBLAS是黑盒性能很好但你很難切入自定義融合邏輯CUTLASS是頂級的代碼生成庫但它的抽象層次極高迭代速度極快想要維護(hù)二次開發(fā)成本不小。CATLASS要做的就是用更輕量的抽象去覆蓋業(yè)務(wù)中最常出現(xiàn)的形狀和融合模式做到夠用、易改、可控。第二分布式訓(xùn)練和推理服務(wù)中經(jīng)常出現(xiàn)非規(guī)則shape和自定義數(shù)據(jù)類型需求模板庫能夠快速適配這是純黑盒庫做不到的。適合誰來參考這套思路我覺得有三類人一類是業(yè)務(wù)團(tuán)隊中負(fù)責(zé)底層算子優(yōu)化、但沒精力完整啃下CUTLASS龐大抽象層的人一類是在做AI推理引擎、正在設(shè)計自己算子層抽象的人還有一類是單純想理解高性能算子如何通過模板拆解實現(xiàn)復(fù)用的人。這篇文章我會把CATLASS的設(shè)計思路、關(guān)鍵實現(xiàn)細(xì)節(jié)、踩過的坑和后續(xù)規(guī)劃完整展開不吹不黑盡量還原我們做這套庫時的真實取舍。2. 核心設(shè)計思路為什么用模板來抽象算子而不是代碼生成或運行時調(diào)優(yōu)2.1 對比三條技術(shù)路線的取舍在CATLASS立項之前團(tuán)隊內(nèi)部其實認(rèn)真討論過三條技術(shù)路線第一是運行時調(diào)優(yōu)路線就是準(zhǔn)備幾十個kernel實現(xiàn)上線前跑一遍自動調(diào)優(yōu)選最優(yōu)配置運行類似cuBLAS的heuristic策略第二是離線代碼生成路線也就是用Python或者外部DSL描述算子的循環(huán)結(jié)構(gòu)和數(shù)據(jù)搬移然后生成CUDA C代碼投入編譯第三就是我們最終選擇的模板元編程路線把算子的結(jié)構(gòu)拆成編譯期常量組合通過模板參數(shù)實例化出不同實現(xiàn)。三者的核心差別在“調(diào)優(yōu)決策發(fā)生在哪一層”。運行時調(diào)優(yōu)最靈活但代價是顯存占用爆炸、啟動延遲上升而且對于融合算子這種需要跨層感知的情況預(yù)置方案很難枚舉齊全。離線代碼生成最“自由”但會引入代碼生成鏈路的維護(hù)成本調(diào)試排錯多一環(huán)而且生成的代碼常常不夠穩(wěn)定容易被編譯器優(yōu)化差異搞崩。模板方案則把變化點收斂到類型參數(shù)上沒有額外代碼生成環(huán)節(jié)編譯器看到的是實實在在的C源碼調(diào)試體驗最接近手寫kernel。但模板方案的缺點也非常明顯抽象層級一旦沒設(shè)計好模板參數(shù)數(shù)量會指數(shù)膨脹代碼可讀性直線下降編譯時間暴漲。這個問題我們吃了不少苦頭后面3.2小節(jié)會詳細(xì)講怎么控制模板復(fù)雜度。選型結(jié)論是核心高頻算子用模板組合邊緣場景和一次性實驗用腳本生成兩條路線并存但主路徑始終是CATLASS模板。2.2 一個核心GEMM模板的參數(shù)拆解拿我們最常跑的FP16 GEMM算子舉例CATLASS的kernel入口長這樣template typename ElementA, // 輸入A矩陣元素類型比如 half typename ElementB, // 輸入B矩陣元素類型 typename ElementC, // 輸出C矩陣元素類型 typename ElementAccum, // 累加器類型通常是 float typename TileShape, // 線程塊級tile形狀比如 Shape128, 128, 32 typename WarpShape, // warp級tile形狀 typename ThreadShape, // 線程級tile形狀 typename StageCount, // 流水線階段數(shù)比如 3 或 4 typename SchedulePolicy, // 主循環(huán)調(diào)度策略 bool SwapAB // 是否交換A/B加載角色 __global__ void GemmKernel(const ElementA* A, const ElementB* B, ElementC* C, ...);第一次看到這一大串模板參數(shù)的同事通常都會懵。我解釋的時候喜歡打個比方這就像做菜菜譜固定但你可以選食材種類Element類型、切菜大小TileShape、灶臺數(shù)量StageCount和顛勺節(jié)奏SchedulePolicy模板就是把這些選擇提前到“點菜下單”階段而不是做菜過程中臨時更改。TileShape里的三個維度分別是線程塊在M維、N維、K維上的分塊大小。選多少不是拍腦袋M維和N維的乘積決定了線程塊并行度要和GPU的SM數(shù)量、寄存器預(yù)算匹配K維則直接影響數(shù)據(jù)復(fù)用率K太小則每次從全局內(nèi)存搬入的數(shù)據(jù)很快被消費完K太大則Shared Memory容納不了所需分塊。我們以A100為例做過一個估算A100的SM Shared Memory是164KB可配置到最多164KB如果TileShape選128x128x32用FP16存儲A和B的tile各占128*32*28KB雙份就是16KB三級流水線就要48KB余量留給double buffer和同步開銷是比較舒服的組合。如果上到256x256x64光A/B tile就是256KB直接爆顯存所以這種激進(jìn)形狀只有在特定大L2的卡上才能考慮。模板設(shè)計中最需要小心的是“可組合性”和“可用性”的平衡。我們的做法是做分層抽象不是一把梭把所有參數(shù)堆在一個模板上。底層是數(shù)據(jù)搬移原語和計算原語中間層是線程塊級調(diào)度和Warp級調(diào)度頂層才拼裝成完整的GemmKernel。這樣底層原語可以獨立測試和重排不同上層策略能復(fù)用同一套搬移代碼開發(fā)新算子的成本從兩周縮減到兩三天。2.3 對比CUTLASSCATLASS做了哪些取舍說到算子模板庫繞不開CUTLASS。CATLASS立項時深度參考了CUTLASS 2.x的設(shè)計在概念層面高度一致比如Tile抽象、Warp布局、Shared Memory迭代器等。但我們在三個點上做了主動簡化第一砍掉了“Collection”和“Complex”這類為追求極致彈性而設(shè)計的抽象層級。CUTLASS為了支持任意算子形態(tài)引入了很多中間類型這對框架維護(hù)者來說是合理的但對業(yè)務(wù)開發(fā)者來說心智負(fù)擔(dān)太重。CATLASS只保留“把它變成GEMM能解決的問題”這條主線其他形態(tài)走適配層轉(zhuǎn)接到GEMM模板上讓90%的日常需求落在一條主路徑上。第二調(diào)度策略沒有做成完全可插拔的泛型而是內(nèi)置了少數(shù)幾種經(jīng)過驗證的SchedulePolicy枚舉。CUTLASS的調(diào)度器是高度模板化的策略類你可以通過不同的策略組合實現(xiàn)warp-synchronous、warp-specialized、ping-pong等模式。CATLASS也支持這些模式但對外只暴露四五種預(yù)設(shè)策略內(nèi)部實現(xiàn)通過if constexpr分發(fā)到不同代碼路徑。這樣可以大幅減少模板實例化分支編譯時間從CUTLASS動輒幾分鐘一次降到幾十秒。第三數(shù)據(jù)類型的適配范圍更聚焦。CUTLASS支持從FP64到int4的廣泛數(shù)據(jù)類型以及各種mixed-precision組合CATLASS首選支持的是FP16和BF16輸入、FP32累加這個AI計算最常見組合FP8還在完善中int8目前走的是另一套SGEMM模擬路徑。聚焦帶來的好處是我們可以在內(nèi)存對齊、向量化加載和歸約順序上針對這幾個類型做激進(jìn)優(yōu)化減少模板分支。3. 核心細(xì)節(jié)解析從Tile切分到寄存器緩存這些地方?jīng)Q定了性能上限3.1 線程塊級、Warp級、線程級的三層映射高性能算子的本質(zhì)是重復(fù)利用數(shù)據(jù)。你從全局內(nèi)存里搬一個數(shù)據(jù)進(jìn)寄存器或Shared Memory總希望能多算幾次再扔。CATLASS的映射策略就是圍繞“復(fù)用”二字展開的。線程塊級ThreadBlock映射負(fù)責(zé)確定一個C tile由哪些線程塊計算典型是128x128的C tile然后由4個warp每個warp負(fù)責(zé)64x64或者8個warp每個warp負(fù)責(zé)32x64劃分。線程塊之間完全獨立不需要通信這是scheduling友好的基礎(chǔ)。Warp級映射決定每個warp內(nèi)部32條線程如何協(xié)作。以WarpShape64, 64為例每個線程最終負(fù)責(zé)的C元素個數(shù)是64*64/32128個。這128個元素不會連續(xù)分配給一個線程而是分散成多個8x8或4x8的微塊交錯放在warp的lane上。這么做的目的是讓同一時刻相鄰lane訪問的Shared Memory地址盡可能分布在不同Bank上減少Bank Conflict。我們在一個內(nèi)部測試?yán)飳Ρ冗^同樣計算量下布局方案從連續(xù)分配改成交錯分配后Shared Memory的訪存效率提升了32%GEMM整體性能漲了8個百分點。線程級映射是最底層的計算粒度決定了每個線程在寄存器里的數(shù)據(jù)布局。這里有一個非常關(guān)鍵的經(jīng)驗**寄存器里的C矩陣布局決定了累加時是否會產(chǎn)生寄存器Bank Conflict也決定了最后寫回全局內(nèi)存時能否走STG.128向量化寫。**我們用float4對齊的布局將每個線程的8個C值組織成兩組float4寫回時剛好一條st.global.v4.f32指令搞定。三層映射之間的關(guān)系可以用俄羅斯套娃來理解塊套warpwarp套線程每層都遵循同一個原則——計算密度和訪存密度的比值要盡量大。如果某一層計算太少就會造成同步開銷相對過高出現(xiàn)“花大量時間等人”的局面。3.2 Mainloop多級流水線怎么排才能讓數(shù)學(xué)單元永不空等GEMM kernel的性能核心在主循環(huán)Mainloop。一次典型的mainloop迭代干三件事從全局內(nèi)存加載A和B的下一塊數(shù)據(jù)到Shared MemoryLoad階段把Shared Memory中的數(shù)據(jù)搬進(jìn)寄存器并做乘累加Compute階段以及維護(hù)各個階段的同步點。CATLASS使用多級流水線隱藏訪存延遲默認(rèn)三級流水較極端的情況用四級。流水線的核心思路是讓Load和Compute重疊。用生產(chǎn)者-消費者模型理解Load階段是生產(chǎn)者Compute階段是消費者。三級流水意味著Shared Memory里同時維護(hù)三份A/B tile一份正在被load填充、一份已經(jīng)就緒等待計算、一份正在被消費。這樣計算單元拿到數(shù)據(jù)后不必等待下一次全局內(nèi)存訪問延遲被掩蓋在流水線里。實現(xiàn)細(xì)節(jié)上我們是靠cuda::pipeline原語 手動cp.async指令配合完成的。這里有個很大的坑cp.async的commit/wait組管理。commit表示一批異步拷貝已經(jīng)發(fā)出wait則等待某批完成。如果批次數(shù)和實際申請的Shared Memory buffer對不上輕則數(shù)據(jù)錯亂重則死鎖。我們的做法是在StageCount模板參數(shù)里帶上對應(yīng)批次數(shù)作為編譯期常量通過#pragma unroll展開循環(huán)確保commit和wait嚴(yán)格一一配對。與之關(guān)聯(lián)的是“屏障”的放置。__syncthreads()是線程塊級屏障但多級流水線中不同warp可能處于不同階段全量同步會把流水線拉平喪失overlap效果。我們的解決方案是在warp specialize模式下只對生產(chǎn)者warp和消費者warp之間的共享buffer做輕量級屏障barrier arrive/wait讓不同warp各忙各的。這算是對CUTLASS中warp-specialized策略的一種簡化實現(xiàn)性能提升顯著。3.3 Shared Memory布局與Bank Conflict規(guī)避實戰(zhàn)Shared Memory是GP U上最緊俏的存儲資源同時也是最容易出現(xiàn)性能陷阱的地方。CATLASS在這塊的實踐可以濃縮成三句話數(shù)據(jù)要對齊訪問要分散padding要到位。先對齊。Bank寬度是4字節(jié)一個warp訪問Shared Memory時硬件會把32條線程的訪問請求按地址分到32個Bank上。如果線程訪問的地址剛好都落在同一個Bank就發(fā)生沖突變成串行訪問。CATLASS中所有Shared Memory數(shù)組都按16字節(jié)對齊分配確保向量化load/store不會跨Bank邊界。再分散。以FP16為例兩個half拼成4字節(jié)正好占一個Bank。當(dāng)線程lane_id想訪問A[row][col]時如果col對所有l(wèi)ane相同就會出現(xiàn)所有線程同時訪問同一行的不同列地址這在Bank層面可能是好的但如果行偏移設(shè)計不好就會撞地址。我們的經(jīng)驗是A矩陣按天真的行主序存儲但訪問時搭配Swizzle模式把地址打散。具體做法對每個tile內(nèi)部把原本行連續(xù)的地址通過異或操作映射到不同的Bank組合這樣連續(xù)lane訪問的地址在Bank上均勻分布。Padding是最粗暴也最有效的兜底方案。我們在共享內(nèi)存數(shù)組的每行末尾加一個元素的padding把行寬從對齊寬度變成非對齊寬度讓連續(xù)行的起始Bank號錯開。這招看似簡單實測下來能將極端情況下的Bank Conflict從8路沖突降到1路性能直接翻倍。不要小看這一行代碼很多開源實現(xiàn)里為了省那一點Shared Memory不加padding結(jié)果性能反而更差。3.4 寄存器緩存與指令級并行壓榨ALU利用率的最后一公里Mainloop計算階段的最后瓶頸往往不在Shared Memory而在寄存器和指令調(diào)度。每個線程從Shared Memory裝載A/B的fragment后乘累加操作分布在多個獨立的依賴鏈上。如果依賴鏈過長每周期ALU可能空等數(shù)據(jù)如果依賴鏈過短且沒有足夠多的尾數(shù)指令流水線又會堵塞。我們在模板中默認(rèn)讓每個線程在K方向上一次處理4個配合FP16的half2向量化累加器保持8到16個獨立fragment。這能讓編譯器有足夠的指令級并行ILP空間去隱藏FMA指令的延遲。寄存器分配上累加器只用float或float4不引入額外寄存器副本搬運用的臨時寄存器用完即棄避免寄存器溢出到Local Memory。這里有一個我們自己踩過的坑某次為了減少Active Warps數(shù)量達(dá)到更高單核頻率把每個線程的C tile從8x8改成16x8導(dǎo)致每個線程的累加器數(shù)量從8個變成32個。寄存器壓力直接爆表kernel occupancy從50%降到了25%最終性能不升反降。后來我們才意識到寄存器緩存不是越多越好而是在不降低占用率的前提下盡量多。GPU是“用并行換延遲”的機器拋棄占用率去追求單線程ILP是舍本逐末。3.5 融合算子的處理思路不是所有融合都要拼進(jìn)GEMMCATLASS里還有一類高頻需求是融合算子典型如GELU和GEMM的融合。很多框架的做法是在GEMM后面接一個獨立的activation kernel數(shù)據(jù)先寫回全局內(nèi)存再讀出來做GELU白白多一遍全局內(nèi)存往返。CATLASS的做法是在GEMM的epilogue階段把C累加器經(jīng)過激活函數(shù)處理后直接寫回省掉中間商。這里要強調(diào)一個設(shè)計原則能融到epilogue里的操作就融進(jìn)去需要跨整個tile統(tǒng)計的操作不要硬融。例如GELU、ReLU、LayerNorm中的per-row均值方差前者元素間無依賴可以逐線程處理放到epilogue非常合適后者需要跨同一行所有線程做歸約如果硬融進(jìn)GEMM需要在epilogue階段額外引入一次跨線程通信Shared Memory占用和同步開銷都會上漲。我們的折衷方案是保留獨立kernel做per-row的統(tǒng)計但讓GEMM算子直接輸出到L2友好的中間布局減少后續(xù)kernel的訪存開銷。融合判斷有一個經(jīng)驗公式當(dāng)融合引入的額外Shared Memory/同步操作帶來的開銷小于省掉一次全局內(nèi)存讀寫帶來的收益時才值得融合。計算時可以粗略估計一次全局內(nèi)存訪問的耗時幾百個cycle再和同步/歸約的開銷做對比心里就有數(shù)了。4. 工具鏈與工程化模板庫要真正落地光有漂亮源碼不夠4.1 編譯期校驗與靜態(tài)斷言把錯誤留在編譯期模板庫最大的痛點之一是錯誤信息晦澀。實例化失敗時編譯器有時只給一個十幾個模板層深的報錯新人基本看不明白。CATLASS的解法是在模板入口放置大量static_assert把形狀合法性、類型組合規(guī)則、對齊要求、Shared Memory大小上限等在編譯期就檢查完畢。例如TileShape的M/N維度必須能被WarpShape整除否則直接斷言失敗并輸出提示信息ElementAccum的字節(jié)數(shù)不能小于輸入類型的字節(jié)數(shù)防止累加精度丟失StageCount對應(yīng)的Shared Memory總量必須小于目標(biāo)架構(gòu)的Shared Memory上限。這些校驗讓大部分錯誤在CI編譯階段就暴露而不是等到算子跑起來才發(fā)現(xiàn)數(shù)值不對或直接非法內(nèi)存訪問。寫static_assert還有一個隱藏好處它能倒逼模板設(shè)計者把“隱式約定”變成“顯式約束”。早期我們有一些模板組合依賴調(diào)用者遵循不成文的規(guī)則比如“K維度必須是16的倍數(shù)”“Shared Memory buffer數(shù)必須是2的冪”結(jié)果不同業(yè)務(wù)團(tuán)隊各寫各的約束經(jīng)常被打破。后來把這些規(guī)則全部落地成static_assert問題立刻少了大半。4.2 自動化基準(zhǔn)測試與性能回歸算子模板庫和普通業(yè)務(wù)代碼有一個本質(zhì)差別普通代碼只要邏輯正確就完成了大半目標(biāo)而算子模板庫必須把“性能正確性”當(dāng)作一等公民來對待。邏輯正確但性能差10倍的算子在業(yè)務(wù)上幾乎不可用。因此CATLASS配套了一套相對嚴(yán)格的基準(zhǔn)測試體系?;鶞?zhǔn)測試要做三件事正確性校驗、性能采集、回歸比對。正確性校驗使用CPU參考實現(xiàn)對每個模板實例生成隨機輸入和邊界case比如全零、全極小、K方向長度極小和極大比對輸出誤差。性能采集使用CUDA Event計時同時配合Nsight Compute的SM占用率、Shared Memory吞吐、指令吞吐等硬件計數(shù)器一同記錄?;貧w比對則是把每次提交的性能數(shù)據(jù)與基線庫中保存的歷史最優(yōu)值做對比性能下降超過容忍閾值就在CI中標(biāo)記失敗。這套體系剛上線時也遭過抵觸跑一批模板實例的benchmark要十幾分鐘CI耗時暴漲。后來我們做了分級提交級只編譯不跑benchmark夜間跑全量benchmark并生成趨勢報告。這樣既保證性能變化能被及時發(fā)現(xiàn)又不會拖慢日常開發(fā)節(jié)奏。4.3 自動調(diào)優(yōu)器讓模板組合在幾百個候選中找到最優(yōu)解模板庫雖然可以通過參數(shù)組合實現(xiàn)不同變體但人工遍歷所有組合不現(xiàn)實。舉例來說TileShape有5個候選WarpShape有8個候選StageCount有3個候選SchedulePolicy有4個候選排列組合就是480種每種跑一遍benchmark要好幾秒人力根本做不完。CATLASS落地了一個簡單的自動調(diào)優(yōu)器用貝葉斯優(yōu)化在參數(shù)空間里搜索最優(yōu)配置并將結(jié)果緩存到配置文件里運行時通過hash后的shapeGpuModel索引直接查表。自動調(diào)優(yōu)器的設(shè)計原則是“離線調(diào)優(yōu)在線查表”。每次調(diào)優(yōu)的結(jié)果都會帶上GPU型號、驅(qū)動版本、計算庫版本作為上下文存入本地數(shù)據(jù)庫。這樣即使換了機器或驅(qū)動也不會錯誤套用不匹配的參數(shù)。調(diào)優(yōu)器還考慮到不同業(yè)務(wù)場景對延遲和吞吐的偏好差異優(yōu)化目標(biāo)函數(shù)支持兩種模式一種是最小化延遲適合在線推理場景一種是最大化吞吐適合離線批處理場景。兩個模式搜索到的參數(shù)往往不同比如延遲優(yōu)先模式傾向于小tile加4級流水吞吐優(yōu)先模式則傾向于大tile加3級流水。5. 規(guī)劃短期痛點、中期能力、長期形態(tài)5.1 短期規(guī)劃補FP8、補齊稀疏和Attention算子先說內(nèi)部最急迫的幾個需求。FP8推理在業(yè)務(wù)側(cè)的呼聲越來越高論壇上關(guān)于FP8格式的討論密度也很高我們計劃在下一個版本讓CATLASS的GEMM模板完整支持E4M3和E5M2兩種FP8格式累加器仍用FP32權(quán)重和數(shù)據(jù)在送入kernel前完成quantize。表面上看只是多加一個Element類型實際上涉及Shared Memory的存儲密度、向量化加載寬度和NVLink傳輸時的位寬對齊等多個地方的調(diào)整工作量不小。稀疏算子是另一條線。我們計劃支持2:4結(jié)構(gòu)化稀疏的GEMM也就是每4個元素里只有2個非零的稀疏模式。CUTLASS已經(jīng)證明這種模式可以利用稀疏張量核心獲得接近2倍的算力提升但模板抽象要處理好“元數(shù)據(jù)布局”和“非零元素選取”兩層邏輯。我們的初步方案是參考CUTLASS的SparseTile設(shè)計但把元數(shù)據(jù)布局從類型參數(shù)中剝離出來留給業(yè)務(wù)側(cè)根據(jù)數(shù)據(jù)分布自行選擇。Attention算子的需求來自我們的LLM推理引擎。目前FlashAttention類kernel在長序列場景下效果很好但它是自成體系的獨立算子和CATLASS的模板體系互不相通。我們打算把attention的前向主循環(huán)抽象成“QK^T分塊乘、Softmax、PV分塊乘”三步分別復(fù)用CATLASS的GEMM主循環(huán)和epilogue機制。這個短期版本的目標(biāo)是能覆蓋主流attention變體性能達(dá)到FlashAttention-2的90%以上。5.2 中期規(guī)劃擴展自動調(diào)優(yōu)能力和多平臺適配自動調(diào)優(yōu)器目前只能搜索有限幾個模板參數(shù)中期我們希望把優(yōu)化空間擴展到“算法選擇”層面。比如同一個GEMM問題在A100上可能最適合wgmma路徑在上一代架構(gòu)上可能最適合simt路徑這兩條路徑在CATLASS內(nèi)部是兩套完全不同的主循環(huán)實現(xiàn)?,F(xiàn)階段選型的邏輯是硬編碼在調(diào)度器里的不夠靈活。我們計劃讓調(diào)優(yōu)器自動從“路徑”維度做選擇并引入離線訓(xùn)練的性能模型來預(yù)估在一個沒見過的新GPU型號上的最優(yōu)配置。多平臺適配也在規(guī)劃中。AMD的ROCm平臺、intel的oneAPI平臺在我們客戶的機器上有現(xiàn)實需求。模板庫的一個天然優(yōu)勢是核心邏輯只依賴并行編程模型語義理論上可以通過封裝層適配到不同后端。當(dāng)然真正落到代碼上cp.async、wgmma這些指令在不同后端上的對應(yīng)實現(xiàn)差異巨大不可能完全無縫遷移。我們的思路是保持CATLASS上層API不變下層把平臺相關(guān)指令封裝成Backend接口先實現(xiàn)HIP后端驗證可行性。這里也提醒一句多平臺適配的投入產(chǎn)出比需要謹(jǐn)慎評估。如果目標(biāo)平臺的業(yè)務(wù)量不大硬適配的成本可能遠(yuǎn)大于收益不如直接讓CUDA版本跑在一個兼容層上。5.3 長期形態(tài)向開源生態(tài)靠攏同時保持內(nèi)部定制能力內(nèi)部庫最大的風(fēng)險是封閉導(dǎo)致衰退。CATLASS長期規(guī)劃中的一項核心動作是選擇一個合適的時機把核心模板層的代碼清理后開源。開源的目的不只是回饋社區(qū)更現(xiàn)實的意義是能引入外部貢獻(xiàn)者的review和測試擴大benchmark覆蓋面和硬件適配范圍。一旦開源我們會在源碼層面明確區(qū)分“核心可移植層”和“內(nèi)部定制層”核心層由社區(qū)共同維護(hù)定制層保留在我們的私有分支上。從另一個角度看模板庫的長期生命力取決于“表達(dá)力”。現(xiàn)在只能描述GEMM-like算子長期我們希望CATLASS能描述更廣泛的算子拓?fù)浔热缍噍斎攵噍敵龅膹?fù)雜算子組合。這會觸及模板元編程的表達(dá)邊界我們也在觀察C20的concepts和編譯期反射P2996提案等能否降低這方面的抽象成本。如果標(biāo)準(zhǔn)落地順利CATLASS的類層次和約束檢查會有一次重構(gòu)機會。6. 避坑指南與經(jīng)驗之談做算子模板庫這件事本身比想象的難6.1 團(tuán)隊協(xié)作中最容易翻車的3個點算子模板庫的開發(fā)和普通應(yīng)用開發(fā)對團(tuán)隊能力的要求完全不同。我復(fù)盤下來最容易翻車的點集中在下面三個地方。第一模板抽象失控。有位同事曾經(jīng)把調(diào)度策略設(shè)計成一整套泛型方案每個warp的調(diào)度狀態(tài)用類型組合描述代碼確實優(yōu)雅但實例化后編譯一個kernel要5分鐘報錯信息長達(dá)三百行沒人改得動。后來我們定了一條硬性規(guī)定每個模板新增前必須寫清“它替調(diào)用者解決了什么問題”如果回答不上來就不允許進(jìn)主干。這條規(guī)定其實來自CTO的一句玩笑話“模板參數(shù)的多少和代碼作者對這問題的理解程度成反比?!钡诙阅芑貧w被忽視。算子庫最容易被盯上的指標(biāo)是單算子性能但一旦模板被很多業(yè)務(wù)復(fù)用一個基礎(chǔ)類型的小改動會影響所有上層算子。曾經(jīng)有一次修改了Shared Memory的Swizzle函數(shù)單獨測新算子是提升的但老算子的吞吐普遍掉了5%。當(dāng)時沒有性能基準(zhǔn)門禁問題上線兩周后才被發(fā)現(xiàn)?,F(xiàn)在我們已經(jīng)強制所有改動必須在未做benchmark的情況下合入。第三文檔和實例代碼跟不上。模板庫的“API可發(fā)現(xiàn)性”天然比普通代碼庫差最好最實用的“文檔”其實是cookbook式的示例程序。我們?yōu)榇司S護(hù)了一個examples目錄每個示例對應(yīng)一個真實業(yè)務(wù)場景比如“動態(tài)shape場景下的GEMM調(diào)用方式”、“融合LayerNorm的推理算子”。每次模板接口變更examples必須同步更新這條規(guī)則雖然簡單但非常有效。6.2 一些“反直覺”但真實有效的細(xì)節(jié)有幾個細(xì)節(jié)是在多次調(diào)優(yōu)中發(fā)現(xiàn)的看起來反直覺但對最終性能影響很大。一個是用__launch_bounds__控制kernel的寄存器上限時數(shù)值上不要卡在編譯器的臨界值。比如某個kernel自然編譯需要96個寄存器設(shè)置__launch_bounds__(256)表示最多允許256個線程/塊編譯器會把寄存器限制在64個以內(nèi)這會導(dǎo)致溢出。反而設(shè)置__launch_bounds__(192)讓編譯器有96個寄存器的余量雖然occupancy低了一些但性能反而更高。經(jīng)驗法則是不要一味追求高occupancy而是要在Registers Per Thread和Occupancy之間找到實際運行最快的平衡點。另一個是主循環(huán)的#pragma unroll不是越高越好。#pragma unroll 8和#pragma unroll 4的性能在某些shape上差異明顯但同一份kernel在A100和H100上的最優(yōu)unroll數(shù)不同。自動調(diào)優(yōu)器除了搜Tile參數(shù)也會把unroll factor作為候選參數(shù)一起搜索。還有一點是關(guān)于L2 Cache的利用。GEMM的A矩陣和B矩陣在全局內(nèi)存里的布局對L2命中率影響很大。把A和B按tile順序做一次重排blocked layout后L2命中率能提高10%到20%。這個優(yōu)化只改數(shù)據(jù)布局不改kernel代碼性價比極高。我們的GemmHost接口默認(rèn)會做這一步但允許業(yè)務(wù)側(cè)通過參數(shù)關(guān)閉以節(jié)省預(yù)處理時間。6.3 常見的錯誤解析與排查思路很多新手在集成CATLASS時遇到“kernel不work”之類的問題這里列幾個高頻case和排查思路按出現(xiàn)頻率排序。第一個是數(shù)據(jù)類型不匹配導(dǎo)致的計算錯誤。比如用BF16輸入但累加器誤設(shè)成half在數(shù)值上不會報錯但精度誤差很容易在幾十步迭代后放大。排查時優(yōu)先打印累加器類型的sizeof再對比輸入類型的數(shù)值范圍。第二個是Shared Memory分配超出上限。這個問題在換GPU型號后最容易出現(xiàn)。排查方法是查啟動時的cudaFuncAttributes.sharedSizeBytes如果超過目標(biāo)卡的限制編譯期不會報錯運行時會直接啟動失敗。我們的模板里對SharedMemory占用做了static_assert至少能把問題定位到模板參數(shù)上。第三個是Bank Conflict導(dǎo)致的性能異常它不會報錯只是變慢。判斷方法是用Nsight Compute看Shared Memory Conflict計數(shù)器如果每周期沖突次數(shù)明顯高于預(yù)期就去檢查Swizzle函數(shù)和Padding是否生效。我們遇到過把Swizzle函數(shù)寫成了“對某些地址是映射到相同bank”的case排查時盯計數(shù)器指標(biāo)會高效很多。第四個是流水線同步問題表現(xiàn)為“偶爾出錯、偶爾正?!薄_@類問題最難排查因為可能與驅(qū)動調(diào)度、內(nèi)核啟動參數(shù)有關(guān)。我們的經(jīng)驗是先做最小化復(fù)現(xiàn)固定一個shape和一組模板參數(shù)反復(fù)跑幾百次然后添加?xùn)艡诖蛴£P(guān)鍵buffer的校驗和縮小問題范圍。大多數(shù)情況下最后定位到cp.async的wait批次計數(shù)錯誤而不是底層的硬件問題。7. 后續(xù)還可以怎么做從模板庫到算子生態(tài)最后聊一點我個人在維護(hù)CATLASS過程中的體會。算子模板庫這件事最難的不是寫出一個性能出色的kernel而是設(shè)計出一套“能在不同業(yè)務(wù)、不同硬件、不同需求之間穩(wěn)定復(fù)用”的抽象層。這條路走到現(xiàn)在我們最大的收獲不是那幾萬行模板代碼而是踩坑后沉淀下來的判斷力什么時候該用泛型去抽象什么時候該簡單堆代碼解決問題。如果你也在做類似的方向我的建議是先從業(yè)務(wù)側(cè)最高頻的10個算子入手把它們的手寫實現(xiàn)抽成可參數(shù)化的模板不要一開始就追求CUTLASS那種大而全的設(shè)計。模板庫是在反復(fù)迭代中慢慢長出來的不是一次設(shè)計出來的。等你的模板參數(shù)開始能覆蓋新出現(xiàn)的需求而不需要改底層時你才算真正摸到了門道。CATLASS目前還在持續(xù)演進(jìn)。最近我決定把自動調(diào)優(yōu)器的搜索結(jié)果做成一版可視化報表方便業(yè)務(wù)團(tuán)隊直接看懂每個配置在什么場景下最優(yōu)。再往后我希望CATLASS能和部署側(cè)打通讓訓(xùn)練腳本里用到的算子形狀能自動映射到調(diào)優(yōu)器產(chǎn)出的最佳配置上真正做到“從模型定義到高性能算子”的端到端自動化。這條路還很長但方向是對的。