vLLM 中 Attention Kernel 如何并行处理多个请求

发布时间:2026/7/31 15:23:25
vLLM 中 Attention Kernel 如何并行处理多个请求 vLLM 中 Attention Kernel 如何并行处理多个请求在使用 vLLM 推理时多个请求会被放入同一个 batch 中执行。但这里很容易产生一个误解多个请求的 token 被打包在一起后Attention 是否会把它们当成一条长序列不同请求之间会不会互相看到答案是不会。vLLM 会把多个请求的 token 放入连续张量以提高 GPU 计算效率与此同时它通过请求边界、序列长度和 KV Block Table 保证每个请求只能访问自己的上下文。更重要的是FlashAttention 通常不会真的构造完整的 Attention 相关性矩阵。所谓“多个下三角矩阵”更多是一个数学上的逻辑视图。GPU kernel 实际上会按 tile 分块计算并在寄存器中即时生成 causal mask。一、多个请求的逻辑 Attention 矩阵假设同时处理两个请求请求 A6 个 token 请求 B4 个 token为了提高 GPU 利用率vLLM 可以把它们打包为packed_tokens [A0, A1, A2, A3, A4, A5, B0, B1, B2, B3]同时记录请求边界start_locations [0, 6] sequence_lengths [6, 4]如果把 Attention 分数画成一个全局矩阵逻辑上是A0 A1 A2 A3 A4 A5 | B0 B1 B2 B3 ----------------------------------- A0 | ✓ | × × × × A1 | ✓ ✓ | × × × × A2 | ✓ ✓ ✓ | × × × × A3 | ✓ ✓ ✓ ✓ | × × × × A4 | ✓ ✓ ✓ ✓ ✓ | × × × × A5 | ✓ ✓ ✓ ✓ ✓ ✓ | × × × × ----------------------------------- B0 | × × × × × × | ✓ B1 | × × × × × × | ✓ ✓ B2 | × × × × × × | ✓ ✓ ✓ B3 | × × × × × × | ✓ ✓ ✓ ✓数学上可以表示为一个分块对角结构\[ S \begin{bmatrix} S_A -\infty\\ -\infty S_B \end{bmatrix} \]其中\(S_A\) 是请求 A 自己的下三角 Attention。\(S_B\) 是请求 B 自己的下三角 Attention。两个请求之间的区域全部被屏蔽。不过GPU 上通常不会真的分配这个全局矩阵。二、Attention Kernel 的启动网格vLLM 中一个比较容易理解的 Triton prefill Attention 实现在vllm/v1/attention/ops/triton_prefill_attention.py其启动网格为grid ( batch, num_heads, triton.cdiv(max_input_len, BLOCK_M), )kernel 内部读取cur_batch tl.program_id(0) cur_head tl.program_id(1) start_m tl.program_id(2)可以近似理解为一个 Triton program 对应一个 CUDA thread block也就是一个 CTA。每个 CTA 负责一个请求 × 一个 Attention Head × 一块 Query 行例如program_id (1, 3, 2)表示这个 CTA 负责第 1 个请求 第 3 个 Attention Head 第 2 个 Query Tile因此不同请求之间不是在一个 CTA 里通过复杂 mask 强行分离而是通常从 CTA 分工开始就已经区分开了。不同请求的 CTA 可以同时被调度到不同 SM 上运行。三、用一个小例子说明 CTA 如何分工为了方便展示假设BLOCK_M 4 BLOCK_N 4真实 kernel 中 tile 大小可能是 64、128 或其他值。仍然使用请求 A6 tokens 请求 B4 tokens最大长度是 6所以 Query 方向需要ceil(6 / 4) 2 个 Query Tile假设只有一个 Attention Head启动网格为grid (2 requests, 1 head, 2 query tiles)一共启动 4 个 CTACTA负责内容(A, h0, tile0)A 的 Query 03(A, h0, tile1)A 的 Query 45(B, h0, tile0)B 的 Query 03(B, h0, tile1)超出 B 的长度被 mask 掉如果一张 GPU 上有 8 个 local Attention Heads那么这些 CTA 会分别针对 8 个 head 执行A2 个 Query Tile × 8 Heads 16 个有效 CTA B1 个 Query Tile × 8 Heads 8 个有效 CTA这些 CTA 不需要按照请求顺序执行。GPU 可能这样调度SM0A / head0 / tile0 SM1B / head5 / tile0 SM2A / head7 / tile1 SM3B / head1 / tile0 ...四、一个 CTA 具体计算哪块矩阵当前 CTA 的 Query 行由下面的代码生成offs_m ( start_m * BLOCK_M tl.arange(0, BLOCK_M) )Key 列位置由下面的代码生成offs_n tl.arange(0, BLOCK_N)1. 第一个 Query TileCTA(A, head0, query_tile0)负责Query positions [0, 1, 2, 3]它读取第一块 KeyKey positions [0, 1, 2, 3]然后计算\[ S_{tile}Q_{0:4}K_{0:4}^{T} \]形状为[4, head_dim] × [head_dim, 4] ↓ [4, 4]causal mask 通过局部位置比较产生pos_q offs_m[:, None] pos_k start_n offs_n[None, :] mask pos_q pos_k得到K0 K1 K2 K3 Q0 ✓ × × × Q1 ✓ ✓ × × Q2 ✓ ✓ ✓ × Q3 ✓ ✓ ✓ ✓然后qk tl.dot(q, k) qk tl.where( mask, qk * softmax_scale, -1.0e8, )因此下三角 mask 并不是提前保存在显存中的矩阵而是在 CTA 内通过query_position key_position即时产生。2. 第二个 Query TileCTA(A, head0, query_tile1)负责Query positions [4, 5]它首先扫描 Key 03K0 K1 K2 K3 Q4 ✓ ✓ ✓ ✓ Q5 ✓ ✓ ✓ ✓这一块完全位于下三角内部因此整块有效。然后扫描 Key 45K4 K5 Q4 ✓ × Q5 ✓ ✓这一块位于对角线上需要逐元素 causal mask。所以一个大下三角矩阵在 tile 层面可以表示为△ · · · ■ △ · · ■ ■ △ · ■ ■ ■ △其中■整块位于下三角区域全部有效。△对角 tile需要逐元素 causal mask。·位于未来区域整个 tile 可以跳过。这也是 FlashAttention 能够减少无效计算的重要原因之一。五、不同请求为什么不会互相访问每个 CTA 首先读取当前请求的信息sequence_length tl.load( sequence_lengths cur_batch ) sequence_start tl.load( start_locations cur_batch )访问 Query 时Q[ sequence_start local_query_position ]访问 Key 和 Value 时K[ sequence_start local_key_position ] V[ sequence_start local_key_position ]对于请求 Asequence_start 0 sequence_length 6因此访问范围是[0, 6)对于请求 Bsequence_start 6 sequence_length 4因此访问范围是[6, 10)更重要的是causal mask 使用的是请求内部的局部位置A 的位置0,1,2,3,4,5 B 的位置0,1,2,3而不是 packed tensor 中的全局位置。所以虽然 B0 在 packed tensor 中位于索引 6但它的局部位置仍然是 0它不会看到 A0A5。六、FlashAttention 不保存完整相关性矩阵朴素 Attention 可以写成scores Q K.T scores causal_mask(scores) probs softmax(scores) output probs V这种实现需要把完整的scores: [sequence_length, sequence_length]写入显存。当序列长度为 50,000 时仅一个 head 的相关性矩阵就包含50000 × 50000 25 亿个元素这显然非常昂贵。FlashAttention 的做法是for each K/V tile: score_tile Q_tile K_tile.T score_tile causal_mask(score_tile) 更新 online softmax 更新 output accumulator output_tile accumulator / softmax_sumkernel 只保留当前 Q Tile 当前 K Tile 当前 V Tile 当前 Score Tile 每行 Running Max 每行 Running Sum 每行 Output Accumulator处理完一个 K/V tile 后当前 score tile 就可以丢弃。七、Online Softmax 如何工作一个 Query Tile 会依次扫描多个 K/V Tile。初始化m_i -inf l_i 0 acc 0其中m_i每个 Query 行目前见过的最大分数。l_isoftmax 指数和。acc加权 Value 的累计结果。对每个 K/V Tilescores Q_tile K_tile.T scores causal_mask(scores)更新最大值m_new max( m_old, rowmax(scores), )计算当前 tile 的指数p exp(scores - m_new)因为最大值可能变化之前的累计结果需要重新缩放alpha exp(m_old - m_new) l_new l_old * alpha sum(p) acc_new ( acc_old * alpha p V_tile )全部 K/V Tile 扫描完成后output acc / l完整公式是\[ m_{new}\max(m_{old},\max S_{tile}) \]\[ \alphae^{m_{old}-m_{new}} \]\[ l_{new} \alpha l_{old} \sum e^{S_{tile}-m_{new}} \]\[ O_{new} \alpha O_{old} e^{S_{tile}-m_{new}}V_{tile} \]这种算法与一次性计算完整 softmax 数学等价但不需要保存完整 Attention Matrix。八、一个 CTA 内的线程如何分工在 Triton 源码中矩阵乘通常只写成scores tl.dot(q, k)它并没有明确规定thread 0 计算 score[0,0] thread 1 计算 score[0,1]Triton 编译器会根据BLOCK_M BLOCK_N head_dim 数据类型 num_warps GPU 架构把 tile 映射到CUDA threadswarpsTensor Core MMA 指令寄存器Shared Memory假设num_warps 8那么一个 CTA 通常包含8 warps × 32 threads 256 threads这些线程协作完成加载 Q Tile 加载 K Tile 执行 Q × Kᵀ 计算每行最大值 计算 softmax 指数和 加载 V Tile 执行 P × V 保存输出可以大致理解为多个 Warp 协作加载 Q/K/V ↓ Q/K 被拆成 Tensor Core Fragment ↓ Warp 执行 MMA 指令 ↓ 每个线程持有部分 Score/Accumulator Fragment ↓ Warp 内或 CTA 内归约每行 Max/Sum ↓ 继续处理下一个 K/V Tile因此不是一个线程负责一个 token也不是一个线程负责 Attention Matrix 的一个完整行更准确的描述是一个 CTA 负责一个矩阵 tile每个线程持有这个 tile 中若干不连续的寄存器 fragment多个 warp 通过 Tensor Core 指令协作完成矩阵乘和归约。具体到“thread 37 最终负责哪些矩阵元素”不能仅从 Triton Python 源码确定因为这个映射由 Triton 编译器和目标 GPU 架构决定。要精确到单线程需要查看编译后的 PTX/SASS。九、Decode 阶段为什么看不到大下三角普通自回归 decode 中每个请求本轮通常只有一个 Query。例如请求 A 上下文长度20,000 请求 B 上下文长度35,000Attention 形状分别为A[1, 20001] B[1, 35001]因为当前 Query 位于序列末尾所以所有历史 Key 都满足key_position query_position对应 mask 是A[✓ ✓ ✓ ✓ ... ✓] B[✓ ✓ ✓ ✓ ... ✓]之所以看不到下三角是因为一个完整 causal Attention 下三角矩阵的最后一行本来就是全部有效。Prefill 的特点是Query Length 接近 Context Length所以会看到明显的下三角。普通 Decode 的特点是Query Length 1 Context Length 很大所以 Attention 更像一个长度很大的向量。十、多 Token 验证时的 Attention 结构假设某个请求已经有20,000 个历史 token本轮需要同时验证 6 个新位置query_length 6 kv_length 20,006逻辑 Attention 结构是20,000 历史 token 本轮 6 token --------------------------------------- Query 0 | 全部可见 | ✓ × × × × × Query 1 | 全部可见 | ✓ ✓ × × × × Query 2 | 全部可见 | ✓ ✓ ✓ × × × Query 3 | 全部可见 | ✓ ✓ ✓ ✓ × × Query 4 | 全部可见 | ✓ ✓ ✓ ✓ ✓ × Query 5 | 全部可见 | ✓ ✓ ✓ ✓ ✓ ✓也就是一个 6×20000 的全有效矩形 一个 6×6 的下三角kernel 可以通过绝对位置生成 maskquery_abs_position context_length query_local_position mask ( key_position query_abs_position )如果 batch 中有多个请求每个请求都有自己的context_length query_start_location sequence_length block_table所以多个这样的 Attention 结构依然互相独立。十一、Paged KV Cache 如何参与计算vLLM 的历史 KV 通常不是按请求连续存放而是分页存放。一个请求内部的逻辑 token 位置logical_position 1024首先计算逻辑 blocklogical_block 1024 // block_size然后读取physical_block block_table[request_id][logical_block]最后得到物理 KV slotphysical_slot physical_block * block_size 1024 % block_sizeAttention CTA 每次加载 K/V Tile 时都通过当前请求的block_table找到对应物理块。因此相同的逻辑位置 1024对于两个请求可能映射到完全不同的物理显存地址request A position 1024 - physical block 37 request B position 1024 - physical block 912这也是多个请求共用一个 KV Cache 内存池却不会混淆的原因。十二、长上下文 Decode 如何增加并行度普通 decode 每个请求只有一个 Query。如果只按请求 × Attention Head启动 CTA那么并发请求少、local head 数少时CTA 数量可能不足。同时一个 CTA 还需要串行扫描数万 token 的 KV Cache。一种优化是把一条长 KV 序列拆成多个 segmentSegment 0KV 04095 Segment 1KV 40968191 Segment 2KV 819212287 ...启动网格增加一个维度grid ( query_blocks, kv_heads, parallel_softmax_segments, )多个 CTA 并行扫描不同 KV Segment。每个 Segment 输出局部最大值 m_s 局部指数和 l_s 局部加权输出 O_s之后第二个 reduction kernel 合并\[ m\max_s m_s \]\[ l\sum_s e^{m_s-m}l_s \]\[ O \frac{ \sum_s e^{m_s-m}O_s }{ l } \]这种方式能把一个很长的 Attention 行拆给多个 CTA提高 SM 并行度。代价是需要额外的中间结果。需要第二个 reduction kernel。多一次全局内存读写和同步。所以只有长上下文、并行度不足时才值得这样做。十三、Dense Attention 与 Sparse Attention 的差异Dense Attention 会让当前 Query 扫描请求内的全部历史 KVQuery × 20,000 Keys Query × 50,000 KeysSparse Attention 会先为每个 Query 选择部分相关位置例如top-k 2048于是 Attention 变成每个 Query 只与选中的 2048 个 Key 计算相关性如果 Query 向量维度是 576则单个 Query/Head 的主要相关性计算近似为[1, 576] × [576, 2048] ↓ [1, 2048]这时逻辑上不再是完整的下三角矩阵而是一个经过索引选择后的稀疏相关性向量。但请求隔离机制仍然一样当前 Query 属于哪个请求 ↓ 查询该请求的 Block Table ↓ 把请求内 top-k 逻辑位置转换为物理 KV Slot ↓ 只访问该请求的 KV Cache十四、总结可以把 vLLM 的多请求 Attention 归纳为以下几层。请求层每个请求拥有自己的request_id sequence_length query_start_location block_table张量层多个请求的 token 沿 token 维打包Q: [total_query_tokens, heads, head_dim]但请求边界仍然保留。CTA 层一个 CTA 通常负责一个请求 × 一个 Head 或 KV Head × 一个 Query Tile × 一段 KV Tile不同请求的 CTA 可以并行运行在不同 SM 上。Tile 层大的 causal 下三角被拆成完整有效 Tile 对角三角 Tile 完全无效 TileMask 层下三角不是预生成的矩阵而是即时计算mask key_position query_positionSoftmax 层使用 online softmax逐个处理 K/V Tile不保存完整相关性矩阵。线程层一个线程不负责一个完整 token也不固定负责一个矩阵元素。一个 CTA 内的多个 warp 通过 Tensor Core MMA 协作计算矩阵 tile每个线程持有一部分寄存器 fragment。最终“多个请求合并计算”的准确含义是多个请求共享一次 kernel launch 和 GPU 调度网格但每个 CTA 根据请求边界和 Block Table 访问独立的 Q/K/V 范围Attention 分数按 tile 计算causal mask 在寄存器中即时产生不同请求之间从始至终不会发生语义上的 Attention。