跳转到内容

输入关键词开始搜索

    资料摘要:personal chatgpt 21 — llama2 源码分析 GQA:Grouped Query Attention

    视频摘要更新 2026-08-02置信度 high待阅读原始来源 ↗#深度#大模型#课程#personal-chatgpt

    本讲由真实视频转录与作者 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 头数匹配后再做注意力。

    • MHA:Q、K、V 都有相同数量的独立头,表达力强但 KV Cache 最大。
    • MQA:所有 Q 头共享唯一一组 K/V,缓存最小,但共享程度最高。
    • GQA:在两者之间折中;Q 头被分组,每组共享一个 KV 头。
    • 本讲强调 GQA 的目标是提升 attention,尤其是自回归推理阶段的效率;它与上一讲 KV Cache 是同一条优化链。

    课程示例:

    • n_heads = 64
    • n_kv_heads = 8
    • n_rep = n_heads // n_kv_heads = 8

    因此第 1–8 个 Q 头共享第 1 个 K/V 头,第 9–16 个 Q 头共享第 2 个 K/V 头,以此类推。转录把它口述为“每一组对应 8 个 query”。

    notebook 用小张量复现源码:

    bs, seqlen, n_kv_heads, head_dim = 2, 10, 4, 6
    n_rep = 3
    x = 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 参数。

    每层每个 token 的 K/V 缓存元素量可写为:

    2nkvdhead2,n_{kv},d_{head}

    若从 MHA 的 n_q 个 KV 头降为 GQA 的 n_kv 个 KV 头,缓存相对比例为:

    nkvnq,压缩倍数=nqnkv\frac{n_{kv}}{n_q},\qquad \text{压缩倍数}=\frac{n_q}{n_{kv}}

    70B 例子中为 8/64 = 1/8,即 KV 头相关缓存约缩小 8 倍。

    self.n_local_heads = args.n_heads // model_parallel_size
    self.n_local_kv_heads = self.n_kv_heads // model_parallel_size
    self.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)

    Attention(Q,K,V)=softmax(QKdhead)V\operatorname{Attention}(Q,K,V)=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d_{head}}}\right)V

    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_headsn_local_kv_heads 都需按 model_parallel_size 划分。
    • 教学实现物化了展开后的形状;生产内核可直接按组访问,未必真的复制数据。
    • 本讲以课程当时的 Llama 2 参数表和 70B 为主,不应据此推断所有模型规模都采用相同 KV 头配置。
    • 前接第 20 讲:KV Cache 解释“缓存什么”;本讲解释“如何减少需要缓存的 K/V 头”。
    • 后接第 22 讲:把 RMSNorm、RoPE、KV Cache、GQA 等子模块串入完整 generate 自回归流程。
    1. 为什么 n_rep = n_heads // n_kv_heads
    2. 64 个 Q 头、8 个 KV 头时,每个 KV 头服务多少个 Q 头?
    3. expand + reshape 在本例中表达的是参数复制还是共享读取?
    4. GQA 为什么能直接降低 KV Cache,而不改变注意力公式?
    5. 模型并行后为什么要分别计算 n_local_headsn_local_kv_heads
    • 滴答清单抓取时学习状态:已完成
    • 视频正文采用自动语音识别;关键术语已用标题、notebook 与源码校正,无法确认的口语细节不扩写。
    • notebook 未执行的 CUDA、权重下载、训练或外部 API 单元,不表述为已复现实验。