· liyu · ai · 24 min read
深入大模型注意力演进:MHA、MQA 与 GQA 的张量计算与优缺点全景剖析
聚焦大模型三大核心注意力机制:多头注意力(MHA)、多查询注意力(MQA)与分组查询注意力(GQA)。从张量维度(Tensor Shapes)与计算流出发,深度解读前向传播中矩阵的行、列、元素物理含义,剖析三者的优缺点与显存带宽权衡。
在构建和部署大语言模型(LLM)时,**注意力机制(Attention Mechanism)**不仅决定了模型捕捉长距离上下文和复杂语义的能力,更是决定推理吞吐量(Throughput)与显存占用(VRAM footprint)的核心瓶颈。
从原始 Transformer 确立的 多头注意力(MHA, Multi-Head Attention),到激进缩减显存开销的 多查询注意力(MQA, Multi-Query Attention),再到如今成为主流大模型标配的 分组查询注意力(GQA, Grouped-Query Attention),这一演进过程深刻体现了算法设计与底层硬件物理限制之间的博弈与妥协。
本文将专门聚焦于 MHA ➔ MQA ➔ GQA 三大核心机制,结合模型训练与推理中张量(Tensor)的实际计算全流程,细致拆解每一个矩阵的维度结构、行代表什么、列代表什么、每个具体元素代表什么,并深入剖析各自的优缺点与工程权衡。
一、 核心符号与张量基础定义
为了精准理解后续的张量流转过程,我们统一约定以下张量符号与维度记号:
B (Batch Size):批次大小,表示同时处理的独立样本/文本序列数。S (Sequence Length):序列长度(训练时为完整上下文 Token 数,推理 Prefill 阶段为 Prompt 长度,Decode 阶段当前生成步的 Token 长度为 1)。D (Hidden Dimension / Embedding Dim):模型的隐藏层维度(如 4096、8192)。H_q:Query(查询)的注意头数量(如 32、64)。H_kv:Key(键)和 Value(值)的注意头数量(MHA 中H_kv = H_q,MQA 中H_kv = 1,GQA 中1 < H_kv < H_q)。d (Head Dimension):单个注意头的特征维度,满足d = D / H_q(如4096 / 32 = 128)。G (Group Size):GQA 中的分组大小,满足G = H_q / H_kv(即每个 KV 头服务多少个 Query 头)。
二、 基础:缩放点积注意力(Scaled Dot-Product Attention)
无论注意力结构如何演进,单个头内部的基础计算公式始终是缩放点积注意力:
# 缩放点积注意力计算
Attention(Q, K, V) = torch.softmax((Q @ K.transpose(-2, -1)) / math.sqrt(d), dim=-1) @ V[Query 向量] ──┐
├─► [点积相似度 Score = Q × Kᵀ] ─► [/ √d 缩放] ─► [Softmax 归一化] ─► [权重矩阵 A] ──┐
[Key 向量] ──┘ │
├─► [输出 Output]
[Value 向量] ─────────────────────────────────────────────────────────────────────────────────────┘
(矩阵乘法 A × V,加权汇总)- 点积打分(
Q × Kᵀ):衡量 Query 和 Key 在特征向量空间中的方向重合度(内积越大,语义相关性越高)。 - 缩放因子(
1 / √d):在特征维度d较大时,两随机向量点积的方差会线性放大至d。除以√d可将方差重置为 1,防止 Softmax 函数因数值过大进入极小梯度的饱和区,避免梯度消失。 - Softmax 归一化:在最后一个维度上求和为 1,转化为概率分布。
- 加权求和(
× V):利用计算出的注意力权重对 Value 矩阵进行线性组合,提取融合上下文的新表示。
三、 多头注意力(Multi-Head Attention, MHA)
1. 架构核心思想
标准 MHA 赋予了每个注意头完全独立的一套 (Q, K, V) 投影空间。H_q = H_kv = H,即 Query、Key、Value 的头数量严格为 1 : 1 : 1 对应。
┌── Head 1 = Attention(Q₁, K₁, V₁) ── (关注句法语义)
├── Head 2 = Attention(Q₂, K₂, V₂) ── (关注代词指代)
Input ──► Linear ───┼── ...
└── Head H = Attention(Q_H, K_H, V_H) ── (关注长程上下文)
│
▼
Concat(Head₁, ..., Head_H)
│
▼
Linear (W_O) ──► Output2. MHA 训练期张量计算全流程与行列元素深度解读
在模型训练阶段(反向传播的前向计算),输入为一批完整的上下文序列 X。以下拆解每一步张量的计算过程及微观语义:
步骤 1:输入张量 X
# 输入张量定义
X: torch.Tensor # 形状: [B, S, D]- 轴 0 (
B):批次样本索引b ∈ [0, B-1]。 - 轴 1(行 Row):序列中的 Token 位置
i ∈ [0, S-1](当前句子中的第i个词元)。 - 轴 2(列 Col):隐藏层特征维度
j ∈ [0, D-1](词嵌入维度的特征通道)。 - 元素含义
X[b, i, j]:第b个样本中第i个 Token 在第j个隐藏特征通道上的激活数值。
步骤 2:线性投影(Linear Projection)
# 线性投影参数与前向计算
W_Q = nn.Linear(D, D, bias=False) # 权重形状: [D, D]
W_K = nn.Linear(D, D, bias=False) # 权重形状: [D, D]
W_V = nn.Linear(D, D, bias=False) # 权重形状: [D, D]
Q = W_Q(X) # 形状: [B, S, D]
K = W_K(X) # 形状: [B, S, D]
V = W_V(X) # 形状: [B, S, D]- 权重矩阵行与列:行代表输入特征维度(
in_features = D),列代表输出特征维度(out_features = D)。 - 元素含义
W[r, c]:从输入特征第r维映射到输出特征第c维的可学习线性变换权重。 - 投影后张量:得到的
Q, K, V仍为[B, S, D]形状。行仍然是 Token 位置i,列是投影后的D维特征。
步骤 3:维度重排与多头拆分(Reshape & Transpose)
# 拆分出 H 个注意头,并转置以便进行并行批矩阵乘
Q = Q.view(B, S, H, d).transpose(1, 2) # 形状: [B, H, S, d]
K = K.view(B, S, H, d).transpose(1, 2) # 形状: [B, H, S, d]
V = V.view(B, S, H, d).transpose(1, 2) # 形状: [B, H, S, d]- 轴 0 (
B):Batch 样本索引。 - 轴 1 (
H):注意头索引h ∈ [0, H-1](代表第h个独立的表征子空间)。 - 轴 2(行 Row):序列位置
i ∈ [0, S-1](发出检索意图的 Tokeni)。 - 轴 3(列 Col):头内特征维度
k ∈ [0, d-1]。 - 元素含义
Q[b, h, i, k]:样本b在第h个头中,Tokeni发出的第k维查询特征分量。
步骤 4:批矩阵乘计算相关性打分(Batched Matmul)
# 计算未归一化的相似度得分矩阵
# 矩阵乘法: [B, H, S, d] @ [B, H, d, S] ──► [B, H, S, S]
Scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d)- 轴 2(行 Row, Query 视角):发起检索的目标 Token
i(i ∈ [0, S-1])。 - 轴 3(列 Col, Key 视角):被检索匹配的上下文 Token
j(j ∈ [0, S-1])。 - 元素含义
Scores[b, h, i, j]:样本b在头h中,Tokeni的 Query 向量与 Tokenj的 Key 向量的点积标量值。数值越大,代表在子空间h中 Tokeni认为 Tokenj与自己越相关。
步骤 5:因果掩码与 Softmax 归一化(Masking & Softmax)
# 因果掩码 (防止未来信息泄露) 与 Softmax
if mask is not None:
Scores = Scores.masked_fill(mask == 0, float('-inf'))
Attn_Weights = torch.softmax(Scores, dim=-1) # 形状: [B, H, S, S]- 行与列:行依然是当前词
i,列是上下文词j。 - 元素含义
Attn_Weights[b, h, i, j]:归一化后的注意力权重(概率标量,满足0 ≤ α ≤ 1)。它表示 Tokeni在当前头h中,分配给 Tokenj的注意力比例。 - 核心性质:矩阵中每一行的所有列元素之和严格等于 1(即
Σ_j Attn_Weights[b, h, i, j] = 1)。
步骤 6:上下文向量加权汇总(Context Aggregation)
# 利用注意力权重对 Value 进行加权汇聚
# 矩阵乘法: [B, H, S, S] @ [B, H, S, d] ──► [B, H, S, d]
Context = torch.matmul(Attn_Weights, V)- 行(轴 2):Token 位置
i。 - 列(轴 3):头内特征维度
k ∈ [0, d-1]。 - 元素含义
Context[b, h, i, k]:Tokeni在头h中聚合了全序列上下文后的新特征向量的第k维分量。其数学表达式为:Context[b, h, i, k] = Σ_j (Attn_Weights[b, h, i, j] × V[b, h, j, k])
步骤 7:多头拼接与最终输出投影(Concat & Output Projection)
# 拼接多头并过输出线性层
Context = Context.transpose(1, 2).contiguous().view(B, S, D) # 形状: [B, S, D]
W_O = nn.Linear(D, D, bias=False) # 形状: [D, D]
Output = W_O(Context) # 形状: [B, S, D]- 行(Row):Token 位置
i。 - 列(Col):最终融合了
H个子空间所有信息并经W_O线性混合后的D维特征。 - 元素含义
Output[b, i, j]:Tokeni在经过当前 MHA 层处理后的最终特征表示。
3. MHA 的优缺点深度剖析
优势(Pros)
- 表征能力最强(Rich Representation):每个头有独立的 Key 和 Value 参数,能够将序列映射到
H个互不干扰的子空间中,分别捕捉不同的语法结构、位置关联与深层语义。 - 训练收敛稳定:参数自由度高,在大规模预训练时能够稳定拟合复杂的语义模式。
劣势(Cons)
- 推理时 KV Cache 显存爆炸(Capacity Wall):
- 在自回归逐 Token 生成的 Decode 阶段,每一层都需要为当前所有并发请求保存所有历史 Token 的独立 Key 和 Value。
- 单层单 Token 的 KV Cache 尺寸为
2 × H × d = 2 × D个元素。 - 对于 70B 模型(80 层,
D = 8192,FP16),并发 8 个 32K 上下文的请求,KV Cache 显存高达1.37 TB,单机多卡根本无法容纳。
- 显存带宽吞吐严重受限(Bandwidth Wall):
- Decode 阶段每次只生成 1 个 Token,算力需求极小,但必须从 GPU 显存(HBM)中读出全量庞大的 KV Cache。
- 访存比(Arithmetic Intensity)极低,计算核心大量时间处于等待读取数据的空闲状态(Memory-Bound)。
四、 多查询注意力(Multi-Query Attention, MQA)
1. 架构核心思想
为了彻底解决 MHA 在推理时的 KV Cache 爆炸,Noam Shazeer 于 2019 年提出了激进的 MQA。
核心思想:保持全部 H_q 个独立的 Query 头不变,但所有 Query 头共同共享唯独 1 个 Key 头和 1 个 Value 头(H_kv = 1)。
Query Heads (H_q 个独立头) Shared Key / Value (全局仅 1 组)
┌─────┬─────┬─────┬─────┬─────┬─────┐ ┌───────────────┐
│ Q₁ │ Q₂ │ Q₃ │ Q₄ │ ... │ Q_H │ │ K_all │
└─────┴─────┴─────┴─────┴─────┴─────┘ └───────────────┘
│ │ │ │ │ │ ┌───────────────┐
│ │ │ │ │ │ │ V_all │
▼ ▼ ▼ ▼ ▼ ▼ └───────────────┘
┌───────────────────────────────────┐ │
│ 每个 Q_i 与唯一的 K_all 分别打分 │◄─────────────────────┘
└───────────────────────────────────┘2. MQA 训练期张量计算全流程与行列元素深度解读
在训练过程中,MQA 的 Key 和 Value 投影矩阵参数量大幅缩减,并通过广播机制完成多头计算:
步骤 1 & 2:输入与参数投影
# MQA 投影层定义
W_Q = nn.Linear(D, D, bias=False) # 形状: [D, D] (输出 H_q * d)
W_K = nn.Linear(D, d, bias=False) # 形状: [D, d] (仅投影 1 个头维度 d)
W_V = nn.Linear(D, d, bias=False) # 形状: [D, d] (仅投影 1 个头维度 d)
Q = W_Q(X) # 形状: [B, S, D]
K = W_K(X) # 形状: [B, S, d] (列维度缩减 H_q 倍)
V = W_V(X) # 形状: [B, S, d] (列维度缩减 H_q 倍)- 权重矩阵参数量对比:
W_K与W_V的参数量从D × D直接降至D × d,仅为 MHA 的1 / H_q。 - 元素含义
K[b, j, k]:样本b中 Tokenj投影到唯一的共享 Key 空间中的第k维特征(0 ≤ k < d)。
步骤 3 & 4:维度重排与广播点积计算
# 维度变换
Q = Q.view(B, S, H_q, d).transpose(1, 2) # 形状: [B, H_q, S, d]
K = K.view(B, S, 1, d).transpose(1, 2) # 形状: [B, 1, S, d]
V = V.view(B, S, 1, d).transpose(1, 2) # 形状: [B, 1, S, d]
# 批矩阵乘 (自动触发 Head 维度广播)
# [B, H_q, S, d] @ [B, 1, d, S] ──► [B, H_q, S, S]
Scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d)- 广播机制(Broadcasting):由于
K在轴 1(Head 维度)大小为 1,PyTorch 会在逻辑上将单份 Key 矩阵复用给所有的 Query 头。 - 元素含义
Scores[b, h, i, j]:样本b中,头h独有的 Query 向量Q[b, h, i, :]与全局共享的 Key 向量K[b, 0, j, :]的内积相似度。
步骤 5 & 6:Softmax 归一化与共享 Value 汇聚
Attn_Weights = torch.softmax(Scores + Mask, dim=-1) # 形状: [B, H_q, S, S]
# 与共享的 V [B, 1, S, d] 广播相乘
Context = torch.matmul(Attn_Weights, V) # 形状: [B, H_q, S, d]- 微观机制:虽然
K和V是全局单组共享的,但由于每个 Query 头Q_h的投影参数不同,计算出的Attn_Weights[b, h, :, :]在每个头上依然是独立的。 - 乘以单组共享的
V ∈ ℝ^(B × 1 × S × d)时再次触发广播,得到每个头专属的Context[b, h, S, d]。
3. MQA 的优缺点深入剖析
优势(Pros)
- KV Cache 显存断崖式缩减:
- 每层每个 Token 只需要存储
2 × 1 × d = 2d个数值。 - 相比 MHA,KV Cache 容量直接缩减为原来的
1 / H_q(例如 32 个头的模型,KV 显存缩减至1/32 ≈ 3.125%)。
- 每层每个 Token 只需要存储
- 推理显存带宽压力暴降,吞吐翻倍:
- Decode 阶段 GPU 仅需读取极小量的全局 KV 向量,访存瓶颈大幅缓解,Batch Size 可扩大数倍。
- 投影权重参数量缩减:
W_K与W_V的参数量减少了(1 - 1/H_q),降低了静态权重显存。
劣势(Cons)
- 模型表达能力受损(Capacity Drop):
- 所有的 Query 头被迫使用完全相同的一套 Key 空间打分和 Value 空间聚合,失去了“多角度提取不同语义特征”的自由度。
- 训练稳定性与下游任务微调精度下滑:
- 实证研究表明,在知识检索密集型或复杂推理任务中,MQA 相比标准 MHA 存在不可忽视的性能掉点(Accuracy Loss)。
五、 分组查询注意力(Grouped-Query Attention, GQA)
1. 架构核心思想
为了结合 MHA 的高表达能力 与 MQA 的极高推理效率,Ainslie 等人在 2023 年提出了 GQA。
核心思想:将 H_q 个 Query 头划分为 H_kv 个组(Groups),组内共享 1 对 Key 和 Value 头。
- 设组内大小为
G = H_q / H_kv。 - 极端情况 1:当
H_kv = H_q(G = 1)时,退化为 MHA; - 极端情况 2:当
H_kv = 1(G = H_q)时,退化为 MQA; - 典型设定:
H_q = 32, H_kv = 8(G = 4)或H_q = 64, H_kv = 8(G = 8)。
Group 1 (共享 K₁, V₁) Group 2 (共享 K₂, V₂) ...
┌─────────┬─────────┬─────────┐ ┌─────────┬─────────┬─────────┐
│ Q₁ │ Q₂ │ Q₃ │ │ Q₄ │ Q₅ │ Q₆ │
└────┬────┴────┬────┴────┬────┘ └────┬────┴────┬────┴────┬────┘
│ │ │ │ │ │
└─────────┼─────────┘ └─────────┼─────────┘
▼ ▼
[与 K₁ 打分, 乘 V₁] [与 K₂ 打分, 乘 V₂]2. GQA 训练期张量计算全流程与行列元素深度解读
以标准的 H_q = 32, H_kv = 8, G = 4(32 个 Q 头,8 个 KV 组,每组 4 个 Q 头)为例拆解训练张量流:
步骤 1 & 2:输入与参数投影
# GQA 投影层定义 (以 32 个 Q 头, 8 个 KV 头为例)
W_Q = nn.Linear(D, D, bias=False) # 形状: [D, D] (输出 32 * d)
W_K = nn.Linear(D, H_kv * d, bias=False) # 形状: [D, 8 * d] (输出 D/4)
W_V = nn.Linear(D, H_kv * d, bias=False) # 形状: [D, 8 * d] (输出 D/4)
Q = W_Q(X) # 形状: [B, S, 32 * d]
K = W_K(X) # 形状: [B, S, 8 * d]
V = W_V(X) # 形状: [B, S, 8 * d]- 参数量:
W_K和W_V的输出维度为H_kv × d = 8 × 128 = 1024,参数量为 MHA 的1/4(即25%)。
步骤 3 & 4:重排、组内扩展与打分计算
# 1. 拆出头维度
Q = Q.view(B, S, H_q, d).transpose(1, 2) # 形状: [B, 32, S, d]
K = K.view(B, S, H_kv, d).transpose(1, 2) # 形状: [B, 8, S, d]
V = V.view(B, S, H_kv, d).transpose(1, 2) # 形状: [B, 8, S, d]
# 2. 组内对齐扩展 (方式 A: repeat_interleave 复制 4 次)
K_expanded = K.repeat_interleave(G, dim=1) # 形状: [B, 32, S, d]
V_expanded = V.repeat_interleave(G, dim=1) # 形状: [B, 32, S, d]
# 3. 标准点积打分
# [B, 32, S, d] @ [B, 32, d, S] ──► [B, 32, S, S]
Scores = torch.matmul(Q, K_expanded.transpose(-2, -1)) / math.sqrt(d)- 组对齐逻辑:
K中的第 0 个组头被复制分配给Q的第 0, 1, 2, 3 个头;第 1 个组头分配给第 4, 5, 6, 7 个头,依此类推。 - 元素含义
Scores[b, h, i, j]:属于第g = floor(h / G)组的 Query 头h发出的向量Q[b, h, i, :]与该组专属的 Key 向量K[b, g, j, :]的内积相似度。
步骤 5 & 6:Softmax 归一化与加权求和
Attn_Weights = torch.softmax(Scores + Mask, dim=-1) # 形状: [B, 32, S, S]
Context = torch.matmul(Attn_Weights, V_expanded) # 形状: [B, 32, S, d]
# 输出投影
Context = Context.transpose(1, 2).contiguous().view(B, S, D) # 形状: [B, S, D]
Output = W_O(Context) # 形状: [B, S, D]- 表征能力保护:虽然同组内的 4 个 Query 头共享相同的 Key/Value,但由于各自打出的
Attn_Weights独立,依然能汇聚出 32 个具有丰富差异的注意力子空间。
3. GQA 的优缺点深入剖析
优势(Pros)
- 性能无损逼近 MHA(Gold Standard):
- 实验表明,将 64 个头压缩为 8 个 KV 组(
G = 8)时,模型在各类常识问答、代码生成、数学推理等基准测试中的得分与标准 MHA 几乎完全无差。
- 实验表明,将 64 个头压缩为 8 个 KV 组(
- 大幅削减 KV Cache 显存与带宽:
- KV Cache 显存缩减至 MHA 的
H_kv / H_q = 1 / G。 - 例如
H_q = 64, H_kv = 8时,显存占用降低 87.5%(仅剩 1/8),使大模型能够在单卡上承载 32K~128K 超长上下文。
- KV Cache 显存缩减至 MHA 的
- 支持 Uptraining(热启动转换):
- 已经预训练好的 MHA 模型,可以通过对同一个组内的
W_K和W_V权重矩阵进行均值池化(Mean Pooling),仅需极少量的微调(约 5% 的原始预训练计算量)就能快速转换为 GQA 模型。
- 已经预训练好的 MHA 模型,可以通过对同一个组内的
劣势(Cons)
- 相比 MQA 仍有一定显存开销:
- 相比 MQA 极简的单一 KV 头,GQA 仍保留了数个组(如 8 组),在极端长序列与极端超高并发下,显存压力虽大为减小但依然存在。
- 组数超参数
G的权衡选型:G设得过大(趋近 MQA)可能影响精度,G设得过小(趋近 MHA)显存收益不足,需根据部署显卡的实际规格(如 24GB、80GB)进行工程权衡。
六、 MHA、MQA 与 GQA 三者全景对比
| 核心对比维度 | 多头注意力 (MHA) | 多查询注意力 (MQA) | 分组查询注意力 (GQA) |
|---|---|---|---|
键值头数量 (H_kv) | H_q (如 32) | 1 (固定为 1) | H_kv 组 (如 8) |
头数配比 (Q : K : V) | 1 : 1 : 1 | H_q : 1 : 1 | (H_q/H_kv) : 1 : 1 |
| W_K, W_V 参数量 | 2 × D × D (100%) | 2 × D × (D/H_q) (~3%) | 2 × D × (H_kv × d) (~25%) |
| 训练期 Q 张量形状 | [B, H_q, S, d] | [B, H_q, S, d] | [B, H_q, S, d] |
| 训练期 K/V 张量形状 | [B, H_q, S, d] | [B, 1, S, d] | [B, H_kv, S, d] |
| 单 Token KV Cache 显存 | 2 × H_q × d (基准) | 2 × 1 × d (1/H_q) | 2 × H_kv × d (H_kv/H_q) |
| 70B 模型单 Token KV 开销 | ~5.24 MB (100%) | ~0.16 MB (3.1%) | ~0.65 MB (12.5%) |
| 显存带宽读取压力 | 极高 (严重受限) | 极低 | 低 |
| 模型特征表征能力 | 最高 (基准) | 略有下降 (复杂任务有损失) | 几乎无损逼近 MHA |
| 工业界主流代表 | 原始 Transformer, GPT-3 | Falcon-40B, PaLM | LLaMA-2/3, Mistral, Qwen2 |
七、 总结与工程选型建议
从 MHA 到 MQA,再到 GQA 的演进,并非简单的参数增减,而是在算法表征精度与**物理硬件瓶颈(显存与带宽)**之间寻找 Pareto 最优解的经典案例:
- MHA 是最纯粹的多子空间表征,但在大模型自回归解码场景下,高昂的 KV Cache 成本使其难以适应长上下文和高吞吐部署;
- MQA 走向了效率的极致,将 KV Cache 压缩到极限,但过强的共享约束损伤了部分表征能力;
- GQA 则以精妙的分组折中,在大幅卸去显存与带宽包袱(减少 75%~87.5%)的同时,几乎完整继承了 MHA 的强大性能,因而成为了现代开源大模型不可动摇的黄金标准。