PyTorch理事會(huì)席位背后:AI芯片軟件棧適配與算子實(shí)現(xiàn)全解析)
1. 從“同桌”這個(gè)詞說起一個(gè)信號(hào)背后的技術(shù)分量“寒武紀(jì)拿下PyTorch最高席位與英偉達(dá)同桌”——這個(gè)標(biāo)題我第一次看到的時(shí)候正在調(diào)一個(gè)模型訓(xùn)練腳本手邊跑著的是一臺(tái)裝了消費(fèi)級(jí)顯卡的機(jī)器。說實(shí)話第一反應(yīng)不是興奮而是“終于”。因?yàn)樽錾疃葘W(xué)習(xí)框架適配這行的人都知道PyTorch的技術(shù)治理結(jié)構(gòu)里能進(jìn)入核心決策層的企業(yè)從來不只是“貢獻(xiàn)了幾行代碼”那么簡單。PyTorch基金會(huì)PyTorch Foundation的治理架構(gòu)是2022年從Meta手里剝離出來、交給Linux基金會(huì)托管之后逐步成型的。它的理事會(huì)Governing Board和技術(shù)咨詢委員會(huì)Technical Advisory Council席位基本代表了全球深度學(xué)習(xí)框架生態(tài)里最有話語權(quán)的一批玩家。英偉達(dá)、AMD、Meta、Google、微軟、亞馬遜這些名字長期占據(jù)核心位置原因很直接它們要么是框架的主要貢獻(xiàn)者要么是硬件后端的主要實(shí)現(xiàn)方要么是最大規(guī)模的使用方。寒武紀(jì)能拿到這個(gè)席位意味著它在PyTorch生態(tài)里的角色從“下游適配者”變成了“上游共建者”。這個(gè)轉(zhuǎn)變的技術(shù)含量比很多人想象的要高得多。我見過太多人把“支持PyTorch”理解成“能跑起來就行”但真正做過框架后端適配的工程師都清楚從“能跑”到“進(jìn)治理層”中間隔著的是對(duì)框架核心抽象的理解深度、對(duì)算子語義的精確實(shí)現(xiàn)、以及對(duì)整個(gè)編譯棧的持續(xù)投入。這篇文章我想聊的不是新聞本身而是這個(gè)信號(hào)背后一個(gè)AI芯片公司的軟件棧到底要具備什么樣的能力才能走到這個(gè)位置。同時(shí)我也會(huì)把PyTorch環(huán)境搭建、芯片適配、算子實(shí)現(xiàn)這些實(shí)操層面的東西拆開講清楚讓不管是剛?cè)腴T的新手還是正在做國產(chǎn)硬件適配的同行都能從中拿到能直接用的東西。2. PyTorch生態(tài)的“席位”到底意味著什么2.1 基金會(huì)治理結(jié)構(gòu)里的技術(shù)話語權(quán)PyTorch基金會(huì)的理事會(huì)成員通常來自幾個(gè)類別創(chuàng)始成員Meta、AMD、AWS、Google、Microsoft、NVIDIA等、一般成員、以及關(guān)聯(lián)成員。每個(gè)席位背后對(duì)應(yīng)的不只是資金贊助更重要的是技術(shù)方向的投票權(quán)和標(biāo)準(zhǔn)制定參與權(quán)。技術(shù)咨詢委員會(huì)TAC的職責(zé)更偏工程側(cè)負(fù)責(zé)審批新的后端接入方案、評(píng)審核心API的變更提案、協(xié)調(diào)跨廠商的算子語義一致性。一個(gè)硬件廠商如果只是“能跑PyTorch”它只需要維護(hù)一個(gè)torch.compile的后端或者一個(gè)PrivateUse1的設(shè)備擴(kuò)展就行。但進(jìn)入TAC或理事會(huì)意味著你要參與決定“下一個(gè)版本的PyTorch設(shè)備抽象層應(yīng)該怎么改”“新的算子注冊機(jī)制要不要兼容舊后端”這類問題。我舉個(gè)具體的例子。PyTorch 2.x引入的torch.compile和TorchInductor對(duì)后端硬件提出了全新的要求。以前你只要實(shí)現(xiàn)一套ATen算子就能跑Eager模式現(xiàn)在還要考慮Dynamo的圖捕獲、Inductor的代碼生成、以及不同后端之間的調(diào)度策略。如果一個(gè)芯片廠商在TAC里有席位它就能在Inductor的后端接口設(shè)計(jì)階段就提出意見而不是等接口凍結(jié)了再去逆向適配。2.2 從“適配”到“共建”的技術(shù)門檻我接觸過不少做國產(chǎn)芯片軟件棧的團(tuán)隊(duì)大家普遍有一個(gè)誤區(qū)覺得把PyTorch的算子庫對(duì)著文檔實(shí)現(xiàn)一遍跑通ResNet和BERT就算“支持PyTorch”了。這個(gè)標(biāo)準(zhǔn)在2019年可能還說得過去放到今天遠(yuǎn)遠(yuǎn)不夠。現(xiàn)在的PyTorch生態(tài)一個(gè)合格的硬件后端至少要覆蓋這幾層ATen算子層這是最基礎(chǔ)的幾百個(gè)算子的語義要跟CUDA后端對(duì)齊包括各種邊界條件、數(shù)據(jù)類型提升規(guī)則、廣播語義。調(diào)度與內(nèi)存管理層PyTorch的CUDACachingAllocator那套內(nèi)存池機(jī)制在非CUDA設(shè)備上怎么復(fù)現(xiàn)直接影響到訓(xùn)練時(shí)的顯存利用率和碎片化程度。圖編譯層Dynamo捕獲的FX圖要能被你的后端正確消費(fèi)。Inductor生成Triton代碼的那套流程如果你的硬件不支持Triton就得自己實(shí)現(xiàn)一套等價(jià)的代碼生成路徑。分布式訓(xùn)練層NCCL在CUDA生態(tài)里的地位不用多說非CUDA設(shè)備要接入PyTorch的分布式接口要么實(shí)現(xiàn)一套兼容NCCL API的通信庫要么走Gloo的路線但性能會(huì)打折扣。量化與推理優(yōu)化層INT8、FP8這些低精度格式的支持以及和torch.ao量化工具鏈的對(duì)接。寒武紀(jì)能進(jìn)理事會(huì)說明它在這些層面至少拿出了讓社區(qū)認(rèn)可的方案。具體是哪些層面公開信息里沒有完整披露但從它之前開源的torch_mluCambricon PyTorch擴(kuò)展和catch寒武紀(jì)的PyTorch后端來看ATen算子覆蓋和調(diào)度層是下了功夫的。2.3 對(duì)普通開發(fā)者的實(shí)際影響你可能會(huì)問一個(gè)芯片公司進(jìn)理事會(huì)跟我一個(gè)天天寫nn.Module的人有什么關(guān)系關(guān)系很直接。最明顯的一點(diǎn)是以后你在PyTorch里用torch.mlu或者類似的設(shè)備抽象時(shí)它的行為會(huì)更接近torch.cuda。以前國產(chǎn)芯片的PyTorch適配經(jīng)常出現(xiàn)“這個(gè)算子不支持”“那個(gè)dtype報(bào)錯(cuò)”的情況很大程度上是因?yàn)檫m配方?jīng)]有參與上游設(shè)計(jì)只能被動(dòng)跟著CUDA后端的實(shí)現(xiàn)走。有了席位之后設(shè)備抽象層的接口設(shè)計(jì)會(huì)更考慮多后端的通用性而不是默認(rèn)以CUDA為唯一參考。另一個(gè)影響是文檔和教程的覆蓋。PyTorch官方教程里如果開始出現(xiàn)非CUDA設(shè)備的示例對(duì)新手來說學(xué)習(xí)成本會(huì)低很多。我見過太多人第一次接觸國產(chǎn)AI芯片時(shí)卡在“怎么把.cuda()換成對(duì)應(yīng)的設(shè)備調(diào)用”這一步上。3. 芯片適配PyTorch的完整技術(shù)路徑拆解3.1 設(shè)備抽象層PrivateUse1機(jī)制怎么用PyTorch從1.13開始正式引入了PrivateUse1這個(gè)設(shè)備類型專門給非CUDA、非ROCm的第三方硬件用。它的設(shè)計(jì)初衷就是讓芯片廠商不用改PyTorch核心代碼通過擴(kuò)展機(jī)制就能注冊自己的設(shè)備。具體來說你需要做這幾件事# 注冊設(shè)備名稱 torch.utils.rename_privateuse1_backend(mlu) # 注冊設(shè)備模塊 torch._register_device_module(mlu, MLUModule) # 生成對(duì)應(yīng)的設(shè)備類型 torch.utils.generate_methods_for_privateuse1_backend()這三行代碼執(zhí)行完之后你就可以像用torch.cuda一樣用torch.mlu了。tensor.mlu()、torch.mlu.current_device()、torch.mlu.synchronize()這些方法都會(huì)自動(dòng)生成。但這里有個(gè)坑rename_privateuse1_backend必須在任何張量創(chuàng)建之前調(diào)用而且一個(gè)進(jìn)程里只能調(diào)用一次。我見過有人在Jupyter Notebook里反復(fù)執(zhí)行注冊代碼結(jié)果第二次就報(bào)錯(cuò)。正確的做法是把它放在包的__init__.py里或者用一個(gè)單獨(dú)的初始化模塊來管理。3.2 算子注冊從ATen到你的硬件PyTorch的算子注冊機(jī)制核心是TORCH_LIBRARY和TORCH_LIBRARY_IMPL這兩個(gè)宏。對(duì)于第三方后端你需要為每個(gè)算子實(shí)現(xiàn)對(duì)應(yīng)的kernel然后注冊到你的設(shè)備類型上。// 以add算子為例 TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) { m.impl(add.Tensor, TORCH_FN(mlu_add_tensor)); m.impl(add.Scalar, TORCH_FN(mlu_add_scalar)); m.impl(add.out, TORCH_FN(mlu_add_out)); }這里的關(guān)鍵是算子變體。PyTorch里一個(gè)add操作可能有十幾個(gè)變體add.Tensor、add.Scalar、add.out、add.Scalar_out、add_.Tensor原地操作等等。你如果只實(shí)現(xiàn)了add.Tensor那用戶寫torch.add(a, b, outc)的時(shí)候就會(huì)報(bào)“未實(shí)現(xiàn)”的錯(cuò)誤。我的經(jīng)驗(yàn)是先把PyTorch的native_functions.yaml里所有標(biāo)記為CompositeExplicitAutograd的算子過一遍這些是可以通過組合其他算子實(shí)現(xiàn)的優(yōu)先級(jí)可以放低。真正要優(yōu)先實(shí)現(xiàn)的是CompositeImplicitAutograd和那些直接對(duì)應(yīng)硬件指令的算子。3.3 內(nèi)存管理別小看CachingAllocatorCUDA生態(tài)里CUDACachingAllocator是PyTorch顯存管理的核心。它通過緩存已分配的內(nèi)存塊避免頻繁調(diào)用cudaMalloc和cudaFree帶來的性能開銷。在非CUDA設(shè)備上如果你直接用malloc和free訓(xùn)練速度可能會(huì)掉30%以上。寒武紀(jì)的torch_mlu里實(shí)現(xiàn)了一套MLUCachingAllocator基本思路和CUDA版本一致維護(hù)一個(gè)按大小分桶的空閑塊列表分配時(shí)優(yōu)先從緩存里找合適大小的塊找不到再向驅(qū)動(dòng)申請。釋放時(shí)不立即歸還給驅(qū)動(dòng)而是放回緩存。這里有個(gè)細(xì)節(jié)值得注意內(nèi)存池的大小和碎片化策略。CUDA的allocator默認(rèn)會(huì)保留所有釋放的塊直到進(jìn)程結(jié)束。如果你的設(shè)備顯存比較小比如推理卡只有16GB可能需要設(shè)置一個(gè)上限超過之后主動(dòng)釋放一些塊。PyTorch提供了torch.cuda.memory._set_allocator_settings這樣的接口第三方后端也可以實(shí)現(xiàn)類似的配置項(xiàng)。3.4 圖編譯與Inductor后端對(duì)接PyTorch 2.x之后torch.compile成了性能優(yōu)化的主要入口。它的工作流程是Dynamo捕獲Python字節(jié)碼生成FX圖AOTAutograd做前向和反向的圖分解Inductor把FX圖 lowering 成Triton代碼或者C代碼。對(duì)于非CUDA設(shè)備你有兩個(gè)選擇實(shí)現(xiàn)一個(gè)Inductor后端繼承torch._inductor.codegen.common.CodeGen實(shí)現(xiàn)自己的調(diào)度和代碼生成邏輯。這條路工作量大但性能上限高。走Triton兼容路線如果你的硬件能跑Triton生成的代碼或者你能把Triton IR翻譯成自己的指令那就可以復(fù)用Inductor的大部分流程。寒武紀(jì)走的是哪條路公開資料里沒有明確說。但從它之前發(fā)布的torch_mlu更新日志來看torch.compile的支持是逐步推進(jìn)的早期版本需要設(shè)置torch._dynamo.config.suppress_errors True來跳過不支持的圖。4. 實(shí)操從零搭建PyTorch環(huán)境并驗(yàn)證芯片適配4.1 環(huán)境準(zhǔn)備Anaconda與Python版本選擇不管你用的是CUDA設(shè)備還是國產(chǎn)芯片Anaconda都是管理Python環(huán)境最省心的方式。我個(gè)人的習(xí)慣是每個(gè)項(xiàng)目一個(gè)獨(dú)立環(huán)境避免依賴沖突。# 創(chuàng)建環(huán)境Python版本建議3.9或3.10 conda create -n pytorch_mlu python3.10 conda activate pytorch_mlu # 安裝PyTorch基礎(chǔ)包 # 注意如果你的芯片廠商提供了定制版PyTorch要用他們的源 pip install torch torchvision torchaudio這里有個(gè)關(guān)鍵點(diǎn)PyTorch版本和芯片驅(qū)動(dòng)版本的匹配。CUDA生態(tài)里PyTorch 2.0需要CUDA 11.7或11.8PyTorch 2.1開始支持CUDA 12.1。國產(chǎn)芯片也有類似的版本對(duì)應(yīng)關(guān)系裝之前一定要看廠商的release note。我踩過的一個(gè)坑是用conda裝PyTorch時(shí)conda會(huì)自動(dòng)裝一個(gè)它認(rèn)為兼容的CUDA runtime但這個(gè)runtime可能和你系統(tǒng)里的驅(qū)動(dòng)版本不匹配。后來我改成用pip裝并且明確指定--index-url指向廠商的包源問題就少了。4.2 驗(yàn)證設(shè)備可用性與基本算子環(huán)境裝好之后第一件事是驗(yàn)證設(shè)備能不能被PyTorch識(shí)別import torch # 檢查設(shè)備是否可用 print(torch.mlu.is_available()) # 如果是寒武紀(jì) print(torch.mlu.device_count()) print(torch.mlu.get_device_name(0)) # 創(chuàng)建一個(gè)張量并移動(dòng)到設(shè)備上 x torch.randn(3, 3) x_mlu x.mlu() print(x_mlu.device) # 跑一個(gè)簡單的矩陣乘法 a torch.randn(1024, 1024).mlu() b torch.randn(1024, 1024).mlu() c torch.mm(a, b) print(c.sum())如果這幾步都能跑通說明基礎(chǔ)的算子注冊和內(nèi)存管理沒問題。接下來要測的是算子覆蓋度。我的做法是拿一個(gè)真實(shí)的模型比如ResNet-50跑一遍前向和反向看哪些算子會(huì)報(bào)“未實(shí)現(xiàn)”。import torchvision.models as models model models.resnet50().mlu() x torch.randn(32, 3, 224, 224).mlu() y model(x) loss y.sum() loss.backward() print(ResNet-50 forward/backward OK)如果這一步報(bào)錯(cuò)錯(cuò)誤信息通常會(huì)告訴你缺哪個(gè)算子。比如aten::adaptive_avg_pool2d沒實(shí)現(xiàn)你就需要去補(bǔ)這個(gè)算子的kernel。4.3 性能對(duì)比別只看“能跑”“能跑”和“跑得快”是兩回事。我見過一些適配方案功能測試全過但訓(xùn)練速度只有CUDA版本的十分之一。問題通常出在幾個(gè)地方算子實(shí)現(xiàn)沒有用上硬件的向量化指令比如矩陣乘法如果只是用for循環(huán)在CPU上算完再拷貝回設(shè)備那速度肯定不行。內(nèi)存拷貝太頻繁每次算子調(diào)用都做一次host-device同步會(huì)把流水線打斷。沒有做算子融合PyTorch Eager模式下conv bn relu是三個(gè)獨(dú)立的kernel調(diào)用。CUDA生態(tài)里有cuDNN做融合非CUDA設(shè)備如果沒做類似的優(yōu)化性能差距會(huì)很大。我一般會(huì)用torch.profiler來看每個(gè)算子的耗時(shí)with torch.profiler.profile( activities[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.MLU], scheduletorch.profiler.schedule(wait1, warmup1, active3), on_trace_readytorch.profiler.tensorboard_trace_handler(./log) ) as prof: for step, data in enumerate(dataloader): if step 5: break train_step(model, data) prof.step()看trace的時(shí)候重點(diǎn)關(guān)注兩件事一是設(shè)備上的kernel執(zhí)行時(shí)間占比二是host和device之間的同步點(diǎn)。如果同步點(diǎn)太多說明你的后端在算子調(diào)度上還有優(yōu)化空間。4.4 分布式訓(xùn)練接入的注意事項(xiàng)單卡跑通之后下一步通常是多卡分布式。PyTorch的分布式接口主要有兩種DistributedDataParallelDDP和FullyShardedDataParallelFSDP。它們底層都依賴通信庫來做梯度同步。CUDA生態(tài)里用NCCL非CUDA設(shè)備要么實(shí)現(xiàn)一套兼容NCCL API的通信庫要么用Gloo。Gloo的問題是它主要針對(duì)CPU優(yōu)化在設(shè)備間通信時(shí)性能損失比較大。我實(shí)測過一個(gè)4卡訓(xùn)練的任務(wù)用Gloo的吞吐量只有NCCL的60%左右。如果你的芯片廠商提供了自己的通信庫接入方式通常是實(shí)現(xiàn)torch.distributed.ProcessGroup的子類然后通過init_process_group的backend參數(shù)指定。這里要注意的是通信和計(jì)算的重疊。DDP默認(rèn)會(huì)在反向傳播的同時(shí)做梯度allreduce如果你的通信庫不支持異步操作這個(gè)重疊就做不起來訓(xùn)練速度會(huì)明顯下降。5. 常見問題與排查技巧實(shí)錄5.1 算子未實(shí)現(xiàn)報(bào)錯(cuò)怎么定位最常見的報(bào)錯(cuò)長這樣RuntimeError: Could not run aten::xxx with arguments from the MLU backend.排查步驟確認(rèn)這個(gè)算子在CUDA后端有沒有實(shí)現(xiàn)。如果CUDA也沒有那可能是PyTorch版本的問題。檢查你的TORCH_LIBRARY_IMPL注冊代碼看算子名和變體是否寫對(duì)了。PyTorch的算子名是大小寫敏感的add.Tensor和add.tensor不一樣。如果算子是通過組合實(shí)現(xiàn)的Composite檢查依賴的子算子是否都已實(shí)現(xiàn)。我整理了一個(gè)速查表報(bào)錯(cuò)信息可能原因解決方法Could not run aten::xxx算子未注冊實(shí)現(xiàn)并注冊對(duì)應(yīng)kernelExpected all tensors to be on the same device張量設(shè)備不一致檢查.to(device)調(diào)用MLU out of memory顯存不足減小batch size或優(yōu)化內(nèi)存池NCCL error/Gloo error通信庫配置問題檢查環(huán)境變量和網(wǎng)絡(luò)配置dtype not supported數(shù)據(jù)類型不支持轉(zhuǎn)換到支持的dtype或?qū)崿F(xiàn)該dtype的kernel5.2 環(huán)境配置的坑驅(qū)動(dòng)、CUDA、PyTorch三者關(guān)系雖然這里聊的是國產(chǎn)芯片但很多人是在CUDA環(huán)境里做開發(fā)然后遷移到國產(chǎn)芯片上。CUDA環(huán)境本身就有不少坑我順帶說一下。驅(qū)動(dòng)版本和CUDA版本的關(guān)系NVIDIA驅(qū)動(dòng)是向下兼容CUDA的但有一個(gè)最低版本要求。比如CUDA 12.1需要驅(qū)動(dòng)版本530。你可以用nvidia-smi看驅(qū)動(dòng)版本用nvcc --version看CUDA版本。如果nvcc顯示的版本和PyTorch編譯時(shí)用的CUDA版本不一致可能會(huì)出現(xiàn)運(yùn)行時(shí)錯(cuò)誤。PyTorch和CUDA的對(duì)應(yīng)關(guān)系PyTorch官網(wǎng)的安裝命令里會(huì)明確寫cu118、cu121這樣的后綴。如果你用conda install pytorch而不指定conda可能會(huì)裝一個(gè)CPU版本或者裝一個(gè)和你驅(qū)動(dòng)不匹配的CUDA版本。Anaconda環(huán)境隔離我強(qiáng)烈建議用conda創(chuàng)建獨(dú)立環(huán)境不要在base環(huán)境里裝PyTorch。因?yàn)椴煌?xiàng)目可能依賴不同版本的PyTorch混在一起遲早出問題。5.3 性能調(diào)優(yōu)的獨(dú)家經(jīng)驗(yàn)做了幾年框架適配我總結(jié)了幾條性能調(diào)優(yōu)的經(jīng)驗(yàn)有些是文檔里不會(huì)寫的第一條先看數(shù)據(jù)加載再看模型計(jì)算。很多人一上來就優(yōu)化算子結(jié)果發(fā)現(xiàn)瓶頸在DataLoader上。用torch.utils.data.DataLoader的時(shí)候num_workers設(shè)成CPU核心數(shù)的一半到三分之二比較合適pin_memoryTrue在CUDA環(huán)境下能加速host到device的拷貝但在非CUDA設(shè)備上不一定有效要實(shí)測。第二條batch size不是越大越好。大batch能提高硬件利用率但也會(huì)增加顯存壓力和通信開銷。我一般會(huì)做一個(gè)batch size的掃描從16開始翻倍看吞吐量的變化曲線找到拐點(diǎn)。第三條混合精度訓(xùn)練要謹(jǐn)慎。AMP自動(dòng)混合精度在CUDA上很成熟但在非CUDA設(shè)備上FP16的算子覆蓋度可能不夠。如果發(fā)現(xiàn)loss變成NaN先檢查是不是某個(gè)算子在FP16下溢出了。第四條算子融合是最大的性能杠桿。如果你們的芯片支持自定義算子融合一定要把convbnrelu、lineargelu這些常見pattern做進(jìn)去。我見過一個(gè)案例光是融合了這幾個(gè)pattern訓(xùn)練速度就提升了40%。5.4 從CUDA遷移到國產(chǎn)芯片的代碼改動(dòng)清單如果你有一個(gè)現(xiàn)成的CUDA項(xiàng)目要遷移到國產(chǎn)芯片需要改的地方其實(shí)不多但每一處都要仔細(xì)# 1. 設(shè)備指定 # 原來 device torch.device(cuda:0) # 改成 device torch.device(mlu:0) # 或廠商指定的設(shè)備名 # 2. 張量移動(dòng) # 原來 x x.cuda() # 改成 x x.mlu() # 3. 分布式后端 # 原來 dist.init_process_group(backendnccl) # 改成 dist.init_process_group(backendcncl) # 或廠商提供的后端名 # 4. 隨機(jī)種子 # 原來 torch.cuda.manual_seed(42) # 改成 torch.mlu.manual_seed(42) # 5. 性能分析 # 原來 with torch.profiler.profile(activities[torch.profiler.ProfilerActivity.CUDA]): # 改成 with torch.profiler.profile(activities[torch.profiler.ProfilerActivity.MLU]):看起來簡單但實(shí)際遷移時(shí)最容易出問題的是第三方庫的依賴。比如apex、deepspeed、flash-attention這些庫它們內(nèi)部有大量CUDA-specific的代碼。如果廠商沒有提供對(duì)應(yīng)的移植版本你可能需要自己改。6. 這件事對(duì)行業(yè)意味著什么6.1 多后端生態(tài)的必然趨勢PyTorch基金會(huì)接納寒武紀(jì)本質(zhì)上反映了一個(gè)趨勢深度學(xué)習(xí)框架正在從“CUDA中心化”走向“多后端并行”。這個(gè)趨勢不是PyTorch一家的事TensorFlow有tf.device的插件機(jī)制JAX有PJRTPortable JAX Runtime大家都在做類似的事情。對(duì)開發(fā)者來說這意味著以后寫代碼時(shí)設(shè)備相關(guān)的部分會(huì)越來越抽象。你可能不需要寫x.cuda()而是寫x.to(device)然后通過配置來決定用哪個(gè)后端。這對(duì)代碼的可移植性是好事但也要求你對(duì)不同后端的特性有基本了解不然性能調(diào)優(yōu)會(huì)無從下手。6.2 國產(chǎn)芯片軟件棧的短板與機(jī)會(huì)說實(shí)話國產(chǎn)AI芯片在硬件參數(shù)上追得很快但在軟件棧上普遍落后。這個(gè)落后不是“能不能跑”的問題而是“好不好用”的問題。具體表現(xiàn)在文檔質(zhì)量參差不齊很多廠商的文檔只告訴你“怎么裝”不告訴你“為什么這么裝”出了問題只能提工單。社區(qū)支持薄弱CUDA生態(tài)里有Stack Overflow、有GitHub上成千上萬的issue國產(chǎn)芯片的社區(qū)還在建設(shè)中。工具鏈不完整性能分析工具、調(diào)試工具、可視化工具這些CUDA生態(tài)里習(xí)以為常的東西在國產(chǎn)芯片上往往缺失。寒武紀(jì)進(jìn)PyTorch理事會(huì)至少說明它在軟件棧上的投入得到了社區(qū)認(rèn)可。這對(duì)整個(gè)國產(chǎn)芯片行業(yè)是一個(gè)正向信號(hào)軟件生態(tài)的建設(shè)開始被放到和硬件同等重要的位置。6.3 給開發(fā)者的建議現(xiàn)在該做什么如果你是一個(gè)深度學(xué)習(xí)開發(fā)者不管你現(xiàn)在用的是CUDA還是國產(chǎn)芯片我有幾個(gè)建議第一不要把設(shè)備相關(guān)的代碼寫死。用device torch.device(...)這樣的方式而不是到處寫.cuda()。這樣以后遷移的時(shí)候改一個(gè)地方就行。第二關(guān)注PyTorch的RFCRequest for Comments。PyTorch的重大變更都會(huì)先發(fā)RFC比如PrivateUse1機(jī)制、torch.compile的后端接口都是在RFC階段就公開討論的。提前了解這些能讓你在適配時(shí)少走彎路。第三動(dòng)手試。如果你手邊有國產(chǎn)芯片的開發(fā)板或者云上的實(shí)例花一個(gè)下午把PyTorch環(huán)境搭起來跑一個(gè)簡單的模型。很多問題只有親手做了才會(huì)遇到看文檔是看不出來的。第四參與社區(qū)。PyTorch的GitHub issue和論壇里關(guān)于非CUDA后端的討論越來越多。你遇到的問題很可能別人也遇到過。把你的解決方案分享出來既幫了別人也讓自己對(duì)問題的理解更深一層。7. 一個(gè)具體的算子適配案例從報(bào)錯(cuò)到跑通7.1 問題現(xiàn)場adaptive_avg_pool2d未實(shí)現(xiàn)我之前幫一個(gè)團(tuán)隊(duì)做模型遷移模型里用了nn.AdaptiveAvgPool2d((1, 1))在CUDA上跑得好好的換到某國產(chǎn)芯片上就報(bào)錯(cuò)RuntimeError: Could not run aten::adaptive_avg_pool2d with arguments from the XXX backend.查了一下這個(gè)算子在PyTorch里的實(shí)現(xiàn)是CompositeExplicitAutograd也就是說它本身不直接對(duì)應(yīng)硬件指令而是通過組合其他算子實(shí)現(xiàn)的。理論上如果基礎(chǔ)算子都實(shí)現(xiàn)了這個(gè)算子應(yīng)該能自動(dòng)工作。但報(bào)錯(cuò)說明要么是組合路徑上的某個(gè)基礎(chǔ)算子沒實(shí)現(xiàn)要么是自動(dòng)微分部分出了問題。7.2 排查過程逐層分解我的排查思路是這樣的第一步確認(rèn)adaptive_avg_pool2d在CUDA后端的實(shí)現(xiàn)方式。翻PyTorch源碼發(fā)現(xiàn)它最終調(diào)用的是adaptive_avg_pool2d_out_cuda里面用了at::native::adaptive_avg_pool2d這個(gè)函數(shù)。第二步檢查這個(gè)函數(shù)依賴哪些基礎(chǔ)算子。主要是mean、view、unsqueeze這幾個(gè)。寫一個(gè)最小復(fù)現(xiàn)腳本import torch x torch.randn(1, 64, 7, 7).mlu() # 手動(dòng)模擬adaptive_avg_pool2d y x.mean(dim[2, 3], keepdimTrue) print(y.shape) # 應(yīng)該是 (1, 64, 1, 1)如果這一步報(bào)錯(cuò)說明mean算子有問題。如果這一步能過那問題出在自動(dòng)微分或者算子注冊上。第三步檢查自動(dòng)微分。adaptive_avg_pool2d的反向傳播需要adaptive_avg_pool2d_backward這個(gè)算子在CUDA后端是單獨(dú)實(shí)現(xiàn)的。如果國產(chǎn)芯片的后端沒有實(shí)現(xiàn)這個(gè)反向算子那前向能跑反向就會(huì)報(bào)錯(cuò)。7.3 解決方案注冊復(fù)合算子確認(rèn)問題之后解決方案有兩種方案一實(shí)現(xiàn)缺失的基礎(chǔ)算子。如果mean沒實(shí)現(xiàn)那就補(bǔ)mean的kernel。這是最徹底的做法但工作量大。方案二注冊復(fù)合算子。在TORCH_LIBRARY_IMPL里把a(bǔ)daptive_avg_pool2d注冊為一個(gè)CompositeImplicitAutograd算子讓PyTorch自動(dòng)用基礎(chǔ)算子組合出前向和反向。TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) { m.impl(adaptive_avg_pool2d, TORCH_FN(at::native::adaptive_avg_pool2d)); m.impl(adaptive_avg_pool2d_backward, TORCH_FN(at::native::adaptive_avg_pool2d_backward)); }這里的關(guān)鍵是at::native::adaptive_avg_pool2d這個(gè)函數(shù)本身是設(shè)備無關(guān)的它內(nèi)部會(huì)調(diào)用mean等基礎(chǔ)算子。只要基礎(chǔ)算子在你的設(shè)備上實(shí)現(xiàn)了這個(gè)復(fù)合算子就能工作。7.4 驗(yàn)證與性能測試改完之后重新跑模型model MyModel().mlu() x torch.randn(16, 3, 224, 224).mlu() y model(x) loss y.sum() loss.backward() print(Forward and backward OK)跑通之后用profiler看一下這個(gè)算子的耗時(shí)。如果發(fā)現(xiàn)adaptive_avg_pool2d的耗時(shí)占比很高那可能需要進(jìn)一步優(yōu)化比如針對(duì)(1, 1)這種輸出尺寸做特化實(shí)現(xiàn)。這個(gè)案例的通用經(jīng)驗(yàn)是遇到算子未實(shí)現(xiàn)先查PyTorch源碼看它是怎么實(shí)現(xiàn)的再?zèng)Q定是補(bǔ)基礎(chǔ)算子還是注冊復(fù)合算子。不要一上來就寫kernel很多時(shí)候組合現(xiàn)有算子就能解決問題。8. 關(guān)于PyTorch版本選擇的一些個(gè)人建議8.1 穩(wěn)定版還是Nightly版PyTorch的發(fā)布節(jié)奏是每季度一個(gè)穩(wěn)定版中間有nightly版。對(duì)于生產(chǎn)環(huán)境我強(qiáng)烈建議用穩(wěn)定版。nightly版雖然能提前用到新特性但API變動(dòng)頻繁而且可能有未修復(fù)的bug。對(duì)于芯片適配來說穩(wěn)定版還有一個(gè)好處廠商的適配通常是跟著穩(wěn)定版走的。你用nightly版可能遇到廠商還沒適配的API變更。8.2 從哪個(gè)版本開始支持PrivateUse1PrivateUse1機(jī)制是PyTorch 1.13正式引入的。如果你用的芯片廠商的適配是基于更早的版本比如1.12那它可能用的是更老的擴(kuò)展機(jī)制比如torch.utils.cpp_extension或者直接改PyTorch源碼。后者的維護(hù)成本很高每次PyTorch升級(jí)都要重新打patch。所以如果你在選擇芯片方案可以問一下廠商你們的PyTorch適配是基于哪個(gè)版本用的是PrivateUse1還是改源碼這個(gè)問題的答案很大程度上反映了廠商軟件棧的成熟度。8.3 長期支持版本的考量PyTorch基金會(huì)從2.0開始對(duì)每個(gè)大版本提供一定的長期支持。但說實(shí)話PyTorch的LTS策略不如Ubuntu那么明確。我的建議是如果你的項(xiàng)目周期比較長選一個(gè)社區(qū)活躍、廠商適配跟得緊的版本然后鎖定這個(gè)版本不要頻繁升級(jí)。我自己的項(xiàng)目里PyTorch版本是寫在requirements.txt里的精確到小版本號(hào)。升級(jí)之前一定會(huì)在測試環(huán)境里跑一遍完整的回歸測試。9. 寫在最后一些零散但有用的經(jīng)驗(yàn)做框架適配這幾年我最大的體會(huì)是軟件棧的成熟度比硬件參數(shù)更能決定一個(gè)芯片好不好用。一個(gè)算力很強(qiáng)的芯片如果PyTorch適配做得稀爛開發(fā)者用起來會(huì)非常痛苦。反過來一個(gè)算力中等的芯片如果軟件棧做得好能覆蓋大部分常用模型那它的實(shí)際可用性反而更高。寒武紀(jì)進(jìn)PyTorch理事會(huì)是一個(gè)積極的信號(hào)但也是一個(gè)起點(diǎn)。進(jìn)了理事會(huì)不等于所有問題都解決了后面還有大量的工程工作要做。對(duì)開發(fā)者來說保持關(guān)注、動(dòng)手嘗試、反饋問題是對(duì)這個(gè)生態(tài)最好的支持。最后分享一個(gè)小技巧如果你在適配過程中遇到了PyTorch的bug或者覺得某個(gè)API設(shè)計(jì)不合理可以在PyTorch的GitHub上提issue。提issue的時(shí)候附上一個(gè)最小復(fù)現(xiàn)腳本說明你的設(shè)備類型和PyTorch版本。我提過幾個(gè)關(guān)于PrivateUse1的issue社區(qū)的響應(yīng)速度比我想象的要快。參與開源社區(qū)其實(shí)沒有想象中那么遙不可及。