连续批处理
静态批处理让所有请求都得等最慢的那一个。改成在每次前向之间进出请求,吞吐能提升数倍。
- 静态 vs 连续
- 迭代级调度
- 不规则 batch
- 消除空泡
为什么重要
问题
S05 留给我们一个令人不适的事实:一次 decode 要读完整个模型才产出一个 token,而 GPU 的算术单元有 99% 以上的时间在闲着。修法几乎不言自明:在同一次前向里处理多个请求,让一次权重读取服务许多 token。batch 32 大约能带来 32 倍吞吐,而每步耗时几乎不变。
然后你动手实现,就撞上了那堵让 LLM 服务不同于其它一切批处理负载的墙:你不知道一个请求会跑多久。
经典批处理组成一批,跑到全部完成,再一起返回结果。放到 LLM 上,这一批就要一直跑到其中最长的那个结束。只需要 5 个 token 的请求会在槽位里干坐 500 步,毫无产出;在这批组成之后一步才到达的请求,只能等整批排空。
核心思路
解法
别再把 batch 当成一个整体。把「准入」和「退出」的决策放到每两次前向之间,而不是每两个 batch 之间。Orca 把这叫做迭代级调度;现在大家都叫它连续批处理。
def run(self):
while self.has_work():
# --- these two steps used to happen once per BATCH ----------
for req in self.running:
if req.finished:
self.retire(req) # free its KV blocks
while self.can_admit():
self.running.append(self.waiting.popleft())
# ------------------------------------------------------------
logits = self.model.forward_batch(self.running) # one pass
for req, row in zip(self.running, logits):
req.append(sample(row, req.params))两行代码挪进了循环里,技术就是这些。这样做的代价之所以付得起,全靠 S06:退出一个请求必须立刻释放它的 KV 显存,准入一个请求必须在不打扰其他人的前提下完成分配。若用连续缓存,中途退出一个请求只会留下一个谁都用不了的洞。
现在的 batch 是不规则的
cu_seqlens 数组,然后据此索引。FlashAttention 的 varlen API 和 PagedAttention 都假定这种形状;连续批处理没法硬塞进朴素实现,原因也正在于此。工作原理
工作原理
为什么在真实流量上收益这么大
收益有多大,取决于输出长度的方差。如果每个请求都恰好生成 100 个 token,静态批处理几乎就是最优的,连续批处理也就几乎买不到什么。
真实流量完全不是这样。输出长度是重尾的:大多数回复很短,少数非常长;在 32 个请求的一批里,均值与最大值之比动辄 10 倍以上。按最大值补齐因此浪费掉大半个 batch。实测提升稳定落在 2–4 倍。
- 典型吞吐提升
- 2–4×
- 在对话形态流量上相对静态批处理
- 而且延迟也变好了
- ↓
- 不必再等下一个 batch 边界
- 不需要额外显存
- 0
- 在分页 KV 之上,它纯粹是调度
批处理解决不了什么
批处理提升的是每张 GPU 的吞吐,不是每个请求的延迟。过了某个点,它还会让单请求延迟更差:大 batch 下前向重新变成算力受限,每步更久,batch 里每个请求都感觉得到。
这是服务领域最根本的一笔权衡,S10 正是为它而存在。总得有人根据此刻排队的请求和你承诺的延迟,决定 batch 该多大。
动手观察
动手试试
十二个长度重尾的请求在前十八步内陆续到达。先用连续模式跑一遍,再切到静态模式跑一遍。正面对比面板会保留两次结果,方便你直接比较。
- 已完成
- 11/12
- 有效 token
- 213
- GPU 利用率
- 83%
- 平均延迟
- 25.0 步数
- 正在生成一个 token
- 填充 —— 已结束,但仍占着槽位
- 空闲槽位
请求一结束就被退出,槽位在下一步立刻被填上。根本不存在 batch 边界,所谓「batch」就是此刻恰好在跑的那些请求。
64 步内完成数
GPU 利用率
平均延迟(步)
连续批处理并不会让任何一次前向变快。它消除的是空隙:填充,以及等待下一个 batch 的时间。在输出长度重尾分布的真实流量上,它稳定带来 2–4× 的吞吐提升,同时延迟也大幅改善。两头都赚的情况很少见,值得留意。
试试把槽位设成 1。两种模式变得完全一样:只有一个槽位时,本来就没有填充可消除。连续批处理是一项占用率优化,而不是 kernel 优化。
亲手实现
实现
不规则 batch 是值得认真写的那一部分。不要用带填充的 [batch, seq, hidden] 张量,而是把每个请求的 token 拼成一条长序列,另带一个偏移数组。
def build_batch(running: list[Request]):
"""Flatten a ragged batch the way real kernels want it."""
token_ids, positions, cu_seqlens = [], [], [0]
for req in running:
new = req.tokens_to_process() # prefill: many; decode: 1
token_ids.extend(new)
positions.extend(range(req.computed, req.computed + len(new)))
cu_seqlens.append(len(token_ids)) # cumulative offsets
return Batch(
token_ids=np.array(token_ids), # [total_tokens] — no padding
positions=np.array(positions),
cu_seqlens=np.array(cu_seqlens), # [num_requests + 1]
block_tables=[r.block_table for r in running],
)采样参数是按请求的,不是按 batch 的
temperature=0 不确定之前,它一直不会露面。$ python code/s09_continuous_batching.py只需要 NumPy — 查看环境准备.
生产实践
在生产环境中
- Orca(OSDI '22) —— 提出迭代级调度和选择性批处理。此后每个引擎都是它的后代。
- vLLM —— 调度器每次
step()跑一遍,每轮迭代产出一组全新的运行中请求。 - TGI、TensorRT-LLM、SGLang、LMDeploy —— 都支持。TensorRT-LLM 叫它 in-flight batching,这个名字更好。
- llama.cpp —— 在 server 程序里通过并行序列支持它。它的主战场是单用户本地推理,那里 batch 就是 1,这一切都用不上。
练习
- 1在固定到达速率下,画出吞吐和 P99 延迟随 batch 大小的变化。找到拐点:仍能满足 50 ms token 间延迟目标的最大 batch。
- 2把输出长度从重尾改成均匀,再跑一遍对比。确认连续批处理的优势基本消失,并用一句话解释原因。
- 3给批量采样器加上按请求的采样参数,并写出能抓住「整个 batch 共用一个 temperature」这个 bug 的测试。
继续学习
接下来
现在我们每一步都在准入请求,可依据是什么?准入多少?生成中途显存耗尽怎么办?来了一个巨大的 prompt 又怎么办?S10 会构建调度器。这些决策都住在那里,一个引擎大部分对外可见的行为也在那里被决定。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
真实流量的哪个性质,让连续批处理值 2–4 倍而不是几乎为零?