跳到正文
LLM 推理
S11批处理与调度·190

分块 Prefill

一次 32k token 的 prefill 会把排在它后面的所有 decode 堵上一秒。把 prefill 切成小块,混进 decode 批次里。

  • token 间延迟毛刺
  • chunk 大小调优
  • 顺带执行
  • TTFT/ITL 权衡

为什么重要

问题

S10 的模拟器把它显示成一根紫色尖峰:某一步的全部 token 预算都给了一个长 prefill,排在它后面的每个 decode 都毫无产出。

用户体验到的是输出冻住了。在 32k token 的 prompt 上,这可能接近一秒,长到有人会刷新页面。而对受影响的用户来说,原因完全看不见:是别人发了一个长 prompt。

重排解决不了它。问题不在于哪个请求跑,而在于 prefill 是不可分割的。8192 个 token 要么一次前向全过去,要么根本过不去。

核心思路

解法

让它可分割。没有任何东西要求 prefill 必须在一次前向里完成;注意力是因果的,所以可以先处理前 1024 个 token、把它们的 KV 写进缓存,之后再在某一步基于第一块留下的缓存处理接下来的 1024 个。

一旦 prefill 可切分,调度器就用全部 decode 加上塞得下的那么多 prefill把每一步填到 token 预算。步长变得几乎恒定,尖峰随之消失。

分块调度python
def schedule(self):
    budget = self.max_num_batched_tokens
    batch = []

    for req in self.running:                     # decodes first, 1 token each
        batch.append(req.decode_slice()); budget -= 1

    for req in self.waiting_or_partially_prefilled:
        take = min(req.prompt_len - req.computed, budget, self.max_chunk)
        if take <= 0:
            break
        batch.append(req.prefill_slice(take))    # a SLICE, not the whole prompt
        req.computed += take
        budget -= take

    return batch

分块同时也消除了饥饿这个 bug

在 S10 中,比 token 预算还长的 prompt 根本永远无法被调度。有了分块,任意长度的 prompt 都可调度,代价只是多花几步。这既是延迟上的改进,也是正确性上的改进。

工作原理

工作原理

示意图同一次长 prefill,两种调度
不分块prefill 8,192 个 token一次前向 · 约 150 msdecode 步恢复用户 1卡住用户 2卡住用户 3卡住每个正在流式接收的用户都会看到 150 ms 的空档,只因为另一个用户发了一个长 prompt。在 32k 上下文下这接近一秒,长到看起来像连接断了。启用分块 PREFILL · CHUNK = 1,024第 0 块第 1 块第 2 块第 3 块第 4 块第 5 块第 6 块第 7 块每一步 = 一块 prefill + 每个在跑 decode 的一个 token。步时间几乎恒定。用户 1用户 2用户 3没有空档。整个长 prefill 期间每个用户都持续收到 token。代价是什么那个长 prompt 的 TTFT 会略微上升:固定开销从 1 次变成 8 次。分块要和 decode 混在同一个参差 batch 里 —— 需要 S09。
总工作量完全相同。分块只是把它重新分配,使任何一步都不会长到让流式用户察觉。

把 prefill 和 decode 混在一起为什么不只是公平,而且高效

这样做还有第二个原因,而且是更有意思的那个。prefill 受算力限制,decode 受显存带宽限制。它们争的是不同的资源。

纯 decode 的一步让算术单元闲着;纯 prefill 的一步把算术单元打满,显存系统却相对空闲。把两者混合,decode 的 token 就搭上了 prefill 读权重的顺风车:在一个 prefill 步里额外加 24 个 decode,边际成本远小于让这些 decode 单独占一步。

所以分块 prefill 通常能在降低延迟方差的同时提高总吞吐。提出分块 prefill 的 SARATHI 论文把这种搭车称为「捎带」(piggybacking);其后续论文 Sarathi-Serve 把它做成了调度器,并把结果命名为「无停顿批处理」。这个名字说的是 decode 永远不会停在 prefill 后面,而不是吞吐上的收益。

如何选择 chunk 大小

chunk 大小和 token 预算其实是同一个旋钮的两面。更小的 chunk 意味着更平滑的 token 间延迟和更多的固定每步开销;更大的 chunk 意味着更好的 prefill 效率,但尖峰也会悄悄回来。

一个合理的起点:选出单步耗时仍能满足 token 间延迟目标的最大 chunk,然后去测。vLLM 默认开启分块 prefill,token 预算默认在几千的量级(V1 在线服务中为 8192)。对大多数模型来说,这是几十毫秒一步。记得核对你的版本,这个默认值改过。

动手观察

动手试试

模拟器分块 prefill 与 token 间延迟
decode 用户的最长卡顿
24 ms
长 prompt 的 TTFT
189 ms
超出 ITL 目标的步数
0
decode 吞吐
1451 tok/s
步时间线 · 柱高 = 该步耗时
ITL 目标 50 ms
  • prefill 计算
  • decode 计算

每一步都携带一小片 prompt,同时为 24 个流式用户各产出一个 token。各步耗时相差无几,因此没人会感到卡顿。

你实际在做的权衡

更小的 chunk

  • + token 间延迟更平滑
  • + 新请求能更早开始
  • − 步数更多,固定开销更大
  • − 长 prompt 的 TTFT 更差

更大的 chunk

  • + prefill 效率更高
  • + 长 prompt 的 TTFT 更低
  • − 延迟毛刺回来了
  • − 当 chunk = prompt 时,等于什么都没做

把 chunk 大小设成等于 prompt 长度,分块路径就退化成不分块路径:两个开关产生完全相同的时间线。分块不是白来的,它把一次巨大的延迟毛刺换成了略长的 TTFT。这几乎总是更划算,因为毛刺每个用户都感受得到,而 TTFT 只有一个人承受。

把 prompt 推到 32k 并关闭分块:一根柱子远远高出 ITL 目标,所有 decode 都被冻在它后面。打开分块,同样的工作变成三十二个普通步骤。然后把 chunk 大小设成等于 prompt 长度,看尖峰回来。分块管的只是步的粒度,别无其他。

亲手实现

实现

微妙之处在注意力掩码上。一个 chunk 的 query 要看到全部此前已缓存的 key,再加上它自己 chunk 内部按因果掩码可见的那些 key。

code/s11_chunked_prefill.py(节选)python
def chunk_attention_mask(chunk_len: int, cached_len: int):
    """
    Queries: the chunk_len new tokens.
    Keys:    cached_len already-computed tokens, then the chunk itself.
    """
    total_keys = cached_len + chunk_len
    mask = np.zeros((chunk_len, total_keys), dtype=bool)

    # every new query may attend to ALL cached keys — no causality needed
    mask[:, :cached_len] = True
    # within the chunk, ordinary causal masking applies
    mask[:, cached_len:] = np.tril(np.ones((chunk_len, chunk_len), dtype=bool))

    return mask

位置 id 是绝对的,不是 chunk 相对的

RoPE 用的是 token 在序列中的绝对位置。按 1024 分块时,第 3 块从位置 3072 开始,而不是 0。写错了,第一块之后的每一块都会被当作 prompt 的开头来旋转。模型于是读到一份不断重新开始的文档,自信地输出胡话,而且不会抛出任何错误。
在本地运行
把同一个 prompt 分别一次性 prefill 和分块 prefill,断言得到的 KV 缓存数值完全一致;然后模拟一份混合负载,报告有无分块时的 token 间延迟分布。
$ python code/s11_chunked_prefill.py
预期输出: 在每种 chunk 大小下分块与不分块的 KV 逐位一致的断言、32k 下最坏 token 间延迟约 25× 的改善,以及一次 chunk 大小扫描,展示 TTFT/ITL 的权衡。

只需要 NumPy — 查看环境准备.

生产实践

在生产环境中

SARATHI 是原始论文,提出了分块 prefill 和 decode 捎带(piggybacking)。其后续论文 Sarathi-Serve 加上了无停顿批处理调度器。vLLM 把它暴露为 enable_chunked_prefill,在 v1 中默认开启,prefill 与 decode 共享同一份预算和同一个 batch。TensorRT-LLM 以「chunked context」之名提供同样的想法。另一条路是压根不混合,让 prefill 和 decode 跑在不同机器上,那就是分离部署,也是 S18 的主题。

练习

  1. 1
    把 chunk 大小从 128 扫到 8192,画出 P99 token 间延迟与 prefill 吞吐的关系。找出 50 ms 延迟目标下的拐点。
  2. 2
    把分块与前缀缓存结合:一个 block 已被缓存的 chunk 应当被完全跳过。测量二者的相互作用;它们组合起来比单独任何一个都更好。
  3. 3
    故意把「chunk 相对位置」这个 bug 实现一遍,然后写出能抓住它的测试。把分块后的 KV 缓存与不分块的 prefill 做对比,是能拿到的最便宜的回归测试,而且能抓住一整类错误。

继续学习

接下来

调度已经不错了。注意力 kernel 本身仍然要在显存里构建一个 N×N 的分数矩阵,而长上下文下这个矩阵既是显存瓶颈,也是速度瓶颈。S12 用在线 softmax 把它去掉。

自测

习题

先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。

自测 1 题 / 共 6

除了平滑延迟,分块 prefill 还解决了什么正确性问题?

得分 0/6