來源最新收集於 4h

FlashAttention-4 加入 Blackwell 低精度支援

閱讀原文: PyTorch Blog
#attention-kernels#mxfp8#gpu-optimization

升級核心可能讓 Blackwell GPU 的注意力效能突破每秒數 PFLOPS。

30 秒速覽

有什麼變化

MXFP8 同時支援注意力的前向與反向計算。

為什麼重要

低精度注意力可提升訓練與推理吞吐量,同時降低記憶體與頻寬壓力。使用 Blackwell GPU 的團隊可能只需升級核心,就能獲得顯著效能提升,而不必更換模型架構。

下一步行動

在 Blackwell 測試節點上使用 FlashAttention-4 MXFP8 執行訓練核心,並比較吞吐量、損失穩定性與記憶體使用量。

誰應關注:Developers & AI Engineers

關鍵要點

  • MXFP8 同時支援注意力的前向與反向計算。
  • 據報前向效能達 2.85 PFLOPS,反向效能達 2 PFLOPS。
  • 內部形狀測試顯示 FA4 MX8 可達 2.54 PFLOPS。
關鍵數字1.3 倍2.7 倍

深度解析

背景與延伸:來自公開資料,非原文內容。引用 13 個來源。

增強重點摘要

  • FlashAttention-4 採用 NVIDIA 的 CuTeDSL(Python 領域特定語言)完全重寫核心,取代了傳統龐大的 C++/CUDA 實作,將 JIT 編譯時間從 55 秒大幅縮短至約 2.5 秒。
  • 專為 NVIDIA Blackwell 架構重新設計,針對 SM100 移除 WGMMA 並引入新 TCGEN05 指令集與每 SM 獨立 256 KB Tensor Memory (TMEM) 硬體子系統進行管線化重構。
  • 克服 Blackwell 的非對稱硬體擴展(Asymmetric Hardware Scaling)瓶頸,解決 Tensor Core 算力翻倍但共享記憶體與指數運算單元未等比提升導致 Softmax 佔用過多運算時間的問題。
  • 相容性僅限於資料中心級 Blackwell GPU(SM100/SM103,如 B100、B200、GB200、B300)及 Hopper 架構,消費級晶片(如 RTX 5090 / SM120)因缺乏 TMEM 硬體而無法運行。
  • 在 HGX B200 基準測試中,FlashAttention-4 相較於 NVIDIA cuDNN (v9.13) 取得 1.2 至 1.3 倍的加速,並比 Triton 實作快 2.1 至 2.7 倍,已被整合進 vLLM 與 SGLang 等主流推論框架。

競品分析

FlashAttention-4 (FA4 MX8)
支援硬體與架構特性
NVIDIA Blackwell 資料中心級 (SM100/SM103) & Hopper;利用 TCGEN05、TMEM 與 TMA
程式語言 / 編譯環境
CuTeDSL (Python DSL);JIT 編譯僅需約 2.5 秒
相對 Blackwell HGX B200 效能表現
基準最高效能(前向達 2.85 PFLOPS,反向達 2.0 PFLOPS)
NVIDIA cuDNN (v9.13)
支援硬體與架構特性
NVIDIA 資料中心與消費級 GPU;官方閉源最佳化核心
程式語言 / 編譯環境
C++ / CUDA 靜態函式庫或封裝
相對 Blackwell HGX B200 效能表現
比 FA4 慢約 1.2× 至 1.3×
Triton Attention
支援硬體與架構特性
跨架構支援(NVIDIA、AMD 等開源編譯器堆疊)
程式語言 / 編譯環境
Python-based Triton JIT
相對 Blackwell HGX B200 效能表現
比 FA4 慢約 2.1× 至 2.7×

技術深入

  • 架構指令轉型:Blackwell SM100 完全棄用了 Hopper 架構的 WGMMA 指令,改用全新 TCGEN05 指令集,並結合每 SM 配備的 256 KB Tensor Memory (TMEM) 與硬體管理的 TMA (Tensor Memory Accelerator)。
  • 非對稱硬體擴展因應:Blackwell 的 FP16/BF16 Tensor Core 算力相較 Hopper 翻倍(達 2.25 PFLOPS),但共享記憶體頻寬與指數運算單元(SFU)未成比例擴展,導致 Softmax 等非矩陣乘法操作在無管線重排時佔比高達 25%–60%;FA4 透過選擇性重新縮放與精細重疊排程消除此瓶頸。
  • CuTeDSL 語言實作:捨棄傳統單體 C++/CUDA 開發模式,改採純 Python 的 CuTeDSL 撰寫核心管線,保留極致硬體裸機效能的同時大幅縮短 JIT 編譯延遲(~2.5 秒)。
  • 深度整合 FlexAttention:直接內嵌於 PyTorch 的 torch.nn.attention.flex_attention API,允許開發者以純 Python 自訂 Mask 與注意力分數計算,並自動轉譯為 FA4 高效核心運行。

前景展望基於引用來源的 AI 分析

消費級 GPU 與資料中心 GPU 在高效能注意力運算上的軟體生態將正式分流
由於 FA4 高度依賴 SM100/SM103 的 TMEM 硬體架構,消費級晶片(SM120)因硬體閹割無法執行該核心,將迫使推論社群維護兩套完全獨立的注意力和量化執行路徑。
MXFP8 將在主流開源大型語言模型推論框架中迅速取代 FP8/BF16 成為預設推論格式
FA4 透過 CuTeDSL 整合進 vLLM 與 SGLang,並在 Blackwell 上提供高達 2.85 PFLOPS 的實測前向吞吐量,能顯著降低資料中心算力持有成本。

時間線

2026-03
arXiv 發表 FlashAttention-4 演算法與非對稱硬體擴展論文 (arXiv:2603.05451)
2026-03
PyTorch 正式發布 FlexAttention 與 FlashAttention-4 整合後端
2026-09
PyTorch 部落格正式發表支援 Blackwell 低精度 MXFP8 的 FlashAttention-4

AI 週報

閱讀本週精選 AI 大事摘要 →

AI 策展新聞聚合。所有內容版權歸原始發布者所有。
原始來源: PyTorch Blog

這是摘要,不是原文。去看原站,或訂閱每週簡報。

每週電子報

每週一封,可隨時退訂。