LogoCookLLM文档
LogoCookLLM文档
首页CookLLM

原理精讲

词元化
Tokenization 基础BPE 算法详解GPT 系列 TokenizerBPE 训练工程化
模型架构
Transformer LM
从 token ids 到 logitsEmbedding 与 LM Head
Attention 机制
Self-Attention 到 GQAAttention Sink
位置编码
位置编码基础RoPE 数学推导RoPE 代码实现长度外推
GPU 编程基础
GPU 架构基础张量布局Triton 入门:向量加法
FlashAttention
Flash Attention 原理详解从朴素实现到 Auto-TuningBlock Pointer 与多维支持Causal Masking 优化Grouped Query Attention反向传播实现
分布式训练
数据并行ZeRO 优化器全分片数据并行张量并行流水线并行多维混合并行
推理优化
KV CacheContinuous BatchingPagedAttention

动手训练

概述
预训练
预训练数据Tokenizer 训练模型架构数据流水线训练循环监控与验证
X (Twitter)
系统工程FlashAttention

Grouped Query Attention

会员专享

实现 GQA/MQA 支持,让多个 Query Head 共享 KV,优化 KV Cache 内存占用。

在权益中心获取代码

在前面的章节中,我们实现了完整的 Flash Attention Kernel,支持任意序列长度、多维并行和 Causal Masking。本章将添加最后一个关键特性:Grouped Query Attention (GQA)。

这是 Llama 2/3、Mistral、Falcon 等主流开源模型的标配技术,能够在几乎不损失模型质量的前提下,将 KV Cache 内存占用减少 4-8 倍。

GQA 快速回顾

Grouped Query Attention (GQA) 通过让多个 Query heads 共享同一组 Key/Value heads 来减少 KV Cache 内存:

登录以继续阅读

这是一篇付费内容,请登录您的账户以访问完整内容。

# MHA: 每个 Q head 有独立的 K/V head
Q: (B, 8, N, D)
K: (B, 8, N, D)  # 8 个 KV heads
V: (B, 8, N, D)

# GQA: 多个 Q heads 共享 KV heads (groups=4)
Q: (B, 8, N, D)
K: (B, 2, N, D)  # 只有 2 个 KV heads → 内存减少 4x



Head 索引映射公式:

kv_head_idx(hq)=⌊hqHQ/HKV⌋\text{kv\_head\_idx}(h_q) = \left\lfloor \frac{h_q}{H_Q / H_{KV}} \right\rfloorkv_head_idx(hq​)=⌊HQ​/H

如果您对 GQA/MQA 的概念、演进历史、内存分析和设计权衡还不熟悉,建议先阅读 Attention 机制详解 中的详细介绍。

本章重点关注如何在 Flash Attention 中零拷贝、零额外内存地实现 GQA。

Causal Masking 优化

为自回归模型实现因果注意力机制,通过跳过上三角计算实现 ~2x 加速。

反向传播实现

实现 Flash Attention 的梯度计算,通过 Recomputation 实现内存高效的训练。

目录

GQA 快速回顾
PyTorch 标准实现的问题
Flash Attention 中的高效实现
核心思想:指针索引而非数据复制
具体例子:指针偏移的魔法
统一支持 MHA/GQA/MQA
完整实现
性能验证
数值正确性测试
性能基准测试
Autotune 的最佳配置
设计权衡与选择建议
质量 vs 内存的权衡
实现要点总结
总结
V: (B, 2, N, D)
# Q heads 0-3 → KV head 0
# Q heads 4-7 → KV head 1
KV
​
hq​
​
⌋