三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

从零开始写Qwen3(六)PagedAttention

从零开始写Qwen3(六)PagedAttention

项目地址:从零开始写Qwen3

概述

在前文中,我们实现了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 = 112KiB2×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.4GiB40K×112KiB=4480MiB4.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,代码如下

fromflash_attnimportflash_attn_varlen_func,flash_attn_with_kvcache...defforward(self,q:torch.Tensor,k:torch.Tensor,v:torch.Tensor):...ifcontext.is_prefill:ifcontext.block_tablesisnotNone:# prefix cachek,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:# decodeo=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_qcu_seqlens_k
    • 累计长度,长度为B+1,第一个值是0,最后一个值是总长度,可以通过两两相减,得到当前请求的长度
  • block_table
    • 大小为B,math.ceil(max_seqlen_k, block_size),表示每个请求KVCache所使用的块序号列表,按照最大长度填充
  • k_cachev_cache
    • 连续KVCache的起始地址,通过块号计算得到对应偏移,从而加载cache

注意到一件事情,解码是不需要累计长度的和最大长度的,因为每个Q长度只有1,通过查询长度就能知道有多少请求,KV也不需要累计长度,因为这里根本没有传入拼接后的KV,而是KVCache的起始地址,直接通过块号查询地址

一个小细节,预填充需要传入累计KV长度,是因为它支持两种模式

  1. block_table,k和v传原始值,而非KVCache,它退化为原始的FlashAttention,因为KV是连续空间,此时需要使用累计长度
  2. block_table,kv传KVCache的起始地址,它通过块号来加载缓存,块号为无效值则停止循环

在没有任何缓存的时候,可以直接使用连续KV,因为分页毕竟还是有些计算开销的

预填充理论上是不使用缓存的,因为预填充是新来的请求,没有任何缓存,但后续产生了前缀缓存和分块预填充这些优化,让预填充也能利用KVCache

  1. 前缀缓存:一个请求如果前面一部分,比如通用提示词,和之前已经算过的请求完全一致,则可以直接把之前算好的KVCache拿过来用,减少重复计算
  2. 分块预填充:单次预填充太长,把预填充分成多次进行,第二次开始就有缓存了

自己来实现一个

整体流程

  1. 进入PagedAttn
  2. 计算QKV投影
  3. 进行ROPE和QKNorm
  4. 把生成的KV写入缓存(不管下一步使用原始KV还是KVCache,总是要写入的)
  5. 执行FlashAttn
  6. 计算O的投影
  7. 返回

基本代码和 FlashAttn一致,只是要多几个地方

  1. 增加分页缓存的读取和写入部分
  2. 增加拆分请求的部分

首先写入缓存非常简单,没有什么计算,单纯的写入,为了简化计算,提前把每个要写入的位置对应的内部索引给算出来,这个对于所有层的所有缓存写入都是一样的,计算一次,给所有层复用

@triton.jitdef_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)ifslot<0:returnoffsets=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.jitdefload_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_SIZEforiintl.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_offsetglobal_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_datareturnresult

这个函数作为 FlashAttn 的内部函数,不单独调用,而是在KV内部循环中用于加载缓存使用

这里做了简化,让KV分块大小正好可以被分页大小整除,方便计算。比如矩阵计算一次加载计算32个词元的长度,而分页大小是16,这就是刚好两个分页,不会出现跨分页的场景

修改了长度解析和加载KV部分,剩下的计算部分完全一致,不用做任何修改

← 返回列表