Skip to content

CS336 Lecture 5:GPU 快不是因為每個 thread 快,而是資料少搬幾次

2026年8月22日 1 分鐘
TL;DR 第五講從 SM、warp 與記憶體階層解釋 GPU,再用低精度、fusion、recomputation、coalescing 與 tiling 統一常見優化;FlashAttention 正是這些原則在 attention 上的組合。
目錄
  1. GPU 選擇 throughput,不是單一 thread latency
  2. 記憶體越靠近 SM 越小也越快
  3. 五種優化其實都在處理資料路徑
  4. 為什麼矩陣 shape 會造成週期性效能
  5. FlashAttention 是前面原則的總和
  6. 讀完後怎麼除錯效能
  7. 材料完整度
  8. 參考資料

🌏 English version

本篇對應 CS336 Spring 2026 Lecture 5: GPUs, TPUs,2026 年 4 月 13 日由 Tatsunori Hashimoto 主講。主要來源是官方 lecture_05.pdf

第二講已用 roofline 說明 compute-bound 與 memory-bound;第五講打開 GPU,看這兩種瓶頸如何從硬體產生。它的目標不是把 CUDA 術語背完,而是讓「為什麼這個 shape 慢」「FlashAttention 為什麼快」不再像魔法。

GPU 選擇 throughput,不是單一 thread latency

CPU 用複雜控制、cache 與少數強核心縮短單一 thread 的完成時間。GPU 放入大量較簡單的計算單元,讓很多 threads 執行相同指令、處理不同資料,追求總吞吐。

GPU 由多個 streaming multiprocessors(SM)組成;thread 組成 block,block 被排到 SM;threads 又以 warp 為執行單位。Warp 內若走不同 conditional branch,硬體要分批執行不同路徑,形成 divergence。條件式不是不能寫,而是它會削弱同指令多資料的優勢。

TPU 的高階原則相近:輕量控制、大型矩陣乘法單元與高頻寬記憶體。差異主要在計算單元組織與裝置互連,而不是「一種能做矩陣乘法、另一種不能」。

記憶體越靠近 SM 越小也越快

Thread 有 registers,block 內可共享 shared memory;再外面是 L2 cache 與 HBM/global memory。跨 block 共用資料通常得經過較慢的全域記憶體。現代 GPU 的矩陣乘法吞吐成長快於記憶體頻寬,因此讓 tensor cores 吃飽資料成為核心問題。

這就是 arithmetic intensity 的硬體版本。資料從 HBM 搬進來一次,若能在 registers 或 shared memory 重用,就能用更多計算攤平搬運成本。每做一個簡單操作就寫回 HBM,則再多 tensor cores 也會等資料。

五種優化其實都在處理資料路徑

低精度讓每個值占用較少 bytes,也讓 tensor cores 提供更高吞吐。FP8、MXFP8 與更低 precision 仍需 scale factors,並非把所有 tensor 無條件轉型。

Operator fusion 把連續 pointwise operations 合成一個 kernel。中間值留在 registers 或 shared memory,不必每一步都寫回 HBM再讀回。

Recomputation 丟棄某些中間 activation,需要時再算。算術變多,記憶體讀寫反而變少;在 memory-bound 區域,重算可能比讀取快。

Memory coalescing 讓同一 warp 的 threads 存取連續位址,配合 DRAM burst 一次取回有效資料。索引相同、只是 traversal direction 不同,就可能造成巨大差異。

Tiling 把矩陣切成能放入 shared memory 的小塊。每個 tile 從 HBM 載入一次後被多次重用,再寫回輸出。Tile size 還要考慮 shape 是否整除、alignment、register pressure 與同時可駐留的 blocks。

為什麼矩陣 shape 會造成週期性效能

方形矩陣只增加一個元素,runtime 也可能突然跳高。原因不一定是 FLOPs,而可能是 tile 邊界、記憶體對齊或 wave quantization:SM 數量固定,最後多出的一小批 blocks 仍需要完整一個排程 wave。

因此 benchmark 不能只測一個漂亮的二次方 shape,也不能把峰值規格當實際速度。要掃過問題尺寸,觀察週期性與斷點,再用 profiler 判斷 kernel、memory transaction 與 occupancy。

FlashAttention 是前面原則的總和

標準 attention 會建立大型 score matrix,套 softmax,再乘上 values。若每一步都把中間矩陣寫入 HBM,序列變長時資料搬運非常昂貴。

FlashAttention 把 Q、K、V 切成 tiles,在較快的 SRAM/shared memory 中分塊計算。Softmax 使用 online normalization,逐塊維護最大值與 normalization sum,不必一次持有完整 score matrix。Backward 時可重算部分值,進一步減少儲存。

它仍然計算 exact attention;加速主要來自 IO-aware algorithm,而不是近似掉一部分 attention。這正好串起本講五項技巧:tiling、fusion、較少 HBM traffic,以及有意識的 recomputation。

讀完後怎麼除錯效能

先用 roofline 判斷 compute 或 memory bound,再檢查 precision 是否走 tensor cores、operations 能否 fusion,以及 activation 是否值得重算。接著檢查存取是否 coalesced、tile 是否符合 shape 與硬體。最後用多個尺寸 benchmark 和 profiler 驗證。

GPU 優化最容易被寫成技巧清單;第五講把它們收斂成一個問題:資料從哪裡來、在哪裡重用、什麼時候被寫回去。

材料完整度

本講有 Spring 2026 當期 schedule 與完整官方 PDF。本文依投影片的硬體、效能技巧與 FlashAttention 三部分整理。

參考資料