分組查詢注意力 是什麼?
Grouped Query Attention:分組查詢注意力 的完整解釋
一種高效的注意力機制,將多個查詢頭共享同一組鍵值頭,減少模型參數和記憶體消耗,同時保持性能不下降。
核心概念
分組查詢注意力(GQA)是對多頭注意力機制的創新優化。為了理解 GQA,先回顧多頭注意力的結構。
在標準多頭注意力中,輸入被投影到多個「注意力頭」,每個頭都有獨立的查詢(Q)、鍵(K)、值(V)投影矩陣。如果有 h 個頭,則有 h 組完整的 Q、K、V。最終輸出是所有頭的結果拼接和投影。
GQA 的創新在於:不是每個 Q 頭都有獨立的 K 和 V 頭,而是多個 Q 頭共享同一組 K 和 V 頭。例如,如果原本有 32 個注意力頭,可能改為 8 個 K/V 頭,每 4 個 Q 頭共享 1 個 K/V 頭。這樣做的結果是:Q 投影的參數數量不變,但 K 和 V 投影的參數減少到原來的 1/4。
GQA 的核心假設是:同一句話中,不同的查詢位置不一定需要完全不同的鍵值信息視角。通過共享,我們可以捕捉到主要的語義相關性,同時大幅減少冗余。
運作原理
數學表達
標準多頭注意力的計算:
Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V
其中 Q、K、V 都是 (batch_size, seq_len, h, d_k) 的形狀,h 是頭數。
GQA 中,K 和 V 的形狀改為 (batch_size, seq_len, g, d_k),其中 g 是 K/V 頭數(g < h)。Q 保持 (batch_size, seq_len, h, d_k)。
計算時,將 h 個 Q 頭分組,每 h/g 個頭共享一個 K/V 頭:
for i in range(h):
k_v_index = i // (h // g)
output[i] = softmax(Q[i] K[k_v_index]^T / sqrt(d_k)) V[k_v_index]
結果是一個折中方案:保留了多頭注意力的多個查詢視角,但減少了鍵值投影的冗余。
訓練vs推理
訓練時,GQA 和多頭注意力的計算複雜度相近,因為 Q、K、V 的計算都需要進行。
推理時,GQA 的優勢變得明顯。在使用 KV 快取加速推理時,K 和 V 從前面層的輸出快取中讀取,而不是重新計算。由於 GQA 的 K 和 V 較少(原本的 1/g),每個 token 的推理時快取大小和計算量都減少了約 g 倍。這對於長序列的自迴歸生成特別重要。
實際應用
GQA 已被多個前沿大型語言模型採納:
Google 系列
Google 的 T5 模型家族中的新版本,以及後來的 LaMDA 都採用或實驗了 GQA。Google 在論文中展示了 GQA 與標準多頭注意力在下游任務上性能相近,同時模型更小、推理更快。
Meta 系列
Meta 的 LLaMA 2 採用了 GQA。相比 LLaMA 1,LLaMA 2 中多頭注意力改為分組查詢注意力,特別是在較大的模型版本(13B、34B、70B)中。這使得 LLaMA 2 模型在部署時更高效。
Mistral 和其他新興模型
許多新興的大型語言模型,特別是那些面向邊緣設備或推理成本敏感應用的模型,都採用了 GQA。例如 Mistral 7B 使用了改進的注意力機制,結合了 GQA 思想。
實際效果
GQA 的應用效果包括:
模型大小:GQA 可減少 K/V 投影矩陣的參數,通常降低模型大小的 5-15%(取決於分組比例)。
推理速度:在自迴歸生成時,GQA 可加快推理速度 1.5-3 倍(取決於批大小和序列長度)。
記憶體使用:KV 快取的記憶體占用可減少約 g 倍(g 是共享比例),這在處理長上下文時尤為重要。
性能損失:實驗表明,在適當的分組比例下(如 1/4),下游任務性能損失小於 1%,有時甚至無損失。
常見誤區
誤區一:GQA 會大幅降低模型性能。實際上,在合適的分組比例下,GQA 對模型性能的影響極小。Google 和 Meta 的實驗表明,在多數下游任務上,使用 GQA 的模型性能與多頭注意力相近,有時甚至略優。
誤區二:GQA 只在推理時有優勢。實際上,GQA 在訓練時也有優勢。由於參數減少,訓練時的梯度計算更快,一定程度上加速訓練。
誤區三:GQA 的分組數越多越好。實際上,分組數(即 g)需要平衡。過多的分組會導致查詢頭失去多樣性,影響表達能力;過少的分組則無法獲得足夠的效率收益。通常 g = 4 或 g = 8 是較好的選擇。
誤區四:GQA 只適合推理,訓練時應該用多頭注意力。實際上,從一開始用 GQA 訓練會更高效,無須先用多頭訓練再轉換。
與相關技術的比較
與多頭注意力的比較:多頭注意力是標準方案,GQA 是其優化版本。GQA 犧牲了部分查詢頭的獨立性,換取更高的效率。在大多數實際場景中,這個折中是值得的。
與多查詢注意力(MQA)的比較:MQA 將所有查詢頭共享一個 K/V 頭,是 GQA 的極端情況(g = 1)。MQA 效率最高,但可能犧牲性能。GQA 是在 MQA 和多頭注意力之間的平衡。
與 Flash Attention 的比較:Flash Attention 是另一種優化,通過改進注意力計算的記憶體層級來加速。GQA 和 Flash Attention 可以結合使用,互補優化。
與稀疏注意力的比較:稀疏注意力限制每個位置只注意部分其他位置。GQA 和稀疏注意力都是降低注意力計算成本的方法,但思路不同。GQA 保持全注意力但減少頭數,稀疏注意力保留全頭但限制範圍。