资料摘要:personal chatgpt 21 — llama2 源码分析 GQA:Grouped Query Attention
本讲由真实视频转录与作者 notebook/代码交叉整理。原始层:🔒 带时间戳视频转录;课程总览:personal chatgpt — LLMs 实践系列。
本讲从 Llama 2 的 Attention 实现解释 GQA:查询头数量保持较多,而多个查询头共享一组 K/V 头,以减少 K/V 投影及推理时 KV Cache 的占用。课程以 Llama 2 70B 的 64 个 Q 头、8 个 KV 头为例,每组 8 个 Q 头共享一组 K/V;源码用 repeat_kv 把 K/V 在逻辑上扩展到与 Q 头数匹配后再做注意力。
1. MHA、MQA 与 GQA
Section titled “1. MHA、MQA 与 GQA”- MHA:Q、K、V 都有相同数量的独立头,表达力强但 KV Cache 最大。
- MQA:所有 Q 头共享唯一一组 K/V,缓存最小,但共享程度最高。
- GQA:在两者之间折中;Q 头被分组,每组共享一个 KV 头。
- 本讲强调 GQA 的目标是提升 attention,尤其是自回归推理阶段的效率;它与上一讲 KV Cache 是同一条优化链。
2. Llama 2 70B 的分组关系
Section titled “2. Llama 2 70B 的分组关系”课程示例:
n_heads = 64n_kv_heads = 8n_rep = n_heads // n_kv_heads = 8
因此第 1–8 个 Q 头共享第 1 个 K/V 头,第 9–16 个 Q 头共享第 2 个 K/V 头,以此类推。转录把它口述为“每一组对应 8 个 query”。
3. repeat_kv 的张量语义
Section titled “3. repeat_kv 的张量语义”notebook 用小张量复现源码:
bs, seqlen, n_kv_heads, head_dim = 2, 10, 4, 6n_rep = 3x = torch.randn(bs, seqlen, n_kv_heads, head_dim)
x = x[:, :, :, None, :].expand(bs, seqlen, n_kv_heads, n_rep, head_dim)x = x.reshape(bs, seqlen, n_kv_heads * n_rep, head_dim)关键点:expand 先增加共享维度,再 reshape 成与 Q 头数一致的形状。教学代码表现为“重复”,其语义是让多个 Q 头读取同一 K/V;不能把它误解为训练出多份独立的 K/V 参数。
4. 与 KV Cache 的关系
Section titled “4. 与 KV Cache 的关系”每层每个 token 的 K/V 缓存元素量可写为:
若从 MHA 的 n_q 个 KV 头降为 GQA 的 n_kv 个 KV 头,缓存相对比例为:
70B 例子中为 8/64 = 1/8,即 KV 头相关缓存约缩小 8 倍。
公式 / 代码要点
Section titled “公式 / 代码要点”self.n_local_heads = args.n_heads // model_parallel_sizeself.n_local_kv_heads = self.n_kv_heads // model_parallel_sizeself.n_rep = self.n_local_heads // self.n_local_kv_heads
keys = repeat_kv(keys, self.n_rep)values = repeat_kv(values, self.n_rep)
GQA 不改变标准注意力公式;变化发生在 K/V 头的数量和共享方式。
00:05:从上一讲 KV Cache 转入 GQA 源码。00:21:明确 GQA 的目标是提升 attention 计算效率。01:01:给出 70B 的 8 个 KV 头示例。01:36:说明最终计算前需要repeat以匹配 Q 头。02:00:将 GQA 定位为 MHA 与 MQA 的折中。03:13:口述每组 8 个 Q 头共享一组 K/V。06:20:回到源码中的repeat_kv操作。
- notebook 只用随机小张量验证 shape 和
expand/reshape,不是完整 Llama 2 70B 推理。 - 真正源码还叠加模型并行;
n_local_heads与n_local_kv_heads都需按model_parallel_size划分。 - 教学实现物化了展开后的形状;生产内核可直接按组访问,未必真的复制数据。
- 本讲以课程当时的 Llama 2 参数表和 70B 为主,不应据此推断所有模型规模都采用相同 KV 头配置。
与前后讲关系
Section titled “与前后讲关系”- 前接第 20 讲:KV Cache 解释“缓存什么”;本讲解释“如何减少需要缓存的 K/V 头”。
- 后接第 22 讲:把 RMSNorm、RoPE、KV Cache、GQA 等子模块串入完整
generate自回归流程。
- 为什么
n_rep = n_heads // n_kv_heads? - 64 个 Q 头、8 个 KV 头时,每个 KV 头服务多少个 Q 头?
expand + reshape在本例中表达的是参数复制还是共享读取?- GQA 为什么能直接降低 KV Cache,而不改变注意力公式?
- 模型并行后为什么要分别计算
n_local_heads和n_local_kv_heads?
- 视频:llama2 源码分析 GQA:Grouped Query Attention
- 转录:🔒 personal chatgpt 21 视频转录
- 课件 / 代码:🔒 llama2_src_grouped_query_attention.ipynb
- 滴答清单抓取时学习状态:已完成。
- 视频正文采用自动语音识别;关键术语已用标题、notebook 与源码校正,无法确认的口语细节不扩写。
- notebook 未执行的 CUDA、权重下载、训练或外部 API 单元,不表述为已复现实验。