在 PyTorch 中進行分析(第 3 部分):注意力就是你要分析的一切

Profiling in PyTorch (Part 3): Attention is all you profile

Profiling in PyTorch (Part 3): Attention is all you profile

這篇部落格文章使用 NVIDIA A100-SXM4-80GB GPU 對 PyTorch 中的注意力機制進行分析,並示範每個變體在分析器追蹤中的呈現方式。

Naive attention

一個由原始操作(matmul, mul, masked_fill, softmax, matmul)構建的天真注意力實現,每次前向傳播會啟動六個 GPU 核心,其中包括由 out‑of‑place masked_fill 造成的意外記憶體複製。

Naive attention with inplace causal masking

masked_fill 替換為就地版本 masked_fill_ 會移除記憶體複製核心,使每次前向傳播的 GPU 核心數從六個減少到五個。

Scaled Dot Product Attention

PyTorch 的 F.scaled_dot_product_attention 提供單行介面,會分派到多個後端;每個後端在分析器中呈現出不同的特徵。

Math backend

當固定在 math 後端時,SDPA 每次前向傳播會啟動約二十個 GPU 核心,在 CUDA 核心上以 FP32 運行,每次呼叫都會重建因果遮罩,並使用 _safe_softmax 來避免 NaN;其速度大約是天真就地版本的 3.7× 慢。

Efficient backend

高效後端(xformers)每次前向傳播會啟動單個融合的 fmha_cutlassF_bf16_aligned_64x64_rf_sm80 核心,將計算保持在 Tensor 核心的 bfloat16 格式,並避免在 HBM 中寫入中間矩陣。

Flash backend

Flash 後端每次前向傳播會啟動單個融合的 pytorch_flash 核心(FlashAttention‑2);儘管在分析器中顯示的估計佔用率較低(約 13%),但由於它使用暫存器和共享記憶體將注意力塊保持在晶片上,因此是最快的後端。

cuDNN backend

cuDNN 後端每次前向傳播會啟動單個生成的核心(例如 cudnn_generated_fort_native_sdpa_sm80_flash_fprop_wmma_f16_knob_6_128x64x64_4x1x1_cga1x1x1_kernel0_0);它避免了 CPU 上的轉置操作,但由於運行時計畫選擇導致 CPU 時間增加(約 214 µs),而其 GPU 時間位於高效後端與 Flash 後端之間。

Everything we covered, at a glance

變體 我們改變的內容 每次前向傳播的核心數 追蹤揭示的內容
Naive attention 手動從原始操作(matmul, mul, mask, softmax, matmul)構建的注意力 6 來自 out‑of‑place masked_fill 的隱藏 Memcpy
Naive in-place masked_fillmasked_fill_ 5 一行程式碼完全移除了 Memcpy 核心。
SDPA math F.scaled_dot_product_attention 固定在 math 後端 20 參考實作:在 CUDA 核心上使用 FP32,每次呼叫重建遮罩,使用 _safe_softmax。正確但大約慢 3.7×。
SDPA efficient 高效後端(xformers) 1 單個融合的 fmha_cutlassF 核心,保持在 Tensor 核心的 bf16 格式。
SDPA flash Flash 後端 1 單個融合的 pytorch_flash 核心(FlashAttention‑2)。儘管顯示的佔用率看似錯誤(13%),但仍是最快的。
SDPA cuDNN cuDNN 後端 1 每個問題生成的核心:無轉置操作,cuLaunchKernelEx,但成本轉移到了較大的 CPU 條形圖。

系列總結

主要的收穫是先猜測分析器追蹤應該顯示什麼,然後檢查追蹤,將任何不符合視為最有趣的見解;這個習慣揭示了隱藏的 Memcpy、math 後端的二十個核心、flash 的低佔用率以及 cuDNN 的 CPU 端成本。

現在您可以將這種猜測‑檢查方法應用於分析您自己的模型。

Sources