vLLM 中 Attention Kernel 如何并行处理多个请求
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_sum
kernel 只保留:
当前 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,000
Attention 形状分别为:
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_size
Attention 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 Keys
Sparse 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
完全无效 Tile
Mask 层
下三角不是预生成的矩阵,而是即时计算:
mask = key_position <= query_position
Softmax 层
使用 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。
更多推荐


所有评论(0)