尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
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×2112KiB而它最大支持长度为 40K也就是40 K × 112 K i B 4480 M i B ≈ 4.4 G i B 40K \times 112KiB 4480MiB\approx 4.4GiB40K×112KiB4480MiB≈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}addressblockIdx×blockSizetokenId % 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,vk_cache,v_cache oflash_attn_varlen_func(q,k,v,max_seqlen_qcontext.max_seqlen_q,cu_seqlens_qcontext.cu_seqlens_q,max_seqlen_kcontext.max_seqlen_k,cu_seqlens_kcontext.cu_seqlens_k,softmax_scaleself.scale,causalTrue,block_tablecontext.block_tables)else:# decodeoflash_attn_with_kvcache(q.unsqueeze(1),k_cache,v_cache,cache_seqlenscontext.context_lens,block_tablecontext.block_tables,softmax_scaleself.scale,causalTrue)这里预填充和解码使用了不同的核因为两种运算的性质不同预填充是计算密集型的解码是访存密集型的需要分开进行优化。这里涉及到的一些参数如下max_seqlen_q批次中的最大q长度用于划分线程块每个请求按照最大请求长度来算除以Q的块大小。不满最大长度的会自动跳过计算max_seqlen_k批次中最大的k长度可能是用于内部循环优化的cu_seqlens_qcu_seqlens_k累计长度长度为B1第一个值是0最后一个值是总长度可以通过两两相减得到当前请求的长度block_table大小为B,math.ceil(max_seqlen_k, block_size)表示每个请求KVCache所使用的块序号列表按照最大长度填充k_cachev_cache连续KVCache的起始地址通过块号计算得到对应偏移从而加载cache注意到一件事情解码是不需要累计长度的和最大长度的因为每个Q长度只有1通过查询长度就能知道有多少请求KV也不需要累计长度因为这里根本没有传入拼接后的KV而是KVCache的起始地址直接通过块号查询地址一个小细节预填充需要传入累计KV长度是因为它支持两种模式无block_tablek和v传原始值而非KVCache它退化为原始的FlashAttention因为KV是连续空间此时需要使用累计长度有block_tablekv传KVCache的起始地址它通过块号来加载缓存块号为无效值则停止循环在没有任何缓存的时候可以直接使用连续KV因为分页毕竟还是有些计算开销的预填充理论上是不使用缓存的因为预填充是新来的请求没有任何缓存但后续产生了前缀缓存和分块预填充这些优化让预填充也能利用KVCache前缀缓存一个请求如果前面一部分比如通用提示词和之前已经算过的请求完全一致则可以直接把之前算好的KVCache拿过来用减少重复计算分块预填充单次预填充太长把预填充分成多次进行第二次开始就有缓存了自己来实现一个整体流程进入PagedAttn计算QKV投影进行ROPE和QKNorm把生成的KV写入缓存不管下一步使用原始KV还是KVCache总是要写入的执行FlashAttn计算O的投影返回基本代码和 FlashAttn一致只是要多几个地方增加分页缓存的读取和写入部分增加拆分请求的部分首先写入缓存非常简单没有什么计算单纯的写入为了简化计算提前把每个要写入的位置对应的内部索引给算出来这个对于所有层的所有缓存写入都是一样的计算一次给所有层复用triton.jitdef_update_paged_kv_cache_kernel(k_cache,v_cache,k,v,slot_mapping,HIDDEN_DIM:tl.constexpr):n_idtl.program_id(0)slottl.load(slot_mappingn_id)ifslot0:returnoffsetstl.arange(0,HIDDEN_DIM)k_cache_ptrk_cache(slot*HIDDEN_DIMoffsets)v_cache_ptrv_cache(slot*HIDDEN_DIMoffsets)k_ptrk(n_id*HIDDEN_DIMoffsets)v_ptrv(n_id*HIDDEN_DIMoffsets)item_ktl.load(k_ptr)item_vtl.load(v_ptr)target_dtypek_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_DIMNUM_HEADS*HEAD_DIM resulttl.zeros((BLOCK_SIZE_N,HEAD_DIM),dtypecache.dtype.element_ty)dim_offsetstl.arange(0,HEAD_DIM)row_offsetstl.arange(0,BLOCK_SIZE_N)i_sizei_end-i_start# 页内起始偏移: 只有 i_start 是 PAGE_BLOCK_SIZE 整倍数时才为 0# (decode 时 causal STAGE 2 的 i_start N_KEY - 1, 不是整倍数)block_offseti_start%PAGE_BLOCK_SIZEforiintl.range(i_start,i_end,PAGE_BLOCK_SIZE):block_idxi//PAGE_BLOCK_SIZE block_idtl.load(block_tablesblock_idx)# 通过 i_end 可以保证 block_id 0, 而且 triton 中无法写 break 和 continue 就不写了loaded_rowsi-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_rowsrow_offsets[:,None]block_offset)block_datatl.load(cacheglobal_row_offsets*HIDDEN_DIMdim_offsets[None,:],mask((row_offsets[:,None])tl.minimum(i_size,loaded_rowsPAGE_BLOCK_SIZE-block_offset))(row_offsets[:,None]loaded_rows),other0.0,)resultblock_datareturnresult这个函数作为 FlashAttn 的内部函数不单独调用而是在KV内部循环中用于加载缓存使用这里做了简化让KV分块大小正好可以被分页大小整除方便计算。比如矩阵计算一次加载计算32个词元的长度而分页大小是16这就是刚好两个分页不会出现跨分页的场景修改了长度解析和加载KV部分剩下的计算部分完全一致不用做任何修改
返回列表