模型層 開放閱讀

多頭注意力

Multi-Head Attention, MHA

概念 ID
multi-head-attention-mha
更新時間
2026-05-29
來源數量
待補

多頭注意力

3秒看懂

多頭注意力是Transformer架構的“並行理解引擎”——讓模型從多個不同表示子空間同時關注輸入的不同位置,從而捕獲詞語間的多重關係(如指代、修飾、語義角色),是當代大語言模型的基石運算元。

3分鐘產業解釋

多頭注意力可以通俗地理解為“多人協作審閱一份檔案”。單個人閱讀可能會遺漏某些細節或偏見;而讓多個審閱人(多個“頭”)分別用不同的關注點(查詢向量)獨立審閱,再把他們的發現彙總起來,就得到更全面的理解。在技術實現上,它平行計算多組縮放點積注意力,每組擁有各自的查詢(Q)、鍵(K)、值(V)投影矩陣,最後拼接輸出並進行線性變換。多頭機制使得模型在不同子空間學習不同型別的依賴關係,例如某個頭專門捕捉長距離語法結構,另一個頭關注區域性語義搭配。產業側,多頭注意力直接決定了大型模型訓練和推論的算力需求,其計算複雜度為 O(n^2\cdot d)(序列長度 n,維度 d),是 GPU/TPU 設計的核心最佳化靶點。當前大量硬體(如 NVIDIA H100 的 Transformer Engine)與軟體架構(FlashAttention 演算法)均圍繞高效實現多頭注意力而演進。

15分鐘專家深入

在現代深度學習棧中,多頭注意力不只是演算法元件,更是系統設計的樞紐。它的計算圖範式深刻影響了並行策略、記憶體層級和數值精度格式。

並行策略對映:張量並行往往沿著注意力頭數維度切分,每個裝置持有部分頭,然後通過 AllReduce 聚合輸出;序列並行則將長序列切片並配合環形注意力(Ring Attention)完成跨裝置的鍵值互動。專家模型(MoE)中的 token 轉發使用 All-to-All 通訊,但注意力機制的集體通訊主要是使用 AllReduce 或 ReduceScatter 等集體通訊進行聚合(在張量並行中聚合部分頭的輸出)。這些通訊模式的選擇與硬體互聯拓撲(NVLink 頻寬、節點間 InfiniBand)共同決定訓練效率。

演算法最佳化進展:標準 MHA 的空間複雜度 O(b\cdot h\cdot n^2) 促使眾多精確近似或 IO 優化出現。FlashAttention 通過分塊(tiling)和重計算規避巨型注意力矩陣的 HBM 視訊記憶體讀寫,已成為事實標準。多查詢注意力(MQA)和分組查詢注意力(GQA)通過共享鍵值投影減少 KV 快取,大幅降低推論記憶體佔用,已被 LLaMA 2/3 等模型採用。這些變形保留了多頭機制的大部分表達能力,同時顯著改善服務延遲和吞吐。

數值格式與量化:注意力計算對數值精度敏感。FP8 混合精度訓練通常在 Softmax 部分維持高精度(如 FP32),否則梯度會不穩。推論側,KV 快取常用 INT8 或 FP8 量化,需要特殊的校準方法(如 SmoothQuant)抑制異常值。2:4 結構化稀疏可在保持硬體加速的同時剪枝一些注意力頭,兼顧效能。

關鍵引數的量級關係:對於千億引數模型,典型頭數在 32–128,每頭維度 128,總隱藏維度 d_{model} = h \cdot d_k 通常為 4096 至 16384。序列長度從 2048 擴充套件到 32768 甚至百萬 token(通過位置編碼外推),使得注意力矩陣呈平方增長,成為長上下文的關鍵瓶頸。

技術原理

縮放點積注意力(單頭)

給定輸入序列表示 X \in \mathbb{R}^{n \times d_{model}},通過權重矩陣 W^Q, W^K, W^V 線性對映得到查詢 Q、鍵 K、值 V,維度 d_k = d_{model}/h

Q = XW^Q, \quad K = XW^K, \quad V = XW^V

注意力權重計算:

\text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

其中 \sqrt{d_k} 作為縮放因子,防止點積過大導致 Softmax 梯度消失。

多頭拼接

h 個頭並行執行上述注意力,拼接結果再進行一次線性變換:

\text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1, \dots, \text{head}_h) W^O
\text{where head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)

投影矩陣形狀:W_i^Q \in \mathbb{R}^{d_{model} \times d_k}W_i^K \in \mathbb{R}^{d_{model} \times d_k}W_i^V \in \mathbb{R}^{d_{model} \times d_v},通常 d_k=d_v=d_{model}/hW^O \in \mathbb{R}^{hd_v \times d_{model}}

複雜度與記憶體

  • 計算複雜度:O(n^2 \cdot d_{model})(包括 QK 乘法和加權 V)
  • 記憶體佔用:儲存注意力矩陣 n \times n 每個頭需要 O(n^2),若快取中間啟用則成為視訊記憶體瓶頸。標準實現中視訊記憶體需求 O(b \cdot h \cdot n^2)b 為 batch size。
               輸入序列 X: [n, d_model]
                      |
        +-------------+-------------+
        |             |             |
    Linear(W_Q)   Linear(W_K)   Linear(W_V)
        |             |             |
    Q: [n,dk]    K: [n,dk]    V: [n,dv]
        |             |             |
        +------> MatMul(Q,K^T) / sqrt(dk)
                      |
              Softmax(axis=-1)
                      |
                  MatMul(Score, V)
                      |
                 head_output: [n, dv]
        (重複h次,拼接) -> [n, h*dv]
                      |
                 Linear(W_O) -> [n, d_model]

掩碼注意力

解碼器中的因果掩碼(上三角矩陣置 -\infty)確保位置 i 只能看到 i 及之前的資訊,維持自迴歸生成。

技術演進史

  • 2014,基礎注意力:Bahdanau 等人在 RNN 編解碼器中引入加法注意力,讓解碼器動態選擇源句相關部分。計算代價高,序列化明顯。
  • 2015,乘法注意力:Luong 提出點積注意力,簡化計算,但仍與 RNN 耦合。
  • 2017,Transformer 與多頭注意力:Vaswani 等釋出《Attention Is All You Need》,完全拋棄迴圈和卷積,提出多頭縮放點積注意力作為唯一特徵互動機制,定義標準 MHA 範式,開啟大型模型時代。
  • 2019–2020,效率最佳化:Reformer 使用區域性敏感雜湊(LSH)注意力降低複雜度;Longformer、BigBird 結合滑動視窗、全域性和隨機稀疏注意力;Linformer 將注意力矩陣低秩分解,理論複雜度降至 O(n)
  • 2021,多查詢注意力(MQA):將所有頭共享一副鍵值投影,大幅減少推論 KV 快取,PaLM、LaMDA 等模型採用。
  • 2022–2023,分組查詢注意力(GQA):介於 MHA 和 MQA 之間,將頭分為若干組共享 KV 投影,平衡效率與效果,LLaMA 2、Mistral 等主流模型採用。
  • 2022 至今,IO 感知精確注意力:FlashAttention(Dao et al.)利用 GPU 共享記憶體和線上 Softmax 減少 HBM 訪問,經歷 V1/V2/V3,支援 FP8 和動態序列,已成為 Transformer 建置標準背板。
  • 2024+,長上下文與系統融合:Ring Attention 將序列分塊並環通訊,實現跨裝置百萬 token 上下文;Striped Attention 非同步重疊通訊與計算;MLA(Multi-head Latent Attention)在 DeepSeek 等模型中應用,通過低秩鍵值聯合壓縮排一步降低推論代價。

技術路線對比(量化表)

以下對比基於公開論文與業界實踐(符號:← 越大越好 / 越小越好,定性標記)。

變種頭數/分組每頭維度KV 快取大小(相對 MHA)推論視訊記憶體佔用(相對)訓練吞吐(相對)長文本質量損失典型模型
MHAh 頭,獨立 KVd/h低(基準)BERT,GPT-2,ViT
MQAh 頭,全部頭共享 KV同上1/h極低1.2–1.5×中等PaLM,LaMDA
GQAg 組,組內共享 KV同上g/h1.1–1.3×低–中LLaMA 2/3,Mistral
FlashAttention同 MHA(精確等價)~相同相同(但視訊記憶體頻寬佔用低)2–3×(長序列)零(等價)幾乎所有現代大型模型

注:具體數字因序列長度、硬體而異,[根據公開基準測試估算],未精確揭露完整型號的 head-to-head 測量資料。

上下游

  • 上游
    • 基礎數學庫(cuBLAS、rocBLAS)提供高效矩陣乘法;
    • 底層硬體(GPU Tensor Core、TPU MXU)支援半精度矩陣乘累加;
    • 編譯器與架構(Triton、TVM、XLA)生成最佳化的融合核;
    • 並行策略庫(NCCL、NVSHMEM)為頭並行提供通訊原語。
  • 下游
    • 自然語言處理:語言模型(GPT、LLaMA、Gemini)、機器翻譯、文本摘要;
    • 計算機視覺:ViT、Swin Transformer 等將影像分塊作為序列輸入;
    • 多模態:CLIP、Flamingo、ImageBind 使用交叉注意力對齊圖文特徵;
    • 語音/生物序列:Whisper、AlphaFold 利用注意力建模時序和殘基關係。

關鍵指標

  • 效率指標
    • TFLOPS(每秒浮點運算次數)在注意力實現上的利用率;
    • 視訊記憶體頻寬利用率(HBM 讀寫佔比);
    • 推論階段首個 token 延遲與每個 token 生成延遲;
    • KV 快取容量(位元組):約 2\cdot L\cdot h\cdot d_k\cdot N\cdot \text{precision}L 層數,N 序列長度。
  • 能力指標
    • 有效上下文長度:注意力機制能維持準確性的最大 token 距離;
    • 檢索精度(如 Needle-in-a-Haystack 測試分數);
    • 不同頭功能的可解釋性(頭重要性、啟用稀疏度)。

供需與市場資料

由於未獲得最新的專門針對多頭注意力市場規模的獨立報告(該元件通常被歸入更大範圍的 AI 訓練/推論晶片與軟體市場),根據行業趨勢定性:隨著大型語言模型引數量突破萬億、上下文視窗向 1M token 延伸,高效的多頭注意力實現成為 AI 加速器的核心競爭力。據 NVIDIA FY2024 財報揭露,資料中心 GPU 營收約 475 億美元(2024 財年),其中很大比例由 Transformer 及注意力計算驅動。訓練成本方面,據公開技術報告估算,在大型 MoE 模型中,注意力計算的浮點運算佔比通常顯著低於 50%,大部分浮點運算消耗在 FFN/專家層。推論側,KV 快取最佳化能直接降低 30% 以上的 GPU 例項成本,因此 MQA/GQA 等變種已快速普及。開源社群對 FlashAttention 的依賴度極高,幾乎所有的 PyTorch / JAX 訓練指令碼預設呼叫最佳化的注意力後端。

代表公司與資本對映

  • Google(Google):多頭注意力原提出者,Transformer 架構至今仍是其 Gemini 系列模型的基礎;TPU 晶片專門針對注意力計算優化了脈動陣列。
  • 輝達(NVIDIA):通過 TensorRT-LLM、cuDNN 和 Transformer Engine 提供高效能多頭注意力實現,硬體 Blackwell 架構引入專門的注意力引擎;其資料中心 GPU 生態受益於模型規模擴大帶來的注意力算力需求激增。
  • OpenAI / Anthropic / Meta 等模型廠商:在內部訓練架構中深度定製注意力核(如提升長上下文效率的演算法),是 MQA/GQA 等變種的採納者和推動者。
  • AI 晶片初創(Cerebras、Groq、SambaNova):各自通過資料流架構、大規模 SRAM 或特定運算元重構來降低注意力延遲,資本市場高度關注其在長上下文推論中的突破。
  • 風險投資趨勢:投資者傾向於支援那些能顯著降低 Transformer 注意力計算成本的基礎設施公司,包括 FlashAttention 商業化變體(如資料庫索引型注意力)、記憶體池化方案以及存內計算架構。

投資邏輯

  1. 算力供需剪刀差:序列長度每翻倍,裸注意力計算量增長 4 倍;而硬體算力增長放緩(摩爾定律瓶頸),導致高效注意力實現成為剛性價值點,相關軟體(FlashAttention 類庫)和硬體(支援更優稀疏/量化的晶片)受益。
  2. 推論成本決定模型商業閉環:PaaS / SaaS 部署大型模型時,注意力推論佔用 30–70% 的 GPU 時間,因此 GQA、KV 快取量化等技術的成熟度直接影響模型盈利性,押注相關技術棧的公司具備降本增效的長期需求。
  3. 專利與標準風險:Transformer 及 MHA 相關基礎專利(如Google的注意力機制專利)申請於2018年,仍在20年有效期內,遠未過期,但特定最佳化實現(如 FlashAttention 的 tiling 策略)可能受版權/專利影響,需要關注開源生態的健康度。
  4. 邊緣智慧:小模型(≤7B)在端側部署時,記憶體約束極為苛刻,推動分組注意力、INT4 KV 快取等方案,利好低功耗存算一體晶片和極致壓縮工具鏈。

常見誤讀糾偏

誤讀1:“多頭注意力就是讓模型同時理解不同含義,頭數越多效果越好” 實際上頭數受限於每頭維度 d_k 不能過小,否則投影矩陣的容量不足以形成有效的注意力子空間。實驗表明,適中的頭數(如 GPT-3 使用 96 頭)已在各項任務上飽和,繼續增加頭數會帶來更多計算開銷且收益遞減。同時,許多頭在訓練後表現為冗餘或可剪枝,並非所有頭都承載獨立語種功能,其作用往往是高度混合的。

誤讀2:“FlashAttention 近似了注意力計算,因此會損失精度” FlashAttention 及其後續版本實現的是數學上完全等價的標準縮放點積注意力,通過分塊和線上 Softmax 在不改變數值結果的前提下最佳化 IO 模式。其輸出的梯度與原始實現一致(相對於數值誤差環境),因此嚴格意義上不存在精度損失,不應與其他近似稀疏/低秩注意力混為一談。

學習路徑

  1. 基礎入門:閱讀《Attention Is All You Need》論文,重點理解 Section 3.2 多頭注意力的公式和圖形。結合 PyTorch 官方 nn.MultiheadAttention 原始碼(或最小實現)手寫一個 MHA 模組。
  2. 核心演算法:學習《FlashAttention: Fast and Memory-Efficient Exact Attention》,理解 Triton 或 CUDA 如何分塊計算 Softmax 並減少 HBM 讀/寫。嘗試分析 HuggingFace Transformers 中 LlamaAttention 類的分組查詢實現。
  3. 系統與並行:閱讀 DeepSpeed Ulysses 序列並行論文,瞭解張量並行中如何沿著頭維度切分注意力;閱讀 vLLM 的 PagedAttention 機制,理解 KV 快取管理如何影響多請求服務。
  4. 產經視角:追蹤 SemiAnalysis、Next Platform 等分析文章,瞭解不同 AI 晶片(如 Groq LPU,NVIDIA H200)在處理注意力時的頻寬和延遲瓶頸。調研開源模型結構報告(如 LLaMA 論文、Mixture-of-Experts 技術報告)中的注意力配置選擇及其工程原因。

一句話總結

多頭注意力以“多視角並行匹配”取代了序列迴圈,是 Transformer 大型模型的感知中樞,其計算效率與記憶體方案直接定義了大型模型的商用成本天花板。

延伸閱讀與來源

  • 原始論文:Vaswani, A. et al. “Attention Is All You Need.” NeurIPS 2017.
  • FlashAttention 系列:Dao, T. et al. “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.” NeurIPS 2022; “FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning.” 2023; “FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision.” 2024.
  • 分組查詢注意力:Ainslie, J. et al. “GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints.” EMNLP 2023.
  • 系統實現:NVIDIA Megatron-LM 張量並行論文(Shoeybi et al., 2019); DeepSpeed Ulysses 序列並行; vLLM PagedAttention (Kwon et al., SOSP 2023).
  • 行業分析:受限於檢索狀態,具體供需數字未直接獲取,可關注 Omdia、IDC 關於 AI 伺服器/晶片的資料;模型效率分析見 SemiAnalysis 的公開文章。
source: 公開揭露與公開資料整理 本頁僅用於產業鏈學習、資訊檢索和研究輔助;不構成投資建議,不預測漲跌,不提供買賣、部位或目標價建議。
完整概念頁 複盤 13 節結構 公司投研頁 沿產業鏈找到受益公司 投資課 把概念轉成可跟蹤模型