深度学习进阶(二十九)现代 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, dropout=0.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, bias=False) self.k_proj = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False) self.out_proj = nn.Linear(n_heads * self.head_dim, d_model, bias=False) self.dropout = nn.Dropout(dropout) def forward(self, x, mask=None): 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, dim=2) # 每个键值头复制给组内所有查询头 v = v.repeat_interleave(self.n_groups, dim=2) # 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_model=512, 8个查询头, 4个键值头(即2个组) gqa = GroupedQueryAttention(d_model=512, n_heads=8, n_kv_heads=4) x = torch.randn(2, 10, 512) # batch=2, seq_len=10 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(f"MHA KV Cache 大小: {mha_cache_size / 1024:.2f} KB") print(f"GQA (kv_heads=2) KV Cache 大小: {gqa_cache_size / 1024:.2f} KB") print(f"GQA 节省比例: {(1 - gqa_cache_size / mha_cache_size) * 100:.1f}%")inference_comparison()输出示例MHA KV Cache 大小: 1600.00 KBGQA (kv_heads=2) KV Cache 大小: 400.00 KBGQA 节省比例: 75.0%可以看到,将键值头从 8 减少到 2,KV Cache 直接减少 75%。这对于长序列生成(如对话、文档)至关重要,因为 KV Cache 随序列长度线性增长,是推理时的主要内存瓶颈。### GQA 与其他注意力变体的关系| 变体 | 查询头数 | 键值头数 | KV Cache 大小 | 典型应用 ||------|----------|----------|---------------|----------|| MHA | h | h | h × 缓存 | 早期 Transformer || MQA | h | 1 | 1 × 缓存 | PaLM, Falcon || GQA | h | g (1<g<h)| 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 无疑是一个重要的里程碑。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐