从零开始写Qwen3(六)PagedAttention
概述
在前文中,我们实现了FlashAttention,它通过融合自注意力中三个矩阵乘法和一个Softmax的操作,消减了 O ( L 2 ) O(L^2) O(L2) 大小的权重矩阵的显存开销
但除了注意力计算本身,显存开销还有一大重要开销,就是KVCache,对于一些小模型来说,KVCache甚至会比模型本身还大。
本文将从KVCache大小计算开始,介绍传统KVCache的缺陷,详细介绍分页注意力的核心原理和实现逻辑,展示现成的 flash-attn 库的使用方法,并自己用 triton 实现一个
KVCache大小计算
KVCache缓存的就是每个词元每层的K/V的表达,一个词元的表征大小为
L
×
H
×
D
×
单个元素字节数
L\times H\times D\times \text{单个元素字节数}
L×H×D×单个元素字节数
KV就是2倍
对于 Qwen3-0.6B而言,有28层,每层KV头为8,每个头128维,使用bf16的数据,从而一个词元就要
2
×
28
×
1024
×
2
=
112
K
i
B
2\times 28 \times 1024 \times 2 = 112KiB
2×28×1024×2=112KiB
而它最大支持长度为 40K,也就是
40
K
×
112
K
i
B
=
4480
M
i
B
≈
4.4
G
i
B
40K \times 112KiB = 4480MiB\approx 4.4GiB
40K×112KiB=4480MiB≈4.4GiB,而模型本身也只有1点多G
而这仅仅只是 40K 的单个请求的长度,现在的大模型动则200K,甚至1M,请求也远不止一个,在这种情况下显存大小显然成为首要限制,甚至比算力还要重要
传统KVCache管理方式的缺陷
大模型的自回归性质决定KVCache会不断变长,如果简单追加,会出现大量复制,开销很大,于是一种简单的做法就是预分配最大长度
这种模式存在一个巨大的问题,就是显存浪费严重,即使是很短的一个问候都要为它分配最大长度的显存,这些占用的显存无法被其他请求使用。就算是Qwen3-0.6B这种很小的模型,一个请求就要占用4G的显存,即使有1T的显存,也只能提供给128个用户使用
分页自注意力的改进
简单追加空间利用率高,但有大量复制操作,预分配最大长度没有复制,效率高,但空间利用率很低。分页注意力就是将两者结合起来,预分配大量连续空间,但将这个连续空间分割为等长的多个块,每个块可以容纳L个词元的缓存,这样需要的时候只用申请L长度的块就行,空间利用率高了很多
但这里有个重要变化,此时一个请求的KVCache的空间不再连续,它可能变成这样:

为了应对内存的不连续,原先的 FlashAttn 也需要做出相对应的改变:原先是对每个块的Q,按块遍历每个KV,这个步骤假设KV是连续的,所以可以根据词元的序号计算出偏移,但现在序号和地址不再对应,需要变成这样
address
=
blockIdx
×
blockSize
+
tokenId % blockSize
\text{address}=\text{blockIdx} \times \text{blockSize} + \text{tokenId \% blockSize}
address=blockIdx×blockSize+tokenId % blockSize
需要得到每个索引对应的块号
实现
packed 模式
在介绍分块注意力具体实现之前,先介绍一下Packed模式
常规模式下, 如果有多个输入,一般都会把它们打包成一个批次(batch),然后一起计算,这样可以减少GPU启动次数,并且在一些计算中,比如矩阵乘法,还能提高计算密度,提高GPU使用效率
然而对于序列任务,比如自然语言处理,打包成批次会有一个问题,因为每个请求的输入长度是不一样长的,但Batch需要各个长度一致。为了把不同长度的输入打包到一起,通常会使用填充,把每个请求填充到一个批次中的最大长度

对于文本生成这种自回归任务,往往都采用左填充的方式,因为这样预填充完生成的才是紧接着的下一个词元的logits。而训练则会使用左填充,因为训练是一次生成所有的logits,而不是每次产生一个,不需要解码步骤

对于训练而言,填充是能高效利用GPU的好方法,但对于生成,它做了填充,浪费了一些显存空间和计算量
另外一种做法就是不要Batch维度,直接在长度维度把多个请求拼接起来

这种不产生任何填充
对于大多数计算,比如矩阵乘法和元素级运算,它们和长度是没有关系的,不需要做任何改动,因为Batch版本的在进行这些计算的时候通常也都是按长度展平的方式计算的。
唯一的区别在自注意力,它需要把多个请求拆开,分开进行计算
虽然packed模式没有填充浪费,但它毕竟实现复杂,会让本就复杂的FlashAttn的反向传播变得更为复杂,而且训练过程用不到KVCache,所以训练许多情况还是会使用Batch模式,通过一些方式,尽可能让长度相同的匹配到一起,减少浪费
fast-attn 库的使用
分页注意力有现成的库来实现,比如 flash-attn ,flashinfer 等,在 nano-vllm 中直接使用了 flash-attn,代码如下
from flash_attn import flash_attn_varlen_func, flash_attn_with_kvcache
...
def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor):
...
if context.is_prefill:
if context.block_tables is not None: # prefix cache
k, v = k_cache, v_cache
o = flash_attn_varlen_func(q, k, v,
max_seqlen_q=context.max_seqlen_q, cu_seqlens_q=context.cu_seqlens_q,
max_seqlen_k=context.max_seqlen_k, cu_seqlens_k=context.cu_seqlens_k,
softmax_scale=self.scale, causal=True, block_table=context.block_tables)
else: # decode
o = flash_attn_with_kvcache(q.unsqueeze(1), k_cache, v_cache,
cache_seqlens=context.context_lens, block_table=context.block_tables,
softmax_scale=self.scale, causal=True)
这里预填充和解码使用了不同的核,因为两种运算的性质不同,预填充是计算密集型的,解码是访存密集型的,需要分开进行优化。
这里涉及到的一些参数如下
max_seqlen_q- 批次中的最大q长度,用于划分线程块,每个请求按照最大请求长度来算,除以Q的块大小。不满最大长度的会自动跳过计算
max_seqlen_k- 批次中最大的k长度,可能是用于内部循环优化的
cu_seqlens_q,cu_seqlens_k- 累计长度,长度为
B+1,第一个值是0,最后一个值是总长度,可以通过两两相减,得到当前请求的长度
- 累计长度,长度为
block_table- 大小为
B,math.ceil(max_seqlen_k, block_size),表示每个请求KVCache所使用的块序号列表,按照最大长度填充
- 大小为
k_cache,v_cache- 连续KVCache的起始地址,通过块号计算得到对应偏移,从而加载cache
注意到一件事情,解码是不需要累计长度的和最大长度的,因为每个Q长度只有1,通过查询长度就能知道有多少请求,KV也不需要累计长度,因为这里根本没有传入拼接后的KV,而是KVCache的起始地址,直接通过块号查询地址
一个小细节,预填充需要传入累计KV长度,是因为它支持两种模式
- 无
block_table,k和v传原始值,而非KVCache,它退化为原始的FlashAttention,因为KV是连续空间,此时需要使用累计长度 - 有
block_table,kv传KVCache的起始地址,它通过块号来加载缓存,块号为无效值则停止循环
在没有任何缓存的时候,可以直接使用连续KV,因为分页毕竟还是有些计算开销的
预填充理论上是不使用缓存的,因为预填充是新来的请求,没有任何缓存,但后续产生了前缀缓存和分块预填充这些优化,让预填充也能利用KVCache
- 前缀缓存:一个请求如果前面一部分,比如通用提示词,和之前已经算过的请求完全一致,则可以直接把之前算好的KVCache拿过来用,减少重复计算
- 分块预填充:单次预填充太长,把预填充分成多次进行,第二次开始就有缓存了
自己来实现一个
整体流程
- 进入PagedAttn
- 计算QKV投影
- 进行ROPE和QKNorm
- 把生成的KV写入缓存(不管下一步使用原始KV还是KVCache,总是要写入的)
- 执行FlashAttn
- 计算O的投影
- 返回
基本代码和 FlashAttn一致,只是要多几个地方
- 增加分页缓存的读取和写入部分
- 增加拆分请求的部分
首先写入缓存非常简单,没有什么计算,单纯的写入,为了简化计算,提前把每个要写入的位置对应的内部索引给算出来,这个对于所有层的所有缓存写入都是一样的,计算一次,给所有层复用
@triton.jit
def _update_paged_kv_cache_kernel(
k_cache, v_cache, k, v, slot_mapping, HIDDEN_DIM: tl.constexpr
):
n_id = tl.program_id(0)
slot = tl.load(slot_mapping + n_id)
if slot < 0:
return
offsets = tl.arange(0, HIDDEN_DIM)
k_cache_ptr = k_cache + (slot * HIDDEN_DIM + offsets)
v_cache_ptr = v_cache + (slot * HIDDEN_DIM + offsets)
k_ptr = k + (n_id * HIDDEN_DIM + offsets)
v_ptr = v + (n_id * HIDDEN_DIM + offsets)
item_k = tl.load(k_ptr)
item_v = tl.load(v_ptr)
target_dtype = k_cache.dtype.element_ty
tl.store(k_cache_ptr, item_k.to(target_dtype))
tl.store(v_cache_ptr, item_v.to(target_dtype))
这里的 slot_mapping 就是提前算好的每个词元位置对应的内部索引,大小为 B, max_seqlens_k
读取缓存则需要和FlashAttn写在一起
@triton.jit
def load_paged_memory(
cache,
block_tables,
i_start,
i_end,
NUM_HEADS: tl.constexpr,
PAGE_BLOCK_SIZE: tl.constexpr,
HEAD_DIM: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
):
"""
cache: 是PagedAttention的K/V缓存, 总体形状为 (NUM_BLOCKS, NUM_HEADS, BLOCK_SIZE_N),
这里传入的时候已经加上了 head 的偏差
block_tables: 是每个block的id, 总体形状为 (BATCH, cdiv(max_seq_len, PAGE_BLOCK_SIZE))
这里传入的时候已经加上了 batch 的偏差
长度不及 max 的会在最后填充 -1 , 但 i_end 不会加载到那里
"""
HIDDEN_DIM = NUM_HEADS * HEAD_DIM
result = tl.zeros((BLOCK_SIZE_N, HEAD_DIM), dtype=cache.dtype.element_ty)
dim_offsets = tl.arange(0, HEAD_DIM)
row_offsets = tl.arange(0, BLOCK_SIZE_N)
i_size = i_end - i_start
# 页内起始偏移: 只有 i_start 是 PAGE_BLOCK_SIZE 整倍数时才为 0
# (decode 时 causal STAGE 2 的 i_start = N_KEY - 1, 不是整倍数)
block_offset = i_start % PAGE_BLOCK_SIZE
for i in tl.range(i_start, i_end, PAGE_BLOCK_SIZE):
block_idx = i // PAGE_BLOCK_SIZE
block_id = tl.load(block_tables + block_idx)
# 通过 i_end 可以保证 block_id > 0, 而且 triton 中无法写 break 和 continue 就不写了
loaded_rows = i - i_start
# 从 -loaded_heads 开始加载 PAGE_BLOCK_SIZE'
# BLOCK_SIZE_N 中加载 [loaded_heads, loaded_heads + PAGE_BLOCK_SIZE)
# 所以 mask 要把前后的给遮掉
# 页内偏移 = block_offset, 页内可加载量 = PAGE_BLOCK_SIZE - block_offset
global_row_offsets = (
block_id * PAGE_BLOCK_SIZE
- loaded_rows
+ row_offsets[:, None]
+ block_offset
)
block_data = tl.load(
cache + global_row_offsets * HIDDEN_DIM + dim_offsets[None, :],
mask=(
(row_offsets[:, None])
< tl.minimum(
i_size, loaded_rows + PAGE_BLOCK_SIZE - block_offset
)
)
& (row_offsets[:, None] >= loaded_rows),
other=0.0,
)
result += block_data
return result
这个函数作为 FlashAttn 的内部函数,不单独调用,而是在KV内部循环中用于加载缓存使用
这里做了简化,让KV分块大小正好可以被分页大小整除,方便计算。比如矩阵计算一次加载计算32个词元的长度,而分页大小是16,这就是刚好两个分页,不会出现跨分页的场景
修改了长度解析和加载KV部分,剩下的计算部分完全一致,不用做任何修改
更多推荐



所有评论(0)