Grouped Query AttentionGQA原理与 PyTorch 实现从多头注意力到 KV-Cache 显存优化【免费下载链接】leetcodeLeetcode solutions项目地址: https://gitcode.com/GitHub_Trending/leetcode1/leetcode分组查询注意力Grouped Query AttentionGQA是当前大模型推理优化中最核心的注意力变体之一它通过让一组查询头共享同一份 K、V 投影在几乎不损失模型质量的前提下大幅压缩 KV-Cache 的内存占用。本指南基于 NeetCode LeetCode 仓库中的 grouped-query-attention.md 系列教程系统讲解 GQA 的数学动机、完整 PyTorch 实现、维度演算与常见实现陷阱读完后你将能够独立实现一个可运行、可替换标准多头注意力的 GQA 模块并理解它在 Llama、Mistral 等生产模型中的落地方式。前置知识在动手实现 GQA 之前需要先掌握三个基础模块它们在本仓库的 articles 系列中都有对应的独立章节多头自注意力Multi-Head Self-AttentionGQA 的本质是修改多头注意力中各个头共享 K、V 的方式因此必须先理解标准 MHA 中每个查询头都有独立 K、V 投影的结构。可参考 multi-headed-self-attention.md。KV-CacheGQA 的首要动机就是降低推理阶段 KV-Cache 的内存开销因此理解缓存大小的来源问题缓存随序列长度线性增长、随头数线性增长是理解 GQA 的前提。可参考 kv-cache.md。张量形状重塑Tensor Reshaping实现 GQA 需要反复使用view、transpose、repeat_interleave在(B, T, D)与(B, heads, T, head_dim)两种格式之间切换这三者缺一不可。核心概念为什么需要共享 K、V标准多头注意力会给每一个查询头分配独立的 K 和 V 投影。在带 KV-Cache 的自回归推理过程中所有头的 K、V 张量都必须被缓存起来而且缓存的内存开销随头数线性增长头越多每个 token 需要缓存的 K、V 就越多长上下文的推理成本就越高。GQA 的思路非常直接让多个查询头共享同一组 K、V。假设模型有 $h$ 个查询头query heads和 $g$ 个 KV 头KV heads那么每个 KV 头负责服务 $h/g$ 个查询头。这样 K、V 投影层的输出维度从 $h \cdot d_h$ 降到 $g \cdot d_h$缓存中每个 token 存储的 K、V 也随之按比例缩减。其中最关键的操作是repeat_interleave它把 $g$ 个 KV 头按序各自重复 $h/g$ 次扩展成 $h$ 个从而让 K、V 的形状与 Q 对齐可以直接套用标准的缩放点积注意力数学公式。GQA 是一个统一的框架两个极端情况分别是当 $g h$ 时每个 KV 头只服务一个查询头GQA 退化为标准的 MHA当 $g 1$ 时所有查询头共享同一个 K、V退化为多查询注意力Multi-Query AttentionMQA。解决方案直觉理解把输入 $x$ 用num_heads个头投影成 Q用num_kv_heads个头投影成 K 和 V。先把三者都重塑成多头格式(B, heads, T, head_dim)然后通过重复每个 KV 头把 K、V 的头数扩展到与 Q 一致从这一步开始剩下的就是带因果掩码的标准缩放点积注意力最后把所有头的输出拼接起来经过输出投影层得到最终结果。完整实现import torch import torch.nn as nn from torchtyping import TensorType class GroupedQueryAttention(nn.Module): def __init__(self, model_dim: int, num_heads: int, num_kv_heads: int): super().__init__() torch.manual_seed(0) self.num_heads num_heads self.num_kv_heads num_kv_heads self.head_dim model_dim // num_heads self.q_proj nn.Linear(model_dim, num_heads * self.head_dim, biasFalse) self.k_proj nn.Linear(model_dim, num_kv_heads * self.head_dim, biasFalse) self.v_proj nn.Linear(model_dim, num_kv_heads * self.head_dim, biasFalse) self.output_proj nn.Linear(num_heads * self.head_dim, model_dim, biasFalse) def forward(self, x: TensorType[float]) - TensorType[float]: B, T, D x.shape # Project to Q, K, V and reshape into heads q self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) k self.k_proj(x).view(B, T, self.num_kv_heads, self.head_dim).transpose(1, 2) v self.v_proj(x).view(B, T, self.num_kv_heads, self.head_dim).transpose(1, 2) # Expand K, V to match Qs num_heads by repeating each KV head repeats self.num_heads // self.num_kv_heads k k.repeat_interleave(repeats, dim1) v v.repeat_interleave(repeats, dim1) # Scaled dot-product attention with causal mask scores (q k.transpose(-2, -1)) * (self.head_dim ** -0.5) mask torch.tril(torch.ones(T, T, devicex.device)) scores scores.masked_fill(mask 0, float(-inf)) weights torch.softmax(scores, dim-1) out (weights v).transpose(1, 2).contiguous().view(B, T, -1) return torch.round(self.output_proj(out), decimals4)逐行拆解这段实现中的关键点投影层的不对称设计q_proj输出num_heads * head_dim维而k_proj、v_proj只输出num_kv_heads * head_dim维。这是 GQA 与 MHA 在结构上的唯一区别也是全部显存收益的来源——K、V 相关的权重矩阵和缓存张量都变小了。repeat_interleave(repeats, dim1)在头维度dim1上把每个 KV 头重复repeats num_heads // num_kv_heads次使 K、V 与 Q 的头数一致。注意这里要求num_heads能被num_kv_heads整除。因果掩码GQA 服务于 decoder-only 的自回归模型torch.tril生成下三角掩码将未来位置置为-inf保证 token 只能关注自身及之前的 token。输出投影拼接后的(B, T, num_heads * head_dim)经过output_proj映射回model_dim与 MHA 的输入输出形状完全一致因此可以作为 MHA 的即插即用替代品。逐步演算以model_dim8、num_heads4、num_kv_heads2为例逐步跟踪每一步的形状变化步骤操作形状Q 投影q_proj(x)后重塑(B, 4, T, 2)— 4 个查询头head_dim2K 投影k_proj(x)后重塑(B, 2, T, 2)— 只有 2 个 KV 头V 投影v_proj(x)后重塑(B, 2, T, 2)— 只有 2 个 KV 头扩展 Krepeat_interleave(2, dim1)(B, 4, T, 2)— 每个 KV 头重复 2 次扩展 Vrepeat_interleave(2, dim1)(B, 4, T, 2)— 此时与 Q 对齐注意力标准缩放点积注意力每个头输出(B, 4, T, 2)拼接 投影合并所有头输出投影(B, T, 8)从这张表可以直观看到K、V 在投影阶段只有 2 个头直到进入注意力计算前才被复制成 4 个头。复制的是张量数据而不是重新计算投影因此省下的不是推理时间而是投影层的参数规模和 KV-Cache 的存储空间。显存节省分析每 token 的缓存开销对比MHA 每个 token 需要缓存 $h \cdot d_h$ 个 K 值和 $h \cdot d_h$ 个 V 值GQA 每个 token 只需要缓存 $g \cdot d_h$ 个 K 值和 $g \cdot d_h$ 个 V 值。在上述例子中K、V 投影每个 token 只产生num_kv_heads * head_dim 2 * 2 4个值而不是 MHA 的num_heads * head_dim 4 * 2 8个值。在 KV-Cache 中这相当于每一层注意力层的缓存内存减半。把这个比例放大到真实模型以 Llama 2 70B 为例它有 64 个查询头和 8 个 KV 头因此 KV-Cache 的内存节省为 8 倍。这也是为什么 Llama 2 70B 这类超大规模模型能在有限的 GPU 显存中支撑更长的上下文——每一层、每一个 token 的 K、V 缓存都被压缩了 8 倍。时间与空间复杂度时间$O(T^2 \cdot d)$。注意力计算本身与 MHA 完全相同K、V 在进入注意力之前已经完成扩展因此时间复杂度不变。空间KV-Cache 存储为 $O(g \cdot T \cdot d_h)$其中 $g$ 是 KV 头数量$d_h$ 是每个头的维度。相比标准 MHA 的 $O(h \cdot T \cdot d_h)$缓存空间缩减为原来的 $g/h$。值得注意的是GQA 的收益完全来自 KV-Cache 空间和投影层参数量的缩减而不改变注意力计算的时间复杂度——扩展后的注意力矩阵大小与 MHA 完全相同。常见陷阱误用repeat而不是repeat_interleaverepeat会整体平铺整个张量而repeat_interleave会逐个元素地重复。当有 2 个 KV 头、4 个查询头时我们需要的是[KV0, KV0, KV1, KV1]每个 KV 头连续地服务自己的一组查询头而不是[KV0, KV1, KV0, KV1]交叉排列。# 错误repeat 平铺整个张量 k k.repeat(1, repeats, 1, 1) # [KV0, KV1, KV0, KV1] -- 分组错误 # 正确repeat_interleave 逐个重复每个元素 k k.repeat_interleave(repeats, dim1) # [KV0, KV0, KV1, KV1] -- 正确的分组两种写法产生的张量形状相同但分组语义完全不同。交叉排列会让每个查询头错误地混合来自不同 KV 头的信息破坏一组查询头共享同一个 KV 头的核心假设。忘记因果掩码GQA 用于 decoder-only 模型如 GPT、Llamatoken 不得关注未来的位置。缺少下三角掩码模型就会破坏自回归生成的性质——在训练时每个位置都能偷看到后面的 token生成时行为就会错乱。# 错误没有因果掩码 weights torch.softmax(scores, dim-1) # 正确在 softmax 之前施加因果掩码 mask torch.tril(torch.ones(T, T, devicex.device)) scores scores.masked_fill(mask 0, float(-inf)) weights torch.softmax(scores, dim-1)重塑时用错头数view/transpose的次序和参数都很关键。Q 有num_heads个头但 K、V 只有num_kv_heads个头。如果 K 的重塑错误地使用num_heads会导致 K 的形状与k_proj的输出维度不匹配shape mismatch或者静默地产生错误维度。# 错误用 num_heads 重塑 K k self.k_proj(x).view(B, T, self.num_heads, self.head_dim) # 形状不匹配 # 正确用 num_kv_heads 重塑 K k self.k_proj(x).view(B, T, self.num_kv_heads, self.head_dim)在 GPT 项目中的落地GQA 是 Llama 2/3、Mistral、Gemma 等主流开源模型采用的注意力变体。在一个完整的 GPT 实现中GQA 需要与上一节的 KV-Cache 结合使用缓存中每一层只存储num_kv_heads份的 K、V而不是num_heads份从而获得与头数比值成正比的内存缩减。本仓库的 kv-cache.md 详细实现了KVCache类与CachedAttention模块——缓存沿序列维度用torch.cat增长配合 GQA 后每次生成新 token 时只需要投影新 token 的 K、V再追加到缓存中同时由于缓存天然只包含过去 token因果性由构造保证甚至不再需要显式的掩码。在模型整体架构中GQA 模块会替换 transformer-block.md 中描述的多头注意力子层与 LayerNorm、残差连接和前馈网络共同构成 transformer block多个 block 堆叠即构成完整的 decoder-only GPT 模型。关键要点GQA 让多个查询头共享 K、V将 KV-Cache 内存缩减为原来的num_heads / num_kv_heads分之一且质量损失极小。repeat_interleave是核心操作它把每个 KV 头扩展为服务其所属查询头组的多个副本使形状兼容标准注意力计算。GQA 统一了 MHA 与 MQA$g h$ 时是 MHA$g 1$ 时是 MQA。大多数生产模型取中间值例如 Llama 2 70B 使用 64 个查询头配 8 个 KV 头。扩展后的注意力计算与标准 MHA 完全一致节省全部来自更小的投影层和更小的缓存而非计算流程的改变。延伸阅读标准多头注意力的完整实现multi-headed-self-attention.mdKV-Cache 缓存机制与CachedAttention实现kv-cache.mdTransformer Block 架构注意力子层 FFN 残差transformer-block.md本仓库文章写作规范articles/README.md【免费下载链接】leetcodeLeetcode solutions项目地址: https://gitcode.com/GitHub_Trending/leetcode1/leetcode创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
