化實(shí)戰(zhàn):torch.compile、Triton與XLA性能調(diào)優(yōu)指南)
1. 性能瓶頸到底在哪先從一份“看起來(lái)很忙”的Profiling說(shuō)起前陣子幫朋友排查一個(gè)訓(xùn)練任務(wù)A100上GPU利用率看著有八成Loss也在降但總感覺(jué)哪里不對(duì)勁。跑了一輪Profile之后發(fā)現(xiàn)實(shí)際計(jì)算核心Kernel的吞吐遠(yuǎn)沒(méi)有跑滿(mǎn)大量時(shí)間花在了小算子的啟動(dòng)和內(nèi)存拷貝上。這種“表面熱鬧、實(shí)際虛高”的利用率其實(shí)很多搞PyTorch訓(xùn)練的人都遇到過(guò)——模型代碼不是瓶頸PyTorch本身的執(zhí)行機(jī)制才是瓶頸。這也正是我當(dāng)時(shí)決定系統(tǒng)整理這套《AI系統(tǒng)性能工程學(xué)習(xí)筆記》的原因所在而第十四篇的內(nèi)容就是圍繞PyTorch的編譯優(yōu)化這條主線拆開(kāi)講講我在實(shí)際項(xiàng)目中反復(fù)驗(yàn)證過(guò)的三個(gè)關(guān)鍵方向。要理解torch.compile的價(jià)值首先要理解PyTorch默認(rèn)的“逐算子執(zhí)行”模式Eager Mode是怎么回事。簡(jiǎn)單來(lái)說(shuō)你用模型定義寫(xiě)出一層卷積、一層歸一化、一個(gè)ReLUEager模式下PyTorch就會(huì)按順序把每個(gè)算子逐個(gè)交給GPU執(zhí)行。每次執(zhí)行都涉及一次Python層的調(diào)用、一次GPU Kernel的Launch、一次數(shù)據(jù)的搬運(yùn)。小算子和短Kernel特別多的時(shí)候啟動(dòng)開(kāi)銷(xiāo)占比就會(huì)直線上升。哪怕每秒鐘能提交幾萬(wàn)個(gè)KernelGPU真正的計(jì)算單元卻經(jīng)常處于“等任務(wù)”的狀態(tài)。那為什么不把所有算子一次性做完因?yàn)樗阕又g存在數(shù)據(jù)依賴(lài)。卷積的輸出要喂給歸一化歸一化的輸出要喂給ReLU每一步都不能跳過(guò)??扇绻茉诒WC依賴(lài)關(guān)系不被打亂的前提下把多個(gè)連續(xù)算子融合成一個(gè)大的Kernel那啟動(dòng)開(kāi)銷(xiāo)就能被壓到很低內(nèi)存讀寫(xiě)也能少幾個(gè)來(lái)回。這個(gè)思路聽(tīng)著不難但落地起來(lái)步驟非?,嵥椤纫治鲇?jì)算圖又要生成高性能Kernel還要在不同硬件上做適配。PyTorch社區(qū)的答案是torch.compile它把這條鏈路做成了一行代碼就能用起來(lái)的方案。坦白說(shuō)torch.compile剛出來(lái)的時(shí)候我是持觀望態(tài)度的。畢竟PyTorch一直以靈活和動(dòng)態(tài)圖著稱(chēng)硬上編譯優(yōu)化很容易破壞調(diào)試體驗(yàn)。但實(shí)測(cè)了幾批CV和NLP模型之后我的態(tài)度發(fā)生了明顯轉(zhuǎn)變?cè)诓恍枰l繁改圖、動(dòng)態(tài)Shape不嚴(yán)重的訓(xùn)練場(chǎng)景下torch.compile帶來(lái)的提升相當(dāng)直觀尤其是在A100、H100這類(lèi)新架構(gòu)上收益常常能到20%到50%。更關(guān)鍵的是它解決的是“全局優(yōu)化”的問(wèn)題而不是零敲碎打地優(yōu)化某個(gè)算子。從這一篇開(kāi)始我會(huì)把torch.compile的機(jī)制、Triton在其中的作用、以及XLA作為另一條編譯器路線的取舍完整地串起來(lái)講。一個(gè)必須明確的點(diǎn)是torch.compile不是銀彈。它適合那些結(jié)構(gòu)穩(wěn)定、算子類(lèi)型清楚的模型如果模型里到處都是Python控制流、動(dòng)態(tài)Shape、自定義算子編譯優(yōu)化能覆蓋的區(qū)間就會(huì)被壓縮得很厲害。所以這篇文章不只是教你怎么用也會(huì)告訴你什么情況下“不要用它”以及在XLA和torch.compile之間到底該怎么選。2. 拆開(kāi)torch.compile的引擎蓋Dynamo、Graph Break 與 Inductortorch.compile能一行代碼接管優(yōu)化本質(zhì)上是因?yàn)樗鼉?nèi)部有一條流水線式的處理路徑。理解這條路徑比死記API參數(shù)有用得多。我會(huì)從自己調(diào)試過(guò)的實(shí)際案例出發(fā)把它的工作過(guò)程拆成三個(gè)階段講。2.1 Dynamo如何在Python的動(dòng)態(tài)世界里“抓到”計(jì)算圖PyTorch模型是用Python寫(xiě)的Python太靈活了一個(gè)if條件、一個(gè)for循環(huán)、一次字典索引都有可能在運(yùn)行時(shí)改變計(jì)算圖的結(jié)構(gòu)。真要等模型跑完再拿到完整靜態(tài)圖那延遲就太大了。torch.compile在0.x到1.x的迭代中逐步用Dynamo作為前端核心思路是在“保證動(dòng)態(tài)行為正確”的前提下捕獲盡可能大的計(jì)算子圖。Dynamo的做法是追蹤字節(jié)碼。它會(huì)攔截Python函數(shù)的執(zhí)行過(guò)程記錄哪些操作是Tensor計(jì)算哪些操作是純Python邏輯。對(duì)于Tensor計(jì)算Dynamo會(huì)把它轉(zhuǎn)換為計(jì)算圖節(jié)點(diǎn)對(duì)于Python邏輯如果無(wú)法翻譯成圖節(jié)點(diǎn)就標(biāo)志為“Graph Break”。Graph Break之后模型執(zhí)行會(huì)退回Eager模式等跑到下一個(gè)可編譯區(qū)域再重新進(jìn)入優(yōu)化路徑。Graph Break不是報(bào)錯(cuò)但它的多少直接決定了優(yōu)化效果。如果一個(gè)模型里有幾十個(gè)Graph Break那torch.compile幾乎等于沒(méi)優(yōu)化因?yàn)榇蟛糠謺r(shí)間都留在了解釋執(zhí)行狀態(tài)。我在實(shí)戰(zhàn)中遇到過(guò)一個(gè)比較典型的例子模型里對(duì)某個(gè)Tensor做了.item()操作然后根據(jù)這個(gè)標(biāo)量去決定是否執(zhí)行某個(gè)分支。這種寫(xiě)法在調(diào)試時(shí)很自然但Dynamo會(huì)在這里斷掉后面一大段分支都變成Eager路徑。解決辦法也很直接——把.item()移到模型外部或者在損失函數(shù)中避免使用同步操作。你可以在每次編譯后查看torch._dynamo.explain的輸出它能夠列出Graph Break的位置和原因這一招在排查性能問(wèn)題時(shí)極其有效。2.2 Inductor從計(jì)算圖到GPU Kernel的“翻譯官”Dynamo抓到計(jì)算子圖后接下來(lái)的工作就交給后端編譯器。PyTorch默認(rèn)的生成后端是Inductor。它的職責(zé)是把計(jì)算圖翻譯成高性能的GPU Kernel代碼——在NVIDIA GPU上Inductor會(huì)把算子融合和代碼生成的任務(wù)進(jìn)一步交給Triton去搞定。Inductor采用的是基于IR的重寫(xiě)方式。它會(huì)先讀入計(jì)算圖嘗試做元素級(jí)融合、Reduce融合、點(diǎn)乘融合等操作。比如說(shuō)對(duì)一個(gè)Tensor先做x 1再乘2再用tanh激活這三步在Eager模式下至少三四個(gè)Kernel但在Inductor里會(huì)被融合成一個(gè)Triton Kernel一次讀寫(xiě)就完成全部計(jì)算。GPU的內(nèi)存帶寬往往是最大約束少一次全量讀寫(xiě)收益就非常明顯。在torch.compile的配置中backend參數(shù)控制使用哪個(gè)編譯器后端。默認(rèn)是inductor但也可以切換為cudagraphs、tvm等。我實(shí)際用下來(lái)Inductor在NVIDIA卡上的兼容性和性能綜合表現(xiàn)最好。cudagraphs的思路是把一系列Kernel的啟動(dòng)信息錄制下來(lái)然后重復(fù)回放減少CPU端的啟動(dòng)開(kāi)銷(xiāo)但對(duì)算子融合無(wú)能為力。如果你的模型本身就是大算子為主瓶頸不明顯那cudagraphs可能就夠了如果模型是小算子密集型的老老實(shí)實(shí)用Inductor。還有一點(diǎn)容易踩坑torch.compile默認(rèn)會(huì)嘗試動(dòng)態(tài)Shape的支持但開(kāi)啟動(dòng)態(tài)Shape等于放棄了一部分融合優(yōu)化。如果你的輸入尺寸在訓(xùn)練中基本固定可以考慮用dynamicFalse或者把輸入Tensor的尺寸約束住讓Inductor生成更激進(jìn)的專(zhuān)用代碼。我們做離線推理優(yōu)化時(shí)就是這么干的一個(gè)固定尺寸的模型編譯后通常能再壓掉10%左右的延遲。2.3 mode參數(shù)該怎么選default、reduce-overhead還是max-autotunetorch.compile的mode參數(shù)是一個(gè)很容易被忽略但影響很大的選項(xiàng)。官方提供了default、reduce-overhead和max-autotune三檔。default模式下Inductor會(huì)做一些低成本的優(yōu)化編譯時(shí)間短但也意味著放棄了部分更激進(jìn)的改動(dòng)reduce-overhead會(huì)在編譯后引入CUDA Graph錄制對(duì)很多小模型有額外收益max-autotune則會(huì)對(duì)生成的Triton Kernel做大量自動(dòng)調(diào)參性能上限最高但編譯時(shí)間可能長(zhǎng)達(dá)幾分鐘到幾十分鐘。我自己的建議是先用default跑通確認(rèn)無(wú)功能問(wèn)題、無(wú)Graph Break導(dǎo)致的性能回退再?lài)L試reduce-overhead。如果你的模型在多個(gè)批量尺寸上都要用不要貿(mào)然開(kāi)max-autotune因?yàn)閍utotune是針對(duì)固定Shape做的。數(shù)據(jù)增強(qiáng)或動(dòng)態(tài)批量導(dǎo)致Shape頻繁變化時(shí)max-autotune的收益會(huì)被命中率稀釋反而浪費(fèi)了編譯時(shí)間。還有一個(gè)重要技巧在A100或者H100上reduce-overhead通常會(huì)讓小批量訓(xùn)練的速度提升非常明顯因?yàn)樗鼫p少了CPU到GPU之間的同步等待。但在V100這些老卡上CUDA Graph帶來(lái)的收益相對(duì)有限因?yàn)橛布旧淼膯?dòng)延遲沒(méi)那么敏感。環(huán)境不同同樣參數(shù)跑出來(lái)的效果可能完全不一樣這也是為什么我建議任何優(yōu)化都要結(jié)合自己的硬件和模型實(shí)測(cè)而不是照搬網(wǎng)上報(bào)告的數(shù)字。3. Triton 在 torch.compile 里的角色寫(xiě)一次、到處亂跑的高性能內(nèi)核聊到Torch編譯優(yōu)化Triton是一個(gè)繞不開(kāi)的名字。很多剛接觸的人會(huì)把Triton誤解為某種第三方算子庫(kù)其實(shí)它在torch.compile中的作用更底層Inductor生成的代碼很大一部分是Triton語(yǔ)言的Kernel??梢园阉譁\地理解為“GPU上的Python”——用一套類(lèi)Python的語(yǔ)法寫(xiě)出能夠直接在CUDA設(shè)備上高效運(yùn)行的GPU Kernel而不需要手動(dòng)管理線程塊、共享內(nèi)存、同步指令那些復(fù)雜細(xì)節(jié)。3.1 Triton Kernel 到底解決了什么問(wèn)題傳統(tǒng)CUDA編程最難的不是“讓程序跑對(duì)”而是“讓程序跑快”。你要手動(dòng)決定每個(gè)線程負(fù)責(zé)哪個(gè)元素要處理內(nèi)存合并訪問(wèn)要在不同層級(jí)的內(nèi)存之間做搬運(yùn)還要處理Bank Conflict這類(lèi)隱藏很深的性能殺手。寫(xiě)出來(lái)的代碼一旦換了GPU架構(gòu)往往又要重新調(diào)整。Triton的思路是把這些底層細(xì)節(jié)抽象成塊級(jí)操作你只需要描述“每個(gè)塊負(fù)責(zé)計(jì)算什么”而塊內(nèi)部怎么劃分線程、怎么分配寄存器、怎么訪問(wèn)顯存由編譯器自動(dòng)決定。這個(gè)抽象對(duì)自動(dòng)調(diào)優(yōu)特別友好。編譯器可以在一次編譯過(guò)程中生成多個(gè)候選版本然后在真實(shí)硬件上跑一遍選擇最快的那一個(gè)。Inductor在生成Kernel時(shí)就會(huì)調(diào)用Triton做這輪調(diào)優(yōu)。你在日志里看到的“Triton kernel”字樣本質(zhì)上就是Inductor針對(duì)某個(gè)融合子圖生成的GPU代碼。實(shí)際效果上我跑過(guò)一個(gè)典型的ResNet風(fēng)格模型Eager模式下有大概600多個(gè)Kernel執(zhí)行開(kāi)啟torch.compile Inductor Triton后Kernel數(shù)量降到兩百左右端到端訓(xùn)練時(shí)間減少了35%。這個(gè)降幅不是因?yàn)槟骋粋€(gè)算子變快了而是因?yàn)榇罅啃ernel被合并成少量大KernelGPU的執(zhí)行效率因此獲得整體提升。用一句話概括就是Triton不是讓某個(gè)算子從10微秒變成1微秒而是讓幾百次啟動(dòng)從“每次都在浪費(fèi)”變成“每次都真正在計(jì)算”。3.2 Triton 在非編譯場(chǎng)景下的用法手寫(xiě)自定義Kerneltorch.compile之外Triton也可以作為獨(dú)立工具來(lái)寫(xiě)自定義算子。PyTorch里寫(xiě)高性能自定義算子傳統(tǒng)路線是寫(xiě)CUDA C擴(kuò)展然后通過(guò)torch.utils.cpp_extension編譯加載。這條路功能上限高但開(kāi)發(fā)效率低——你得同時(shí)掌握C、CUDA和PyTorch的C接口。Triton提供了一個(gè)折中方案用Python編寫(xiě).triton.kernel裝飾的Kernel函數(shù)運(yùn)行時(shí)自動(dòng)編譯成GPU代碼。調(diào)用的方式跟普通PyTorch函數(shù)一樣傳Tensor進(jìn)去就行。我舉一個(gè)實(shí)際項(xiàng)目里的例子我們要做一個(gè)自定義的注意力掩碼算子標(biāo)準(zhǔn)PyTorch實(shí)現(xiàn)里涉及多次reshape和mask操作顯存開(kāi)銷(xiāo)高。用Triton重寫(xiě)后一個(gè)Kernel內(nèi)完成了mask和softmax的部分融合顯存占用顯著下降速度也快了不少。當(dāng)時(shí)只花了半天時(shí)間就寫(xiě)完了而如果用CUDA C可能要兩三天起步。如果你打算自己動(dòng)手寫(xiě)Triton Kernel建議從簡(jiǎn)單的逐元素算子開(kāi)始練手比如把“ReLU 縮放 偏移”融合成一個(gè)自定義Kernel。把基礎(chǔ)語(yǔ)法、tl.load和tl.store的使用方式搞明白后再?lài)L試更復(fù)雜的Reduce算子。Triton的官方教程里有一個(gè)softmax例子是非常好的入門(mén)材料——同樣一個(gè)功能分別用PyTorch原生實(shí)現(xiàn)和Triton Kernel實(shí)現(xiàn)對(duì)比兩者的時(shí)間消耗你會(huì)很快理解編譯優(yōu)化的核心價(jià)值在哪里。3.3 Triton 版本兼容與安裝坑位Triton目前最常見(jiàn)的安裝方式是隨torch一同安裝。torch.compile在NVIDIA后端會(huì)依賴(lài)Triton所以你在用conda或pip安裝較新版本的PyTorch時(shí)Triton通常已經(jīng)是配套的。但如果你之前手動(dòng)裝過(guò)舊版Triton或者環(huán)境里有多個(gè)PyTorch版本很容易出現(xiàn)“torch版本和Triton版本不匹配”的警告。我在排查環(huán)境時(shí)發(fā)現(xiàn)一個(gè)規(guī)律每當(dāng)PyTorch發(fā)布新版本社區(qū)里就會(huì)出現(xiàn)一批“安裝完torch.compile報(bào)錯(cuò)找不到Triton”或“編譯Kernel報(bào)錯(cuò)版本太舊”的帖子。這時(shí)候首先要做的不是重裝Triton而是確認(rèn)當(dāng)前PyTorch要求的Triton版本范圍。最穩(wěn)妥的方案是直接使用官方推薦的安裝指令pip或conda重新裝一遍PyTorch讓依賴(lài)自動(dòng)拉齊。Google Colab和Kaggle Notebook這類(lèi)云端環(huán)境偶爾也會(huì)出現(xiàn)預(yù)裝Triton版本和torch不匹配的情況重置運(yùn)行時(shí)或升級(jí)torch就可以解決。如果你在CPU-only的機(jī)器上跑torch.compile會(huì)發(fā)現(xiàn)很多Triton相關(guān)功能不可用或者編譯速度極慢。這不是你的代碼有問(wèn)題而是Triton的GPU后端需要CUDA編譯器工具鏈。CPU環(huán)境下Inductor會(huì)嘗試生成C Kernel功能上能跑通但優(yōu)化幅度和GPU場(chǎng)景沒(méi)有可比性。所以建議在動(dòng)手實(shí)踐之前先確認(rèn)自己的環(huán)境有可用的NVIDIA GPU并且PyTorch的CUDA版本和驅(qū)動(dòng)是匹配的。4. XLA 后端另一條編譯器路線的優(yōu)勢(shì)與代價(jià)PyTorch的編譯優(yōu)化不止torch.compile一條路徑。XLAAccelerated Linear Algebra是一個(gè)從TensorFlow生態(tài)里沉淀下來(lái)的編譯器框架也可以作為PyTorch的后端使用。你只需要把模型轉(zhuǎn)換為torch_xla下的執(zhí)行設(shè)備就可以讓計(jì)算圖經(jīng)過(guò)XLA的優(yōu)化后再編譯為對(duì)應(yīng)的硬件指令。這個(gè)方案的好處是跨硬件能力很強(qiáng)——從NVIDIA GPU到Google TPUXLA都有對(duì)應(yīng)的編譯目標(biāo)。4.1 XLA 的工作原理和典型場(chǎng)景XLA的核心思路是“拿到完整的計(jì)算圖再做全局優(yōu)化”。它會(huì)把輸入的計(jì)算圖做算子融合、內(nèi)存規(guī)劃、布局優(yōu)化等處理然后為特定硬件生成可執(zhí)行文件。跟Inductor偏向于“為單個(gè)GPU卡上的小規(guī)模融合”相比XLA的優(yōu)化更傾向于整圖級(jí)別的改寫(xiě)。我最早接觸XLA是在原生TensorFlow時(shí)代那時(shí)候用XLA跑Transformer類(lèi)模型速度提升經(jīng)常是成倍的。后來(lái)PyTorch生態(tài)繁榮起來(lái)torch_xla項(xiàng)目把這種能力搬到了PyTorch里。如果你有TPU資源PyTorch模型可以直接通過(guò)XLA后端在TPU上運(yùn)行這一點(diǎn)是torch.compile目前完全覆蓋不到的。但XLA也有一些明顯的代價(jià)。最大的問(wèn)題是編譯時(shí)間。之前用XLA跑一個(gè)較大規(guī)模的BERT模型編譯階段可能就要幾分鐘到十幾分鐘。如果是訓(xùn)練任務(wù)每個(gè)Step都生成同樣的計(jì)算圖編譯一次就夠了這個(gè)成本可以接受如果是推理任務(wù)每次請(qǐng)求都要處理動(dòng)態(tài)Shape或者新的模型結(jié)構(gòu)編譯開(kāi)銷(xiāo)就會(huì)成為很大的負(fù)擔(dān)。這也是為什么XLA更適用于“固定圖、長(zhǎng)時(shí)間反復(fù)執(zhí)行”的場(chǎng)景。4.2 torch.compile 和 XLA 怎么選很多剛接觸這兩個(gè)概念的人會(huì)糾結(jié)“到底用哪一個(gè)”。我的判斷標(biāo)準(zhǔn)非常簡(jiǎn)單你的目標(biāo)硬件是什么如果你的執(zhí)行環(huán)境是NVIDIA GPU并且模型在PyTorch生態(tài)內(nèi)能跑通默認(rèn)選擇torch.compile。它和PyTorch的接口耦合更緊密調(diào)試信息和工具鏈也更成熟。如果你的目標(biāo)是TPU或者你的模型是從TensorFlow/JAX遷移過(guò)來(lái)的那XLA就是理所當(dāng)然的選擇。另外有些場(chǎng)景會(huì)同時(shí)用到兩者。比如在NVIDIA GPU上做模型開(kāi)發(fā)驗(yàn)證然后跑到TPU上做大規(guī)模訓(xùn)練。我的經(jīng)驗(yàn)是先各自跑通最小實(shí)驗(yàn)再統(tǒng)一對(duì)比指標(biāo)。快速看一張我整理的對(duì)比表對(duì)比維度torch.compile InductorXLA 后端目標(biāo)硬件主要面向NVIDIA GPUNVIDIA GPU / TPU / CPU接入成本一行代碼需要切換設(shè)備為XLA并處理數(shù)據(jù)搬運(yùn)編譯時(shí)間通常幾十秒到幾分鐘大模型可能較長(zhǎng)動(dòng)態(tài)Shape支持較好但允許一定回退一般盡量避免調(diào)試體驗(yàn)較好支持回退Eager模式相對(duì)繁瑣報(bào)錯(cuò)和符號(hào)化程度較高生態(tài)現(xiàn)狀PyTorch官方主推在TensorFlow/JAX場(chǎng)景下更普遍這張表不是絕對(duì)的但它能幫你快速判斷自己該往哪個(gè)方向投入。實(shí)際上我見(jiàn)過(guò)不少團(tuán)隊(duì)為了“追趕熱點(diǎn)”硬上XLA結(jié)果模型在GPU上反而更慢了。因?yàn)閄LA在GPU上的優(yōu)化效果很多時(shí)候并不比Inductor更優(yōu)反而因?yàn)檎麍D編譯的調(diào)度開(kāi)銷(xiāo)在小模型和短任務(wù)上半點(diǎn)便宜都占不到。4.3 XLA 使用中的常見(jiàn)坑位使用torch_xla的時(shí)候最容易踩的坑是數(shù)據(jù)類(lèi)型和Shape不一致帶來(lái)的額外編譯。XLA不喜歡動(dòng)態(tài)Shape。同一個(gè)模型如果每次傳入的sequence長(zhǎng)度都不同XLA就會(huì)反復(fù)做編譯或選擇“動(dòng)態(tài)Shape路徑”性能大打折扣。解決辦法是在數(shù)據(jù)加載階段做好padding把序列長(zhǎng)度統(tǒng)一到同一個(gè)batch內(nèi)的最大值。另一個(gè)坑是分布式訓(xùn)練時(shí)的數(shù)據(jù)同步。torch_xla有自己的分布式接口直接從原生PyTorch的DistributedDataParallel遷移過(guò)來(lái)可能會(huì)碰到collective通信實(shí)現(xiàn)不一致的問(wèn)題。如果只是小規(guī)模實(shí)驗(yàn)單卡訓(xùn)練問(wèn)題不大一旦上多卡建議按照torch_xla官方文檔里的分布式示例調(diào)整代碼而不是硬套原來(lái)的DDP邏輯。最后提一句torch_xla的版本更新頻率通常滯后于PyTorch主版本。如果你正在用很新的PyTorch nightly版裝torch_xla時(shí)最好看一下官方兼容矩陣避免出現(xiàn)API對(duì)不上的問(wèn)題。我自己就遇到過(guò)“小版本不兼容模型A能跑B不能跑”的怪問(wèn)題最終排查出來(lái)是編譯版本不一致導(dǎo)致的。5. 實(shí)操記錄從Eager到torch.compile的一次完整調(diào)優(yōu)這里我會(huì)用一個(gè)簡(jiǎn)化但完整的例子展示我實(shí)際調(diào)優(yōu)一個(gè)CV識(shí)別模型的步驟。這個(gè)模型不復(fù)雜但足以說(shuō)明編譯優(yōu)化的整體流程和注意點(diǎn)。你完全可以把這套流程套用到自己的模型上。5.1 第一步先跑通Baseline拿到可量化指標(biāo)任何優(yōu)化工作第一步永遠(yuǎn)是拿到準(zhǔn)確、可復(fù)現(xiàn)的Benchmark基線。不要上來(lái)就改代碼加編譯否則優(yōu)化前后對(duì)比會(huì)出現(xiàn)很大的噪音。我通常固定隨機(jī)種子、固定輸入Tensor的Shape、固定優(yōu)化器參數(shù)并且把Warmup步數(shù)留足再統(tǒng)計(jì)穩(wěn)定的Step時(shí)間。下面是簡(jiǎn)化示例import torch import torch.nn as nn import time class SimpleModel(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(3, 64, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(64) self.conv2 nn.Conv2d(64, 128, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(128) self.pool nn.AdaptiveAvgPool2d((1, 1)) self.fc nn.Linear(128, 10) def forward(self, x): x torch.relu(self.bn1(self.conv1(x))) x torch.relu(self.bn2(self.conv2(x))) x self.pool(x).flatten(1) x self.fc(x) return x model SimpleModel().cuda().train() optimizer torch.optim.SGD(model.parameters(), lr0.01) x torch.randn(32, 3, 224, 224).cuda() target torch.randint(0, 10, (32,)).cuda() criterion nn.CrossEntropyLoss() # warmup for _ in range(10): optimizer.zero_grad() loss criterion(model(x), target) loss.backward() optimizer.step() # benchmark: 50 iters torch.cuda.synchronize() start time.time() for _ in range(50): optimizer.zero_grad() loss criterion(model(x), target) loss.backward() optimizer.step() torch.cuda.synchronize() avg_step (time.time() - start) / 50 print(fEager avg step: {avg_step * 1000:.2f} ms)跑這類(lèi)腳本時(shí)注意每次循環(huán)都執(zhí)行torch.cuda.synchronize()否則計(jì)時(shí)結(jié)果會(huì)不準(zhǔn)確。GPU執(zhí)行是異步的CPU端的time.time()只能記錄到“提交任務(wù)”的時(shí)間而不是真實(shí)執(zhí)行結(jié)束的時(shí)間。5.2 第二步開(kāi)啟torch.compile并觀察Graph Break在Baseline跑通以后第二步就是一行接入compile并觀察Dynamo的捕獲情況。代碼改動(dòng)很簡(jiǎn)單model torch.compile(SimpleModel().cuda().train(), modereduce-overhead)跑同樣的Benchmark流程。如果代碼沒(méi)有特殊控制流torch.compile一般能直接接管大部分計(jì)算路徑。為了確認(rèn)優(yōu)化覆蓋到了多少比例我強(qiáng)烈建議先跑一次帶有解釋信息的診斷from torch._dynamo import explain compiled_model torch.compile(model, modereduce-overhead) explanation explain(compiled_model, x) print(explanation)你會(huì)看到Graph Break的具體位置、覆蓋算子的比例、以及被優(yōu)化掉的節(jié)點(diǎn)數(shù)。如果Graph Break數(shù)量特別多先不要慌逐條分析原因。常見(jiàn)的原因就那幾類(lèi)調(diào)用了.item()、print到Tensor、使用了不支持的第三方庫(kù)函數(shù)、或者對(duì)Tensor做了Python層面的條件判斷。我當(dāng)時(shí)調(diào)這個(gè)簡(jiǎn)化模型時(shí)最典型的問(wèn)題是在forward里加了幾個(gè)用于調(diào)試的assertDynamo遇到這些Python斷言就Graph Break了。把這些調(diào)試邏輯拿到模型外部之后Graph Break數(shù)量從好幾個(gè)降到零整體性能立刻上了一個(gè)臺(tái)階。5.3 第三步對(duì)照組交叉驗(yàn)證區(qū)分“真實(shí)提升”和“測(cè)量噪音”三步跑完通常得到一張類(lèi)似這樣的表模式平均Step時(shí)間相對(duì)Eager提升Eager18.6 ms基線torch.compile(default)13.8 ms~26%torch.compile(reduce-overhead)12.9 ms~31%torch.compile(max-autotune)12.7 ms~32%這個(gè)例子里reduce-overhead和max-autotune的差距已經(jīng)很小max-autotune的編譯時(shí)間卻可能是后者的幾倍。所以在實(shí)際業(yè)務(wù)里我會(huì)優(yōu)先選擇reduce-overhead因?yàn)樗骖櫫司幾g速度和性能收益。需要提醒的是這類(lèi)數(shù)字受很多因素影響GPU型號(hào)、PyTorch版本、CUDA版本、輸入Shape、Batch Size、是否開(kāi)啟AMP等。同一份代碼在不同機(jī)器上可能跑出完全不同的對(duì)比結(jié)果。我自己有過(guò)一次經(jīng)歷在V100上torch.compile只帶來(lái)5%的提升換到A100上同樣的模型直接提升40%。原因很簡(jiǎn)單——新架構(gòu)對(duì)融合Kernel的執(zhí)行效率更高舊架構(gòu)受限于寄存器、調(diào)度能力收益自然不明顯。5.4 實(shí)戰(zhàn)補(bǔ)遺混合精度和編譯的配合如果你在訓(xùn)練中還開(kāi)啟了AMP自動(dòng)混合精度建議把優(yōu)化順序排成“先AMP再compile”。因?yàn)锳MP本身就能把部分計(jì)算切到FP16/FP16帶來(lái)顯著的顯存和速度收益。此時(shí)再疊加torch.compile還能在融合Kernel里自動(dòng)處理FP16/FP32之間的轉(zhuǎn)換細(xì)節(jié)進(jìn)一步減少讀寫(xiě)開(kāi)銷(xiāo)。有一個(gè)細(xì)節(jié)值得單獨(dú)說(shuō)在Eager模式下AMP會(huì)和torch.cuda.amp.GradScaler配合使用避免梯度下溢。在torch.compile下這個(gè)邏輯本質(zhì)不變但Dynamo會(huì)拿到包含AMP包裝層的計(jì)算圖。正常情況下沒(méi)有問(wèn)題但如果你的模型或損失函數(shù)里有自定義的autograd.Function可能需要額外檢查一下是否兼容編譯路徑。出現(xiàn)問(wèn)題的時(shí)候日志里會(huì)明確提示某個(gè)算子不支持編譯這時(shí)候可以把那個(gè)算子所在的子模塊排除出編譯范圍比如用torch.compile(model, disable...)或者把模塊內(nèi)部改成torch._dynamo.mark_dynamic之類(lèi)的注解去處理。6. 踩坑筆記Graph Break、CUDA Graph 與編譯時(shí)間失控理論聊得再多實(shí)踐中該踩的坑一個(gè)都不會(huì)少。這一節(jié)我把自己在torch.compile、Triton和XLA身上踩過(guò)的高頻坑集中列出來(lái)。每一類(lèi)坑背后都是我曾花掉的幾個(gè)小時(shí)能讓你避開(kāi)的話這篇文章就值回票價(jià)了。6.1 Graph Break導(dǎo)致的“越優(yōu)化越慢”最迷惑人的問(wèn)題就是“編譯之后反而比Eager更慢”。這種情況九成以上都能追溯到Graph Break過(guò)多或編譯覆蓋范圍過(guò)小。一旦模型在關(guān)鍵循環(huán)里頻繁回退到Eager那么每次切換到編譯區(qū)又要額外付出圖捕獲和編譯檢查的代價(jià)結(jié)果就是“反復(fù)橫跳”性能自然惡化。排查的方法特別簡(jiǎn)單用torch._dynamo.explain把優(yōu)化報(bào)告打印出來(lái)里面會(huì)清晰地列出每個(gè)Graph Break的文件名、行號(hào)和原因。我遇到過(guò)一個(gè)項(xiàng)目模型里用了一個(gè)外部庫(kù)的某個(gè)函數(shù)對(duì)Tensor做掩碼處理這個(gè)函數(shù)內(nèi)部包含Python層面的循環(huán)和條件分支Dynamo拿它沒(méi)辦法整段都變成了Graph Break。后來(lái)我用原生PyTorch算子重寫(xiě)了那段邏輯Graph Break立刻消失整體訓(xùn)練時(shí)間縮短了接近一半。實(shí)際操作中還有一個(gè)比較隱蔽的trigger是torch.no_grad()和model.eval()的組合使用。推理階段在關(guān)閉梯度后Dynamo的捕獲反而可能因?yàn)榘b關(guān)系變得更復(fù)雜。如果遇到推理階段Graph Break異??梢栽囋嚢裯o_grad移到調(diào)用compile的模型之前或者將推理函數(shù)整體包進(jìn)一個(gè)torch.no_grad()裝飾的函數(shù)里。6.2 CUDA Graph和動(dòng)態(tài)Shape的沖突reduce-overhead模式會(huì)啟用CUDA Graph技術(shù)。CUDA Graph的精髓是“先錄制后回放”第一步執(zhí)行時(shí)把所有Kernel的調(diào)用信息記錄下來(lái)之后就用極低的CPU開(kāi)銷(xiāo)反復(fù)提交。但錄制意味著Kernel的Shape和輸入Tensor的指針信息基本固定。如果之后某一步傳入的輸入Shape變了CUDA Graph回放了錯(cuò)誤配置就會(huì)導(dǎo)致崩潰、顯存錯(cuò)誤或嚴(yán)重性能下降。我的經(jīng)驗(yàn)是訓(xùn)練中若Batch Size固定、圖像分辨率固定GPU Graph幾乎無(wú)副作用。但你一旦在訓(xùn)練過(guò)程中改變了輸入分辨率比如圖像縮放、動(dòng)態(tài)裁剪或者模型中依賴(lài)數(shù)據(jù)的Shape計(jì)算比如序列長(zhǎng)度變化就必須謹(jǐn)慎。更穩(wěn)妥的方案是把變長(zhǎng)輸入先整理成固定Shape的Batch盡量把動(dòng)態(tài)變化控制在模型外部。如果實(shí)在無(wú)法避免動(dòng)態(tài)Shape那reduce-overhead模式建議先不要用改用default模式跑通驗(yàn)證再考慮進(jìn)一步調(diào)優(yōu)。還有一個(gè)細(xì)節(jié)在使用CUDA Graph時(shí)如果模型內(nèi)部有隨機(jī)性操作比如Dropout錄制后的回放階段會(huì)把隨機(jī)數(shù)生成路徑也一并固定導(dǎo)致結(jié)果不可復(fù)現(xiàn)或者統(tǒng)計(jì)分布偏移。PyTorch的編譯器會(huì)對(duì)這類(lèi)隨機(jī)操作做特殊標(biāo)注但手動(dòng)寫(xiě)的一些自定義Triton Kernel如果內(nèi)部調(diào)用了隨機(jī)API需要格外小心。最好在自定義Kernel里顯式傳入隨機(jī)狀態(tài)或者在模型外部控制隨機(jī)性。6.3 編譯時(shí)間太長(zhǎng)是硬件問(wèn)題還是模型問(wèn)題編譯時(shí)間很多時(shí)候讓人抓狂。max-autotune模式在大模型上動(dòng)輒十幾分鐘、幾十分鐘這在開(kāi)發(fā)迭代階段會(huì)非常影響效率。我的做法是“開(kāi)發(fā)用default上線用max-autotune”開(kāi)發(fā)階段頻繁改代碼沒(méi)必要花大量時(shí)間等待編譯模型定型、進(jìn)入長(zhǎng)期固定訓(xùn)練或反復(fù)推理時(shí)再逐步升級(jí)到更激進(jìn)的編譯選項(xiàng)。另外還有一個(gè)容易忽視的原因Inductor為每個(gè)編譯的圖生成多個(gè)候選Triton Kernel并逐一套用不同參數(shù)運(yùn)行驗(yàn)證這個(gè)“autotune”過(guò)程需要在GPU上真實(shí)跑一遍。如果你同時(shí)開(kāi)了多個(gè)進(jìn)程做編譯GPU資源會(huì)被爭(zhēng)搶autotune的時(shí)間也會(huì)被明顯放大。建議在關(guān)鍵編譯任務(wù)執(zhí)行時(shí)保持同一張卡上只有一個(gè)編譯進(jìn)程。如果你發(fā)現(xiàn)編譯總是卡在某個(gè)特定階段還可以開(kāi)啟TORCH_LOGSinductor環(huán)境變量查看詳細(xì)的編譯日志定位到底是圖優(yōu)化階段耗時(shí)還是Triton Kernel自動(dòng)調(diào)優(yōu)階段耗時(shí)。拿到日志后再?zèng)Q定是換mode還是調(diào)整Shape設(shè)定會(huì)高效得多。6.4 自定義算子無(wú)法編譯三個(gè)層面的應(yīng)對(duì)順序PyTorch生態(tài)里難免有第三方或自制算子。遇到Dynamo不認(rèn)識(shí)的函數(shù)時(shí)第一選擇是改成原生算子組合第二選擇是給這個(gè)算子注冊(cè)一個(gè)“偽量化”的圖捕獲規(guī)則讓編譯期可以忽略其內(nèi)部細(xì)節(jié)只把它當(dāng)作一個(gè)不可拆分的節(jié)點(diǎn)第三選擇才是回到Eager模式把整個(gè)模型或某個(gè)子模塊排除在編譯范圍外。在實(shí)際項(xiàng)目中我見(jiàn)到最多的是很多人一遇到自定義算子不支持就慌張地把torch.compile全局關(guān)閉。如果因?yàn)閹讉€(gè)點(diǎn)損失整張圖的優(yōu)化潛力非??上?。更好的做法是把自定義算子封裝在獨(dú)立的子模塊里用torch.compile只編譯模型中其余穩(wěn)定的部分。PyTorch本身是支持部分模塊編譯的善用這個(gè)能力很多兼容性問(wèn)題都可以繞過(guò)去。7. 更進(jìn)一步在真實(shí)業(yè)務(wù)里如何把編譯優(yōu)化推向生產(chǎn)環(huán)境如果你試過(guò)了torch.compile指標(biāo)也漂亮接下來(lái)要思考的是“怎么讓它在生產(chǎn)環(huán)境穩(wěn)定跑起來(lái)”。這涉及的環(huán)境問(wèn)題比模型代碼本身更多比如服務(wù)化部署時(shí)的請(qǐng)求Shape變化、多卡并行時(shí)的組網(wǎng)方式、以及Cluster里多個(gè)任務(wù)對(duì)GPU的爭(zhēng)搶。7.1 訓(xùn)練場(chǎng)景怎么穩(wěn)定落地訓(xùn)練場(chǎng)景相對(duì)簡(jiǎn)單一些。因?yàn)橛?xùn)練數(shù)據(jù)的Shape通常固定模型結(jié)構(gòu)也穩(wěn)定torch.compile的收益可以平滑落地。需要額外關(guān)注的是多機(jī)多卡的通信開(kāi)銷(xiāo)。當(dāng)你把編譯后的模型接入DDP或FSDP時(shí)通信模式和Kernel執(zhí)行順序可能會(huì)和普通Eager模型略有差異如果出現(xiàn)卡頓或利用率波動(dòng)可以先關(guān)掉編譯模式做一下對(duì)照實(shí)驗(yàn)判斷問(wèn)題是出在編譯優(yōu)化還是通信調(diào)度。FSDP的數(shù)據(jù)流比較復(fù)雜Dynamo在捕獲時(shí)會(huì)看到很多與通信相關(guān)的集合操作。新版本PyTorch已經(jīng)對(duì)這類(lèi)帶通信操作的計(jì)算圖做了更好的適配但如果你用的是較老版本遇到FSDP torch.compile性能不升反降的情況可以嘗試把模型內(nèi)部的通信包裝層移動(dòng)到編譯區(qū)之外或者用官方文檔推薦的版本組合。7.2 推理場(chǎng)景怎么處理動(dòng)態(tài)負(fù)載推理任務(wù)要面對(duì)的最大變量是請(qǐng)求Shape不固定。線上服務(wù)經(jīng)常一個(gè)請(qǐng)求是短文本下一個(gè)就是長(zhǎng)文本還有各種不同Batch Size的準(zhǔn)備策略。如果每個(gè)Shape都觸發(fā)一次torch.compile編譯開(kāi)銷(xiāo)會(huì)完全吃掉性能收益。所以生產(chǎn)落地時(shí)一般會(huì)做“按Shape緩存編譯結(jié)果”的設(shè)計(jì)只對(duì)常見(jiàn)的幾個(gè)Shape組合做預(yù)編譯其他Shape走Eager路徑。我在一個(gè)B端推理服務(wù)里用的策略是在服務(wù)啟動(dòng)階段用幾個(gè)典型Shape預(yù)熱編譯結(jié)果然后用一個(gè)哈希表記錄“Shape簽名 - 編譯后模型”。請(qǐng)求進(jìn)入時(shí)先查表命中就使用編譯版本未命中則走Eager。這樣既保證了服務(wù)質(zhì)量又把編譯成本限制在可接受的范圍內(nèi)。論文里有一個(gè)詞叫“Shape Bucketing”說(shuō)的基本就是這個(gè)思路——把連續(xù)的Shape空間離散化成若干桶以兼容性能和靈活性。7.3 對(duì)“性能工程”這件事本身的思考在我寫(xiě)這套學(xué)習(xí)筆記的整個(gè)過(guò)程中一個(gè)反復(fù)出現(xiàn)的主題是性能優(yōu)化沒(méi)有銀彈。torch.compile很好用但它只是把計(jì)算圖表示、算子調(diào)度和硬件指令生成之間的一層又一層的復(fù)雜性封裝了起來(lái)。真正的性能工程能力不是在某個(gè)模型上跑通一個(gè)API而是能快速定位瓶頸、判斷優(yōu)化手段的適用邊界以及在效率和穩(wěn)定性之間做出合適的權(quán)衡。這需要一種“分層診斷”的思維先看數(shù)據(jù)加載和CPU側(cè)是否飽和再看GPU Util是否真實(shí)有效再看計(jì)算圖里是否存在大量小Kernel最后才輪到是否引入編譯器。整個(gè)順序反了你很可能花了一整天時(shí)間調(diào)編譯參數(shù)最后發(fā)現(xiàn)瓶頸在DataLoader根本沒(méi)喂飽GPU。8. 番外幾個(gè)你一定會(huì)用到的環(huán)境與版本問(wèn)題因?yàn)檫@套筆記發(fā)布后經(jīng)常有讀者私信問(wèn)我環(huán)境配置的問(wèn)題這里把跟torch.compile、Triton、XLA相關(guān)的環(huán)境適配問(wèn)題單獨(dú)拿出來(lái)說(shuō)一遍免得大家在前面代碼沒(méi)跑起來(lái)的時(shí)候就被環(huán)境卡住。8.1 PyTorch與CUDA的版本匹配不同版本的PyTorch對(duì)CUDA的支持版本不同。我個(gè)人的習(xí)慣是直接去PyTorch官網(wǎng)的Get Started頁(yè)面選擇合適的操作系統(tǒng)、包管理工具和CUDA版本然后復(fù)制對(duì)應(yīng)的安裝命令。不要自己手動(dòng)從一堆索引里拼輪子那樣很容易引入版本沖突。如果你是離線環(huán)境需要先在一臺(tái)聯(lián)網(wǎng)的機(jī)器上把對(duì)應(yīng)的wheel包下載好然后再搬到目標(biāo)機(jī)器安裝。下載時(shí)注意選擇cu118、cu121這類(lèi)后綴確保和機(jī)器上實(shí)際安裝的驅(qū)動(dòng)兼容。驅(qū)動(dòng)本身要符合CUDA運(yùn)行時(shí)的最低版本要求否則即使torch裝好了torch.cuda.is_available()也可能返回False。8.2 如何確認(rèn)Triton已經(jīng)正確安裝和可用最簡(jiǎn)單的方式是在Python環(huán)境里執(zhí)行以下代碼import triton print(triton.__version__)如果打印出版本號(hào)基本就說(shuō)明裝好了。接下來(lái)再跑一個(gè)小Kernel驗(yàn)證確保運(yùn)行路徑也正確。比如官方文檔里的向量加法示例或者一個(gè)最簡(jiǎn)單的torch.compile測(cè)試。只import成功不一定代表能正常編譯Kernel因?yàn)門(mén)riton還依賴(lài)本機(jī)的編譯器工具鏈。在容器化環(huán)境中經(jīng)常出現(xiàn)“宿主機(jī)能跑但容器里不能跑”的情況絕大多數(shù)是因?yàn)槿萜麋R像里缺少libgomp等動(dòng)態(tài)庫(kù)或者LD_LIBRARY_PATH沒(méi)有正確包含CUDA的庫(kù)路徑。遇到這類(lèi)問(wèn)題時(shí)先檢查基礎(chǔ)鏡像是否包含完整的CUDA runtime再檢查nvcc是否可用。8.3 我常用的一個(gè)快速驗(yàn)證腳本最后分享一個(gè)我?guī)缀趺看芜w移環(huán)境后都會(huì)跑的快速驗(yàn)證腳本內(nèi)容很簡(jiǎn)單但它能在十分鐘內(nèi)暴露80%的常見(jiàn)環(huán)境問(wèn)題import torch print(PyTorch version:, torch.__version__) print(CUDA available:, torch.cuda.is_available()) if torch.cuda.is_available(): print(CUDA version:, torch.version.cuda) print(GPU name:, torch.cuda.get_device_name(0)) x torch.randn(1000, 1000, devicecuda) y torch.matmul(x, x) print(matmul ok:, y.shape) try: import triton print(Triton version:, triton.__version__) except ImportError as e: print(Triton import error:, e) def simple_add(a, b): return a b compiled_add torch.compile(simple_add) out compiled_add(x, x) print(torch.compile ok:, out.shape)如果這個(gè)腳本全綠那說(shuō)明環(huán)境基本可用如果哪一步紅了就跟著堆棧信息去查對(duì)應(yīng)依賴(lài)。每次跑到全綠我才會(huì)開(kāi)始正式的模型優(yōu)化工作。這套流程幫我省去了無(wú)數(shù)次“改了模型卻不知道是環(huán)境問(wèn)題還是代碼問(wèn)題”的無(wú)謂排查。這篇筆記從torch.compile的內(nèi)部機(jī)制講到Triton如何生成GPU Kernel再對(duì)比了XLA這條跨硬件編譯器路線最后落到工程環(huán)境里的落地實(shí)踐。這些內(nèi)容都是我實(shí)際跑模型、調(diào)GPU、被各種版本和Graph Break折騰完之后沉淀出來(lái)的。按我自己的經(jīng)驗(yàn)真正理解這三條主線之后你再去看PyTorch相關(guān)的性能優(yōu)化問(wèn)題視角會(huì)和以前明顯不一樣——不再只是調(diào)幾個(gè)參數(shù)看數(shù)字而是能判斷瓶頸在哪一層、哪條優(yōu)化路徑最適合當(dāng)前場(chǎng)景、換了硬件后又會(huì)發(fā)生什么變化。這一篇先寫(xiě)到這第十五篇打算展開(kāi)聊聊分布式訓(xùn)練里通信和計(jì)算的流水線重疊那個(gè)方向也是我們生產(chǎn)環(huán)境里吃了不少苦頭才摸清門(mén)道的話題。