作者: 引线小白-本文永久链接:https://www.limoncc.com/post/64e24a5816e7035f/
知识共享许可协议: 本博客采用署名-非商业-禁止演绎4.0国际许可证
一、基础概念
什么是 BMM (Batch Matrix Multiplication)。BMM 是批矩阵乘法,指的是对同一批次中的多个矩阵同时执行矩阵乘法。在 PyTorch 中,torch.matmul 或 @ 运算符在处理三维及以上张量时,遵循以下规则:把前 N−2 个维度当作“批次维度”,把最后 2 个维度当作“矩阵维度”进行乘法。它并不是一个独有的数学定义,而是深度学习框架为了高效处理“多组独立矩阵乘法”而提供的一种运算,通常对应 API 名称 bmm 或 matmul 的批处理能力。它并不是一个独有的数学定义,而是深度学习框架为了高效处理“多组独立矩阵乘法”而提供的一种运算,通常对应 API 名称 bmm 或 matmul 的批处理能力。下面重点讲它的形状约定和在不同场景下的行为。
1.1、严格 BMM 约定
PyTorch 中的 torch.bmm(input, mat2) 是最经典的 BMM 约定,规则非常严格:
- 输入必须是两个 3 维张量
- 两个张量的第一维(batch 维)必须相等。
- 不支持广播,如果 batch 大小不同会直接报错。
1 | # PyTorch 示例 |
1.2、广义批量矩阵乘法(如 torch.matmul / tf.matmul / numpy.matmul)
现代框架中的 matmul 或 @ 运算符支持任意多维张量,把最后两个维度作为矩阵,前面的所有维度都视为批维度,并遵循广播(broadcasting)规则。约定:
- 如果张量维度 >2,则前导维度(除去最后两维以外的部分)会进行广播。
- 矩阵维度遵循:(…, M, K) × (…, K, N) → (…, M, N)
1 | # 示例:PyTorch 中的 matmul |
这里 batch 维不必相等,只要可广播即可。这正是 Transformer 多头注意力中同时处理多个 head 和 batch 的基础。
1.3、应用场景
全连接层的批量计算
输入 (B, features) 可以看作 (B, 1, K),权重 (K, N) 拓展为 (1, K, N),用 BMM 得到 (B, 1, N),等价于 linear 的批处理。注意力机制中的批量点积
Query: (B, num_heads, seq_len, d_k)
Key 转置: (B, num_heads, d_k, seq_len)
两者使用 matmul 直接得到注意力分数 (B, num_heads, seq_len, seq_len),同时利用了 head 和 batch 的双重批量。图神经网络中的边特征聚合
邻接矩阵的批次运算,多个图样本同时进行消息传递。
1.4、与爱因斯坦求和(einsum)的关系
BMM 用 einsum 可以表示为:
- torch.bmm(a, b) 对应 torch.einsum(‘bmk,bkn->bmn’, a, b)
- 广义批量乘法则是 ‘…mk,…kn->…mn’,省略号代表广播的前导维度。
理解 einsum 后,BMM 的约定其实就浓缩为:共享 batch 下标,矩阵内积求和。
1.5、总结
BMM 的核心约定就是:把独立的一批矩阵乘法打包,用统一的操作并行计算,要求(或通过广播使)除最后两个矩阵维度以外的所有前导维度对齐,内部仍遵守二维矩阵乘法规则。如果你在代码中遇到 bmm 或 matmul,根据是否需要广播和输入维度数来选择即可。
二、KV Cache 原理
在大型语言模型(LLM)的自回归生成过程中,推理速度面临极大的挑战。核心瓶颈在于:每生成一个新 Token,都需要对所有历史 Token 进行 Attention 计算。如果不加优化,生成长度为 $N$ 的序列,总计算量将是 $O(N^2)$。KV Cache 和 GQA (Grouped-Query Attention)是解决这一瓶颈的两大利器:前者通过空间换时间避免重复计算,后者通过结构优化大幅缩减缓存体积。
根据 Attention 机制,当前 Token 的 Query 会和所有历史 Token 的 Key 计算相似度,再和历史 Token 的 Value 加权求和。
- 无 Cache 的痛点:生成第 $t$ 个 Token 时,我们需要把 $1$ 到 $t-1$ 的 Token 重新送入模型算出 $K_{1..t-1}$ 和 $V_{1..t-1}$。这导致历史 Token 被反复计算了 $N$ 次。
- KV Cache 的核心思想:既然 Attention 计算只需要 $K$ 和 $V$,我们可以在生成第 $t-1$ 个 Token 时,把对应的 $K_{t-1}$ 和 $V_{t-1}$ 缓存下来。生成第 $t$ 个 Token 时,只需计算当前 Token 的 $Q_t, K_t, V_t$,然后将 $K_t, V_t$ 拼接到历史的 Cache 中即可。
推理的两个阶段
- Prefill 阶段(预填充):输入 Prompt,并行计算所有 Token 的 KV 并缓存。此阶段属于计算密集型。
- Decode 阶段(解码):逐个生成 Token,每步读取历史 KV Cache,并追加当前步的 KV。此阶段属于显存访问密集型。
KV Cache 并非没有代价,它消耗巨大的显存。对于模型层数 $L$,隐藏层维度 $d_{model}$,序列长度 $N$,精度为 $b$ 字节(如 FP16 为 2 字节):
$$\begin{align}
\text{KV Cache Size} = 2 \times L \times N \times d_{model} \times b
\end{align}$$
以 LLaMA-2 70B 为例:$L=80, d_{model}=8192, N=4096, b=2$,单条序列的 KV Cache 约需 10GB 显存!为了减少 KV Cache 体积,GQA 应运而生。
三、从 MHA 到 GQA
3.1、标准 Multi-Head Attention (MHA)
在 MHA 中,有 $h$ 个 Query 头,$h$ 个 Key 头,$h$ 个 Value 头。对于输入 $\bm{X} \in \mathbb{R}^{N \times d}$:
$$\begin{align}
\bm{Q} = \bm{X}\bm{W}_Q, \quad \bm{K} = \bm{X}\bm{W}_K, \quad \bm{V} = \bm{X}\bm{W}_V
\end{align}$$
其中 $\bm{W}_Q \in \mathbb{R}^{d \times (h \cdot d_k)}$, $\bm{W}_K \in \mathbb{R}^{d \times (h \cdot d_k)}$, $\bm{W}_V \in \mathbb{R}^{d \times (h \cdot d_v)}$。
Attention 计算公式为:
$$ \text{Attention}(\bm{Q}_i, \bm{K}_i, \bm{V}_i) = \mathrm{softmax}\left(\frac{\bm{Q}_i \bm{K}_i^T}{\sqrt{d_k}}\right) \bm{V}_i $$
KV Cache 体积:与头数 $h$ 成正比。
3.2、Multi-Query Attention (MQA)
MQA 极端地让所有 Query 头共享1个 Key 和 Value 头。
$$\begin{align}
h_K = h_V = 1
\end{align}$$
这极大地节省了 Cache,但由于参数量骤减,可能导致模型质量下降。
3.3、Grouped-Query Attention (GQA) —— 完美的折中
GQA 将 $h$ 个 Query 头分为 $g$ 个组,每个组共享1个 Key 头和1个 Value 头。定义每组包含 $m = h/g$ 个连续的 Query 头。则,对于索引为 $i\in [0, h-1]$ 的 Query 头,它属于第 $ \lfloor i/m \rfloor $ 组,因此使用的 KV 头索引也是这个组号:
$$\begin{align}
j = \left\lfloor \frac{i}{m} \right\rfloor = \left\lfloor \frac{i \cdot g}{h} \right\rfloor
\end{align}$$
即 Query 头 0 到 $m-1$ 共享 KV 头 0,Query 头 $m$ 到 $2m-1$ 共享 KV 头 1,依此类推。
$$\begin{align}
\text{Attention}(\bm{Q}_i, \bm{K}_j, \bm{V}_j) = \mathrm{softmax}\left(\frac{\bm{Q}_i \bm{K}_j^T}{\sqrt{d_k}}\right) \bm{V}_j
\end{align}$$
KV Cache 体积:缩小为 MHA 的 $\frac{g}{h}$ 倍!
例如 LLaMA-2 70B:$h=64, g=8$,KV Cache 体积缩小为原来的 1/8!
四、代码实现:带 KV Cache 的 GQA
4.1、层内处理
下面使用 PyTorch 从零实现一个带有 KV Cache 的 GQA 模块。核心逻辑拆解
- 1、投影输出:计算当前输入的 $Q, K, V$。注意 $K, V$ 的形状是 [batch, seq_len, num_kv_heads, head_dim]。
- 2、Cache 拼接:将算出的 $K, V$ 与历史 Cache 拼接。
- 3、GQA 扩展:将 $K, V$ 的 num_kv_heads 维度扩展为 num_q_heads,以便与 $Q$ 进行标准的 Batch Matrix Multiplication (BMM)。
- 4、计算 Attention:标准的缩放点积注意力。
1 | import torch |
4.2、层外处理
堆叠后就是这样了。
1 | self.layers = nn.ModuleList([LLMBlock(l, config) for l in range(self.num_hidden_layers)]) |
在外层需要处理位置编码,我们使用past_key_values 存储所有层的KV缓存。
1 | # 不存在,就先预留位置 |
然后向每层传递
1 | kv_caches = [] |
当然这里是极简实现,transformers库的实现要复杂的多,考虑到了更多工程问题。但是核心也就这些。在真实的生产环境中,除了 GQA,我们还需要解决 Cache 显存碎片问题。vLLM 提出的 PagedAttention 借鉴了操作系统的虚拟内存分页机制,将不连续的 KV Cache 显存块管理起来,结合 GQA,实现了极高的吞吐量。
五、前缀缓存(Prefix Caching)
KV Cache 如何通过空间换时间避免自回归生成时的重复计算。然而,在实际的大模型应用场景(如 RAG 检索增强生成、多轮对话、Agent 系统提示词)中,我们常常面临一个新的痛点:不同请求之间,往往包含大量完全相同的文本前缀。
如果每个请求都独立计算这部分相同前缀的 KV Cache,不仅浪费海量算力,更极大地增加了系统的首字响应时间(TTFT)。Prefix Caching(前缀缓存)正是为解决这一痛点而生。
5.1、Prefix Caching 核心思想
既然相同输入 Token 必然生成相同的 KV Cache,何不将其缓存起来供所有请求共享?
Prefix Caching 将 KV Cache 的生命周期从“单个请求的生存期”提升到了“全局跨请求的生存期”。当新请求到达时,系统首先检查其前缀是否已经被缓存:
- 命中:直接从显存/内存中加载对应的 KV Cache,只需对新接入的 Token 计算 Prefill。
- 未命中:正常计算,并将计算出的 KV Cache 写入缓存池。
物理实现的关键:PagedAttention(分页注意力)
在传统的连续 KV Cache 存储中,不同请求的 Cache 在显存中是分散且大小不一的,无法直接共享。Prefix Caching 的工业级实现(如 vLLM)必须依赖 PagedAttention:
- 将 KV Cache 切分为固定大小的 *Block(页),类似操作系统的虚拟内存分页。
- 相同前缀的 KV Cache 指向同一组物理 Block。
- 采用 Copy-on-Write(写时复制)机制:当请求 B 在前缀后生成新 Token 时,新 Token 的 KV Cache 会被写入新分配的 Block,而不会覆盖共享的前缀 Block。
代码实现:带 Prefix Caching 的 GQA
为了直观展示原理,以下代码不涉及复杂的底层显存分页管理,而是用 PyTorch 和字典模拟 Prefix Cache 的逻辑匹配与复用过程。
我们复用上一讲的 GQAttentionWithCache,并在此基础上构建一个 PrefixCacheManager。
5.2、PyTorch 模拟代码
1 | import torch |
核心逻辑拆解:
- 哈希匹配:PrefixCacheManager 对 Token ID 序列求 SHA-256 哈希。只要用户的 Token 完全相同,哈希值就一致。
- MISS 分支:请求 A 首次到达,缓存为空。模型被迫对前 10 个 Token 做 Prefill,然后将计算出的 past_kv 存入 Manager。
- HIT 分支:请求 B 到达,发现前缀哈希命中。直接取出 past_kv,将其作为 kv_cache 参数传入 Attention,仅对后面的 3~4 个新 Token 执行 Prefill。
5.3、工程实践:从逻辑到物理的跨越
上述代码演示了逻辑原理,但在真实的生产环境(如高并发服务)中,存在极大的工程挑战:
显存碎片与 PagedAttention
真实场景中,前缀长度千变万化。如果为每个前缀分配连续的 Tensor 显存,显存会迅速被碎片化,导致 OOM。
解法:vLLM 引入了 PagedAttention。将 KV Cache 分割为固定大小的 Block(如 16 个 Token 一个 Block)。前缀缓存以 Block 为单位存储。请求 B 命中前缀时,只需在页表中映射指向这些物理 Block,无需拷贝数据。
写时复制
请求 B 在前缀之后生成了新 Token,其 KV Cache 需要追加。由于前缀 Block 是共享的,绝对不能直接修改。
解法:新增的 KV 写入新分配的 Block 中,逻辑上通过链表/页表将其与共享的前缀 Block 串联起来,形成完整的 KV Cache 逻辑视图。
驱逐策略
显存有限,不可能缓存所有历史前缀。当显存不足时,需要淘汰旧缓存。
解法:类似操作系统的 LRU(最近最少使用)策略。vLLM 的 PrefixCaching 调度器会监控 Block 的引用计数和时间戳,优先淘汰没有请求在使用且最久未访问的前缀 Block。
RadixAttention (SGLang)
相比于单纯的前缀匹配,SGLang 提出了更激进的 RadixAttention(基数树注意力)。它将所有请求的 Token 序列在一棵 Radix Tree 上进行前缀匹配。不仅系统提示词可以共享,多轮对话的历史记录、Agent 中间步骤的公共子序列,都能在树形结构中找到最长公共前缀并复用 KV Cache。
| 版权声明 | ![]() |
| 由引线小白创作并维护的柠檬CC博客采用署名-非商业-禁止演绎4.0国际许可证。 本文首发于柠檬CC [ https://www.limoncc.com ] , 版权所有、侵权必究。 | |
| 本文永久链接 | https://www.limoncc.com/post/64e24a5816e7035f/ |
| 如果您需要引用本文,请参考: |
| 引线小白. (May. 19, 2026). 《大语言模型研究15——注意力机制优化之KV缓存》[Blog post]. Retrieved from https://www.limoncc.com/post/64e24a5816e7035f |
| @online{limoncc-64e24a5816e7035f, title={大语言模型研究15——注意力机制优化之KV缓存}, author={引线小白}, year={2026}, month={May}, date={19}, url={\url{https://www.limoncc.com/post/64e24a5816e7035f}}, } |
