vLLM 中 Attention Kernel 如何并行处理多个请求
在使用 vLLM 推理时,多个请求会被放入同一个 batch 中执行。但这里很容易产生一个误解:
多个请求的 token 被打包在一起后,Attention 是否会把它们当成一条长序列?不同请求之间会不会互相看到?
答案是不会。
vLLM 会把多个请求的 token 放入连续张量,以提高 GPU 计算效率;与此同时,它通过请求边界、序列长度和 KV Block Table 保证每个请求只能访问自己的上下文。
更重要的是,FlashAttention 通常不会真的构造完整的 Attention 相关性矩阵。所谓“多个下三角矩阵”,更多是一个数学上的逻辑视图。GPU kernel 实际上会按 tile 分块计算,并在寄存器中即时生成 causal mask。
一、多个请求的逻辑 Attention 矩阵
假设同时处理两个请求:
请求 A:6 个 token 请求 B:4 个 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 或其他值。
仍然使用:
请求 A:6 tokens 请求 B:4 tokens最大长度是 6,所以 Query 方向需要:
ceil(6 / 4) = 2 个 Query Tile假设只有一个 Attention Head,启动网格为:
grid = (2 requests, 1 head, 2 query tiles)一共启动 4 个 CTA:
| CTA | 负责内容 |
|---|---|
(A, h0, tile0) | A 的 Query 0~3 |
(A, h0, tile1) | A 的 Query 4~5 |
(B, h0, tile0) | B 的 Query 0~3 |
(B, h0, tile1) | 超出 B 的长度,被 mask 掉 |
如果一张 GPU 上有 8 个 local Attention Heads,那么这些 CTA 会分别针对 8 个 head 执行:
A:2 个 Query Tile × 8 Heads = 16 个有效 CTA B:1 个 Query Tile × 8 Heads = 8 个有效 CTA这些 CTA 不需要按照请求顺序执行。GPU 可能这样调度:
SM0:A / head0 / tile0 SM1:B / head5 / tile0 SM2:A / head7 / tile1 SM3:B / 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 Tile
CTA:
(A, head0, query_tile0)负责:
Query positions = [0, 1, 2, 3]它读取第一块 Key:
Key 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 Tile
CTA:
(A, head0, query_tile1)负责:
Query positions = [4, 5]它首先扫描 Key 0~3:
K0 K1 K2 K3 Q4 ✓ ✓ ✓ ✓ Q5 ✓ ✓ ✓ ✓这一块完全位于下三角内部,因此整块有效。
然后扫描 Key 4~5:
K4 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 ]对于请求 A:
sequence_start = 0 sequence_length = 6因此访问范围是:
[0, 6)对于请求 B:
sequence_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,它不会看到 A0~A5。
六、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_i:softmax 指数和。acc:加权 Value 的累计结果。
对每个 K/V Tile:
scores = 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}) \]\[ \alpha=e^{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 threads
- warps
- Tensor 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 可以通过绝对位置生成 mask:
query_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首先计算逻辑 block:
logical_block = 1024 // block_size然后读取:
physical_block = block_table[request_id][logical_block]最后得到物理 KV slot:
physical_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 序列拆成多个 segment:
Segment 0:KV 0~4095 Segment 1:KV 4096~8191 Segment 2:KV 8192~12287 ...启动网格增加一个维度:
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 扫描请求内的全部历史 KV:
Query × 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。