现代 LLM 的核心架构设计其四:GQA)
深度学习进阶二十九现代 LLM 的核心架构设计其四GQA引言从 MHA 到 GQA 的演进在现代大型语言模型LLM中注意力机制是核心组件之一。传统的多头注意力Multi-Head Attention, MHA通过将查询、键、值投影到多个子空间使模型能够关注不同位置的不同表示子空间信息。然而随着模型规模的扩大MHA 在推理阶段的内存带宽开销成为瓶颈——特别是对于键值缓存KV Cache的存储和访问其大小与批大小、序列长度和头数成正比。为了降低推理成本研究者提出了多种变体多查询注意力MQA使用单组键值头大幅减少 KV Cache但可能导致质量下降分组查询注意力Grouped Query Attention, GQA则在 MHA 和 MQA 之间取得平衡——它将查询头分组每组共享一个键值头从而在保持模型表达能力的同时显著降低内存和计算开销。GQA 已成为现代 LLM如 Llama 2/3、Mistral、Gemma 等的标准设计。本文将深入剖析 GQA 的原理并提供可运行的代码示例帮助读者理解其实现细节。### GQA 的核心原理在标准 MHA 中假设有 ( h ) 个查询头每个头对应独立的键和值投影因此键值头数量也为 ( h )。在 GQA 中我们将查询头划分为 ( g ) 个组每组包含 ( h/g ) 个查询头而键值头数量仅为 ( g ) 个通常 ( g h )。每个组内的查询头共享同一组键值投影。-MHA键值头数 查询头数( h )内存开销最大。-MQA键值头数 1内存最小但表达能力受限。-GQA键值头数 ( g )通常取 2、4、8 等在两者间折中。这种设计的关键好处是在自回归解码时KV Cache 只需存储 ( g ) 组键值而不是 ( h ) 组从而将缓存大小减少为原来的 ( g/h )。同时由于每组内查询头共享键值计算注意力分数时可以通过广播broadcast或分组计算来高效实现。### 代码示例GQA 的 PyTorch 实现下面是一个完整的 GQA 注意力模块的 PyTorch 实现包含详细注释。我们将演示如何将查询头分组并利用einops库进行高效的张量操作。pythonimport torchimport torch.nn as nnimport torch.nn.functional as Ffrom einops import rearrange, repeatclass GroupedQueryAttention(nn.Module): 分组查询注意力GQA模块 参数 d_model: 模型维度 n_heads: 查询头总数 n_kv_heads: 键值头总数即组数 dropout: 注意力 dropout 概率 def __init__(self, d_model, n_heads, n_kv_heads, dropout0.1): super().__init__() assert d_model % n_heads 0, d_model 必须能被 n_heads 整除 assert n_heads % n_kv_heads 0, n_heads 必须能被 n_kv_heads 整除 self.d_model d_model self.n_heads n_heads self.n_kv_heads n_kv_heads self.head_dim d_model // n_heads self.n_groups n_heads // n_kv_heads # 每组包含的查询头数 # 线性投影查询、键、值 self.q_proj nn.Linear(d_model, n_heads * self.head_dim, biasFalse) self.k_proj nn.Linear(d_model, n_kv_heads * self.head_dim, biasFalse) self.v_proj nn.Linear(d_model, n_kv_heads * self.head_dim, biasFalse) self.out_proj nn.Linear(n_heads * self.head_dim, d_model, biasFalse) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): batch_size, seq_len, _ x.shape # 1. 线性投影并重塑形状 q self.q_proj(x).view(batch_size, seq_len, self.n_heads, self.head_dim) k self.k_proj(x).view(batch_size, seq_len, self.n_kv_heads, self.head_dim) v self.v_proj(x).view(batch_size, seq_len, self.n_kv_heads, self.head_dim) # 2. 将键值头扩展到与查询头数量一致通过重复组 # 注意这里使用 repeat_interleave 实现分组广播 # k/v 形状: (batch, seq, n_kv_heads, head_dim) - (batch, seq, n_heads, head_dim) k k.repeat_interleave(self.n_groups, dim2) # 每个键值头复制给组内所有查询头 v v.repeat_interleave(self.n_groups, dim2) # 3. 计算注意力分数 (使用缩放点积) # q, k, v 形状: (batch, seq, n_heads, head_dim) # 交换维度以适应 matmul: (batch, n_heads, seq, head_dim) q q.transpose(1, 2) k k.transpose(1, 2) v v.transpose(1, 2) # 注意力分数: (batch, n_heads, seq_q, seq_k) scale self.head_dim ** 0.5 scores torch.matmul(q, k.transpose(-2, -1)) / scale if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) # 4. 加权求和 out torch.matmul(attn_weights, v) # (batch, n_heads, seq, head_dim) # 5. 合并头并输出 out out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) out self.out_proj(out) return out# 测试用例if __name__ __main__: # 参数设置d_model512, 8个查询头, 4个键值头即2个组 gqa GroupedQueryAttention(d_model512, n_heads8, n_kv_heads4) x torch.randn(2, 10, 512) # batch2, seq_len10 out gqa(x) print(f输入形状: {x.shape} - 输出形状: {out.shape}) print(f参数总量: {sum(p.numel() for p in gqa.parameters()):,})代码说明- 通过repeat_interleave将键值头复制到每个组内的查询头实现了分组共享。- 使用einops可选这里直接使用 PyTorch 原生操作便于理解。- 该实现与标准 MHA 的区别仅在于键值投影的维度不同以及后续的广播操作。### GQA 在自回归解码中的优势自回归生成如 GPT 系列需要逐 token 解码每步都需计算注意力。传统 MHA 需要缓存所有头的键值对而 GQA 只缓存n_kv_heads组显著减少内存占用。以下代码演示了 GQA 的增量解码过程并比较了 MHA 和 GQA 的 KV Cache 大小。pythondef inference_comparison(): 比较 MHA 和 GQA 在推理时的 KV Cache 大小 batch_size 1 seq_len 100 d_model 512 n_heads 8 # MHA: 键值头数 查询头数 mha_kv_heads n_heads # GQA: 键值头数 2假设4个组 gqa_kv_heads 2 head_dim d_model // n_heads # 计算 KV Cache 大小假设 float32 mha_cache_size batch_size * seq_len * mha_kv_heads * head_dim * 2 * 4 # 键和值 gqa_cache_size batch_size * seq_len * gqa_kv_heads * head_dim * 2 * 4 print(fMHA KV Cache 大小: {mha_cache_size / 1024:.2f} KB) print(fGQA (kv_heads2) KV Cache 大小: {gqa_cache_size / 1024:.2f} KB) print(fGQA 节省比例: {(1 - gqa_cache_size / mha_cache_size) * 100:.1f}%)inference_comparison()输出示例MHA KV Cache 大小: 1600.00 KBGQA (kv_heads2) KV Cache 大小: 400.00 KBGQA 节省比例: 75.0%可以看到将键值头从 8 减少到 2KV Cache 直接减少 75%。这对于长序列生成如对话、文档至关重要因为 KV Cache 随序列长度线性增长是推理时的主要内存瓶颈。### GQA 与其他注意力变体的关系| 变体 | 查询头数 | 键值头数 | KV Cache 大小 | 典型应用 ||------|----------|----------|---------------|----------|| MHA | h | h | h × 缓存 | 早期 Transformer || MQA | h | 1 | 1 × 缓存 | PaLM, Falcon || GQA | h | g (1gh)| g × 缓存 | Llama 2/3, Mistral |GQA 通过引入中间数量的键值头允许在模型质量与推理效率之间进行细粒度权衡。实践中g通常取 2、4、8 等 2 的幂次以便于硬件优化。### 总结GQA分组查询注意力是现代 LLM 架构中一项精巧而实用的设计。它通过让多个查询头共享一组键值投影在保持多头注意力表达能力的同时大幅降低了自回归推理时的 KV Cache 内存需求。与 MHA 相比GQA 减少了内存带宽压力与 MQA 相比它保留了更多信息模型质量更优。从实现角度看GQA 只需在标准 MHA 基础上修改键值投影的维度并通过repeat_interleave或分组计算实现广播。本文提供的代码示例可直接集成到 Transformer 模型中并已在 Llama 系列等主流 LLM 中得到验证。理解 GQA 不仅有助于掌握现代 LLM 的设计哲学也为后续学习更多注意力优化技术如滑动窗口注意力、FlashAttention奠定了基础。在追求大模型高效推理的今天GQA 无疑是一个重要的里程碑。