vLLM 中 Attention Kernel 如何并行处理多个请求
2026/7/31 15:22:48 网站建设 项目流程

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。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询