回到文章
「PyTorch 效能分析」系列旨在讓您熟悉閱讀效能分析器追蹤圖和表格。在第一部分中,我們分析了加法和乘法等基本數學運算。我們看到了效能分析器表格如何揭示效能瓶頸,以及追蹤圖如何顯示演算法隨時間運行的順序。
在第二部分中,我們將加法和乘法包裝成一個 PyTorch 線性層。接著,我們將多個線性層堆疊起來(一個多層感知器),並對其進行了效能分析。在此過程中,我們也分析了融合核心和手動優化核心的表現。
從 Transformer 架構的角度來看,我們接下來要分析的邏輯步驟是另一個基本演算法:注意力機制。儘管它以其平方時間複雜度而聞名,但存在許多巧妙的技巧可以緩解這個問題並使其快速運行。我們的目標不是詳細介紹每個技巧,而是要觀察每個技巧在效能分析器下的不同表現。
本部落格文章的腳本位於此處:04_a_naive_attention.py、04_b_inplace_ops_attention.py、04_c_sdpa_attention.py 和 04_d_kernels_attention.py。像以前一樣,將它們在單獨的標籤頁中打開並在閱讀時逐步查看程式碼會很有幫助。
我們使用 NVIDIA A100-SXM4-80GB GPU 來運行這些腳本。在 Hugging Face 基礎設施上設置 GPU 並使用 Spaces 的開發模式來實驗這些腳本非常容易。也可以使用 Hugging Face Jobs 管道來運行這些腳本。
**樸素的注意力機制**
注意力機制使用查詢 (q)、鍵 (k) 和值 (v)。它們之間的互動可以寫成一系列簡短的步驟:
1. 建立注意力分數 `scores`:`matmul(q, k.T)`
2. 縮放分數:`scores * scale`
3. 對分數應用因果遮罩:`scores.masked_fill(mask, "-inf")`
4. 使用 softmax 正規化分數以獲得注意力權重 `attn`:`softmax(scores)`
5. 使用這些權重重新加權值:`matmul(attn, v)`
因此,注意力機制實際上是原始運算的集合。其中一些我們已經知道(矩陣乘法),其餘的也很容易辨識。讓我們在 PyTorch 中編寫一個樸素的注意力模組並對其進行效能分析。
```python
class NaiveCausalAttention(nn.Module):
def __init__(self, head_dim):
super().__init__()
self.scale = 1.0 / math.sqrt(head_dim)
def forward(self, q, k, v, mask):
scores = torch.matmul(q, k.transpose(-2, -1))
scores = scores * self.scale
scores = scores.masked_fill(mask, float("-inf"))
attn = torch.softmax(scores, dim=-1)
out = torch.matmul(attn, v)
return out
```
在打開追蹤圖之前,讓我們像往常一樣練習並猜測我們應該看到什麼。追蹤此模組的 `forward` 傳遞,我們預期會看到:
* 一個矩陣乘法核心 (`q . k.T`)
* 一個乘法核心 (縮放)
* 一個用於遮罩的運算
* 一個 softmax 核心
* 一個矩陣乘法核心 (`attn . v`)
```bash
uv run 04_a_naive_attention.py
uvx trace-util -f traces/ -b <hf_uname>/traces
```
圖 1 顯示了效能分析的 CPU 通道(GPU 通道已摺疊,以免過於龐大)。在 `attn_fwd`(我們註釋的 `forward` 呼叫)中,我們可以看到與我們猜測完全相同的運算。矩陣乘法現在是老朋友了,而新的運算也很容易辨識:
* `mul`:縮放
* `masked_fill`:因果遮罩
* `softmax`:softmax 核心
現在讓我們展開 GPU 通道,看看實際啟動了哪些核心。
圖 2 顯示了 GPU 通道與 CPU 通道並排。讓我們放大 GPU 通道上的一個 `attn_fwd` 區塊,逐一查看核心。
圖 3 讓我們可以讀取一個效能分析步驟的個別核心:
* 矩陣乘法 (查詢和鍵)
* 乘法 (縮放)
* 記憶體複製 🤔
* 因果遮罩
* softmax (產生注意力權重)
* 矩陣乘法 (注意力權重和值)
其中五個是預期的。記憶體複製是其中一個不尋常的,那麼它從何而來?線索是 PyTorch 具有原地操作。當您以普通(非原地)方式對張量進行操作時,PyTorch 通常會進行複製,對其應用請求的運算,然後返回該複製。按照運算順序,這裡的罪魁禍首是我們的 `masked_fill`。
如果我們將其替換為原地操作會怎樣?
**帶有原地因果遮罩的樸素注意力機制**
我們只將 `masked_fill` 更改為 `masked_fill_`(注意後面的底線,這是 PyTorch 原地操作的慣例),然後運行相同的腳本。
```python
def forward(self, q, k, v, mask):
# q, k, v: [batch, heads, seq, head_dim]
scores = torch.matmul(q, k.transpose(-2, -1)) # [batch, heads, seq, seq]
scores = torch.mul(scores, self.scale)
- scores = scores.masked_fill(mask, float("-inf"))
+ scores.masked_fill_(mask, float("-inf"))
attn = torch.softmax(scores, dim=-1)
out = torch.matmul(attn, v) # [batch, heads, seq, head_dim]
return out
```
讓我們看看追蹤圖,看看是否有什麼變化。
```bash
uv run 04_b_inplace_ops_attention.py
uvx trace-util -f traces/ -b <hf_uname>/traces
```
原地版本(圖 5)在遮罩步驟中包含的 CPU 運算遠少於非原地版本(圖 4)。這是一個令人鼓舞的訊號。讓我們展開 GPU 通道以確認那裡發生了什麼。
在 GPU 通道上,`Memcpy` 核心徹底消失了(圖 6 和圖 7)。透過一行程式碼的更改,我們在每次 `forward` 傳遞中減少了一個核心。這本身可能看起來不多,但請記住這是一個單一的注意力操作。在基於 Transformer 的大型模型(LLMs、擴散模型等)的背景下,它在每個層中重複一次,而且有很多層,因此節省的開銷會迅速累積(如果這能讓您加薪,至少與我們分享 10% 感覺很公平)。
非原地操作是 PyTorch 的預設值是有原因的。為了計算梯度,自動微分必須記住它在 `forward` 傳遞中看到的張量值,因為許多 `backward` 公式會重複使用它們。原地操作會覆寫記憶體中的這些值,因此 `backward` 傳遞將讀取錯誤的數字。
由於我們在 `torch.no_grad` 下運行 `forward`,因此原地操作對我們來說是安全的,沒有 `backward` 傳遞,也沒有什麼可以損壞。值得注意的是,原地操作不僅可以節省時間(如我們案例中所示),還可以節省記憶體(因為沒有額外的複製),這對於像 logits 這樣的大型張量來說非常有用!
**縮放點積注意力 (Scaled Dot Product Attention)**
我們剛剛從原始操作中構建了注意力機制,甚至減少了一個 `Memcpy`。好消息是 PyTorch 團隊已經為我們完成了所有這些工作,並將整個管道打包成一個單一函數:
```python
from torch.nn import functional as F
F.scaled_dot_product_attention(q, k, v, is_causal=True)
```
這單行程式碼取代了我們手寫的模組,`is_causal=True` 甚至省去了我們手動建立遮罩的麻煩。值得停下來欣賞這個呼叫隱藏了多少東西。它隱藏的不僅僅是程式碼行數。縮放點積注意力 (SDPA) 沒有單一的實作。在底層,它會分派到幾個後端之一,並選擇支援我們輸入(資料類型、頭部維度、遮罩、硬體等)的最快的一個。
官方的 SDPA 教學文章引導我們完成這個選擇過程,後端本身列在 `torch.nn.attention.SDPBackend` 列舉中:
```python
from torch.nn.attention import SDPBackend
BACKENDS = {
"math": SDPBackend.MATH,
"flash": SDPBackend.FLASH_ATTENTION,
"efficient": SDPBackend.EFFICIENT_ATTENTION,
"cudnn": SDPBackend.CUDNN_ATTENTION,
}
```
通常 SDPA 會為我們選擇,但我們可以使用 `torch.nn.attention.sdpa_kernel` 上下文管理器來固定特定的後端。這就是我們在腳本中所做的。這讓我們可以單獨分析每個後端,並讀取它們在追蹤圖中顯示的不同之處。讓我們一個一個來。
**Math 後端**
```bash
uv run 04_c_sdpa_attention.py --backend math
uvx trace-util -f traces/ -b <hf_uname>/traces
```
在我們打開任何東西之前,讓我們猜測一下。我們已經用一行程式碼取代了手寫的注意力機制(矩陣乘法、乘法、遮罩、softmax、矩陣乘法),所以我們應該預期追蹤圖會更簡單、更快。更少的核心、更少的 CPU 分派,甚至可能是一個融合核心。讓我們先檢查效能分析表。
| 指標 | 哪裡看? | 樸素原地 | SDPA math |
|---|---|---|---|
| `*_fwd CUDA time avg` | `*_fwd` 運算的「CUDA time avg」欄位 | 1.955 毫秒 | 7.239 毫秒 |
| `Self CUDA time total` | 效能分析表底部 | 7.194 毫秒 | 27.279 毫秒 |
這是我們第一個驚訝,這單行程式碼慢了 3.7 倍。
打開追蹤圖(圖 9)顯示了為什麼警報響起,`math` 後端每次 `forward` 傳遞啟動了 20 個 GPU 核心,而不是我們樸素注意力實作啟動的 5 個(圖 8)。這與我們猜測的完全相反。讓我們找出原因。
**Tensor 核心閒置**
在第二部分中,我們學會了像指紋一樣讀取核心名稱。讓我們在這裡運用這個習慣:
| 執行 | 矩陣乘法核心 |
|---|---|
| 圖 10:樸素注意力 | `s16816` |
| 圖 11:SDPA 與 math 後端 | `sgemm` |
我們用於擷取這些追蹤圖的 A100 GPU 搭載了 Tensor 核心,這是用於加速矩陣乘法的專用硬體,已知比普通 CUDA 核心快得多。要了解這在這裡為何重要,了解 GPU 內部結構會有所幫助。串流多處理器 (SM) 是 GPU 的計算單元,每個 SM 有兩種算術單元:CUDA 核心和 Tensor 核心。
CUDA 核心是通用型的,一次處理少量元素,而 Tensor 核心則在單一指令中乘積累積整個小型矩陣塊。所以問題很簡單:「每個後端是否真的使用了快速路徑?」
核心名稱回答了這個問題。樸素核心中的 `s16816`(圖 10)是 bfloat16 Tensor 核心矩陣乘法的簽名(16x8x16 Tensor 核心指令),因此樸素版本正在使用快速路徑。`sgemm`(圖 11)是經典的單精度 (FP32) 矩陣乘法,運行在普通的 CUDA 核心上。
換句話說,`math` 後端根本沒有觸及 Tensor 核心:為了犧牲速度換取數值精度,它將張量向上轉型為 FP32(即使輸入是 bf16,也使資料移動量加倍),並退回到較慢的 CUDA 核心。
**因果遮罩的建立**
在樸素版本中,我們只建立了一次因果遮罩並重複使用它。在這裡,我們傳遞了 `is_causal=True`,而 `math` 後端在每次呼叫時都為我們實例化了一個。您可以在 CPU 通道上觀察到它發生:
圖 12 顯示了遮罩運算的 CPU 通道。
這是我們在圖 12 中看到的:
`aten::ones` -> `aten::tril` 建立一個 `[seq, seq]` 下三角矩陣
`aten::scalar_tensor` -> `aten::fill_` 建立 `-inf` 填充值
`aten::where` 將其轉換為加性偏差 (0 或 -inf)
在 GPU 上,這顯示為 `triu_tril_kernel`、幾個 `where` 核心和一個 `add_`。讓我們不再考慮遮罩的便利旗標並沒有消除工作,它只是將工作下移了一層,在每次 `forward` 傳遞時從頭開始重建遮罩。
**安全的 Softmax**
我們手寫的版本呼叫了普通的 `aten::softmax`。`math` 後端呼叫了 `aten::_safe_softmax`,其差異再次以額外核心的形式可見(圖 13):
圖 13 顯示了安全 softmax 與通用 softmax 相比的額外核心。
一個完全被遮罩(每個條目都是 -inf)的行將使普通的 softmax 計算 `exp(-inf)/sum(exp(-inf)) = 0/`
