來自 Hugging Face 的 PyTorch Profiling 系列第三部曲,透過 profiler trace 比較五種 attention 實現在 GPU 層面的核心數量、記憶體行為與效能差異。

核心概念

本文核心方法論是「先猜後看」——開啟 profiler 前先陳述預期核心數量,再比對 trace 找出不匹配。所有重要發現都來自這個「猜測 vs 現實」的偏差。

Attention 基本五步驟: Query × Key → 縮放 → 因果掩碼 → Softmax → 加權值向量

天真實現的隱藏成本

最基本的手寫 attention 啟動 6 個 GPU 核心,其中一個是 masked_fill 隱含的記憶體複製(非原位操作觸發複製)。改為 masked_fill_(原位版本)後,單行修改省去一整個核心,降至 5 個。原位操作在 autograd 下會破壞梯度,僅適用於 torch.no_grad() 情境。

SDPA 四種後端對比

PyTorch 的 scaled_dot_product_attention 介面背後有四個後端:

後端 核心數 主要特徵
Math(參考) 20 FP32 upcast、CUDA core、每次重建掩碼
Efficient(xformers) 1 融合 fmha_cutlassF_bf16,Tensor core
Flash(FlashAttention-2) 1 pytorch_flash,最快,13% occupancy
cuDNN(生成式) 1 執行時生成 kernel,無 transpose,但 CPU 成本高

Math backend 比天真實現慢 3.7 倍——原因是強制 FP32 upcast,繞過 Tensor core,改用 CUDA core 的 sgemm 路徑。

FlashAttention:低 occupancy 為何反而正確

Flash backend 的 GPU occupancy 僅 13%,看似低效,但這是刻意設計:每 block 使用 255 個寄存器,一個 SM 只能容納 2 個 block(12.5% occupancy)。這些寄存器用於執行 online softmax,讓完整的 [seq, seq] 注意力分數矩陣永遠不寫入 HBM(全局記憶體),根本消除記憶體頻寬瓶頸。

低 occupancy ≠ 低效率。高 occupancy 靠多 block 隱藏延遲;Flash 靠片上數據重用消除 HBM 往返。

cuDNN 的成本轉移

cuDNN backend 有 0 個 transpose 操作,但 CPU profiler 顯示 214μs 的厚條帶,遠高於 Flash 的 138μs。工作沒消失,被藏進 _cudnn_attention_forward 的 knob 選擇與執行時 kernel 生成。cuLaunchKernelEx(driver-level API)還讓 CUPTI 無法追蹤 occupancy,顯示為 0%——這是測量盲點,非效能問題。

關鍵要點

  • 一行原位操作消除一個核心masked_fillmasked_fill_ 省去 Memcpy,節省記憶體與時間
  • 核心融合是效能倍增器:從 6 個核心(天真)到 1 個(融合後端),每次 kernel launch 都有啟動開銷
  • bf16 vs FP32 走不同路徑:Math backend 強制 FP32 upcast,繞開 Tensor core 快速路徑,慢 3.7×
  • Flash 的 HBM 規避是真實節省:完整注意力矩陣從不觸及全局記憶體,而非隱藏成本
  • 工作只會移動不會消失:cuDNN 的 0 transpose 把成本轉到 CPU 的 knob 選擇,profiler 看不見

實務應用

後端選擇依據:

  • 通用訓練推論 → Flash(最低 HBM 流量、序列長度穩定)
  • 大 head dimension 或非標準形狀 → cuDNN(有 knob 預調優)
  • 偵錯精度問題 → Math(保留 _safe_softmax、FP32 數值穩定)
  • 明確使用 xformers → Efficient(bf16,單融合核心)

建立 profiling 習慣的三步驟:

  1. 開啟 trace 前寫下「我預期看到 N 個核心、X 個 op」
  2. 打開 trace 比對現實
  3. 追問每一個不匹配的「為什麼」直到答案清晰

這套方法不需要深厚 GPU 硬體知識,只需仔細觀察的紀律。

相關頁面:torch.profiler 入門指南:矩陣乘法帶你解碼 GPU 效能軌跡 | 非同步連續批次推論:LLM 推論的 CPU GPU 並行加速 | Delta Weight Sync in TRL:異步強化學習的百倍同步壓縮

延伸觀點

FlashAttention 演進:FA-2 到 FA-3

多篇論文共同確認了 FlashAttention 系列的演進路徑。FlashAttention-3(arXiv 2407.08608)針對 Hopper/H100 架構利用非同步技術(TMA、WGMMA)與 FP8 量化,在 H100 上達到 840 TFLOPS(85% 理論上限),比 FA-2 快約 2 倍。本文剖析的 FA-2 「低 occupancy 高效能」邏輯在 FA-3 中進一步延伸:H100 的 Warp Specialization 讓生產者/消費者 warp 異步重疊計算與記憶體存取,occupancy 看起來更低,但實際吞吐更高。

編譯器加速注意力變體(Flashlight,arXiv 2511.02043)

另一個趨勢是跳脫固定後端,改用 PyTorch compiler 基礎設施(TorchInductor + Triton)自動生成融合 kernel。Flashlight 將多種注意力變體(稀疏、線性、帶寬約束)統一接入編譯路徑,解決 SDPA 四個後端覆蓋不到的非標準形狀問題。這與 FlexAttention 的設計理念一致:讓開發者用 Python 描述注意力變體,底層自動生成最佳化 kernel。

跨配置 LLM 推論分析(Dooly,arXiv 2605.07985)

傳統 profiler(包括本文使用的 PyTorch profiler)在不同 batch size 或硬體配置間缺乏可遷移性,需要重新 profiling。Dooly 透過捕捉 kernel 執行模式,在不同配置下複現推論行為,讓開發者無需重新佈建測試環境即可預測效能——補足了本文方法論的一個實務缺口。

反向連結

以下頁面引用了本頁: