分块 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 预算。步长变得几乎恒定,尖峰随之消失。
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
工作原理
工作原理
把 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)。对大多数模型来说,这是几十毫秒一步。记得核对你的版本,这个默认值改过。
动手观察
动手试试
- decode 用户的最长卡顿
- 24 ms
- 长 prompt 的 TTFT
- 189 ms
- 超出 ITL 目标的步数
- 0
- decode 吞吐
- 1451 tok/s
- 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。
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 相对的
$ python code/s11_chunked_prefill.py只需要 NumPy — 查看环境准备.
生产实践
在生产环境中
SARATHI 是原始论文,提出了分块 prefill 和 decode 捎带(piggybacking)。其后续论文 Sarathi-Serve 加上了无停顿批处理调度器。vLLM 把它暴露为 enable_chunked_prefill,在 v1 中默认开启,prefill 与 decode 共享同一份预算和同一个 batch。TensorRT-LLM 以「chunked context」之名提供同样的想法。另一条路是压根不混合,让 prefill 和 decode 跑在不同机器上,那就是分离部署,也是 S18 的主题。
练习
- 1把 chunk 大小从 128 扫到 8192,画出 P99 token 间延迟与 prefill 吞吐的关系。找出 50 ms 延迟目标下的拐点。
- 2把分块与前缀缓存结合:一个 block 已被缓存的 chunk 应当被完全跳过。测量二者的相互作用;它们组合起来比单独任何一个都更好。
- 3故意把「chunk 相对位置」这个 bug 实现一遍,然后写出能抓住它的测试。把分块后的 KV 缓存与不分块的 prefill 做对比,是能拿到的最便宜的回归测试,而且能抓住一整类错误。
继续学习
接下来
调度已经不错了。注意力 kernel 本身仍然要在显存里构建一个 N×N 的分数矩阵,而长上下文下这个矩阵既是显存瓶颈,也是速度瓶颈。S12 用在线 softmax 把它去掉。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
除了平滑延迟,分块 prefill 还解决了什么正确性问题?