跳到正文
LLM 推理
S09批处理与调度·244

连续批处理

静态批处理让所有请求都得等最慢的那一个。改成在每次前向之间进出请求,吞吐能提升数倍。

  • 静态 vs 连续
  • 迭代级调度
  • 不规则 batch
  • 消除空泡

为什么重要

问题

S05 留给我们一个令人不适的事实:一次 decode 要读完整个模型才产出一个 token,而 GPU 的算术单元有 99% 以上的时间在闲着。修法几乎不言自明:在同一次前向里处理多个请求,让一次权重读取服务许多 token。batch 32 大约能带来 32 倍吞吐,而每步耗时几乎不变。

然后你动手实现,就撞上了那堵让 LLM 服务不同于其它一切批处理负载的墙:你不知道一个请求会跑多久。

经典批处理组成一批,跑到全部完成,再一起返回结果。放到 LLM 上,这一批就要一直跑到其中最长的那个结束。只需要 5 个 token 的请求会在槽位里干坐 500 步,毫无产出;在这批组成之后一步才到达的请求,只能等整批排空。

核心思路

解法

别再把 batch 当成一个整体。把「准入」和「退出」的决策放到每两次前向之间,而不是每两个 batch 之间。Orca 把这叫做迭代级调度;现在大家都叫它连续批处理。

重构后的循环python
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 是不规则的

batch 里每个请求都处在不同位置,上下文长度也各不相同,根本没有矩形张量可建。kernel 接收的是一个扁平化的 token 缓冲区,外加一个记录累积序列长度的 cu_seqlens 数组,然后据此索引。FlashAttention 的 varlen API 和 PagedAttention 都假定这种形状;连续批处理没法硬塞进朴素实现,原因也正在于此。

工作原理

工作原理

示意图同一份流量下的静态批处理 vs 连续批处理
静态批处理 —— BATCH SIZE 4A (3)B (14)C (4)D (5)E 于 t=2 到达F 于 t=3 到达batch 边界56 个槽位-步里有 30 个是填充。E 干等了 12 步、F 干等了 11 步,而三个槽位什么也没干。连续批处理 —— 同样的槽位,同样的流量槽位 0槽位 1槽位 2槽位 3A 结束 → E 进入C 结束 → F 进入D 结束 → 槽位 3 空闲30 个填充槽位-步缩成 16 个空闲槽位-步:那是可用容量,不是填充,新请求一到就能填上。E 和 F 分别提前 11 步和 10 步进入,六个请求在 t=14 就全部完成,而不是 t=22。循环里到底改了什么1. 退出已完成的请求2. 准入等待中的请求3. 对当前在跑的所有请求做一次前向第 1、2 步从「每 batch 一次」变成了「每 token 一次」。整个想法就这么多。它叫迭代级调度(iteration-level scheduling)。它成立的前提是每请求的 KV 显存能廉价地出现和消失 —— 所以 S06 必须排在前面。代价:batch 现在参差不齐,不同请求处在不同位置,kernel 必须能处理变长。
同样六个请求、同样四个槽位。静态批处理在填充上浪费了 30 个槽位步,还让 E 和 F 分别干等 12 步和 11 步。连续批处理把浪费缩到 16 个空闲槽位步,那是随时可用的容量;E 和 F 分别提前 11 步和 10 步进入,t=14 时全部完成,而不是 t=22。

为什么在真实流量上收益这么大

收益有多大,取决于输出长度的方差。如果每个请求都恰好生成 100 个 token,静态批处理几乎就是最优的,连续批处理也就几乎买不到什么。

真实流量完全不是这样。输出长度是重尾的:大多数回复很短,少数非常长;在 32 个请求的一批里,均值与最大值之比动辄 10 倍以上。按最大值补齐因此浪费掉大半个 batch。实测提升稳定落在 2–4 倍。

典型吞吐提升
2–4×
在对话形态流量上相对静态批处理
而且延迟也变好了
不必再等下一个 batch 边界
不需要额外显存
0
在分页 KV 之上,它纯粹是调度

批处理解决不了什么

批处理提升的是每张 GPU 的吞吐,不是每个请求的延迟。过了某个点,它还会让单请求延迟更差:大 batch 下前向重新变成算力受限,每步更久,batch 里每个请求都感觉得到。

这是服务领域最根本的一笔权衡,S10 正是为它而存在。总得有人根据此刻排队的请求和你承诺的延迟,决定 batch 该多大。

动手观察

动手试试

十二个长度重尾的请求在前十八步内陆续到达。先用连续模式跑一遍,再切到静态模式跑一遍。正面对比面板会保留两次结果,方便你直接比较。

模拟器静态 vs 连续批处理
批处理方式
第 64 / 64 步
已完成
11/12
有效 token
213
GPU 利用率
83%
平均延迟
25.0 步数
执行时间线 · 连续批处理
slot0
slot1
slot2
slot3
  • 正在生成一个 token
  • 填充 —— 已结束,但仍占着槽位
  • 空闲槽位

请求一结束就被退出,槽位在下一步立刻被填上。根本不存在 batch 边界,所谓「batch」就是此刻恰好在跑的那些请求。

同样流量、同样槽位下的正面对比

64 步内完成数

continuous11
static7

GPU 利用率

continuous83%
static47%

平均延迟(步)

continuous25.0
static31.9

连续批处理并不会让任何一次前向变快。它消除的是空隙:填充,以及等待下一个 batch 的时间。在输出长度重尾分布的真实流量上,它稳定带来 2–4× 的吞吐提升,同时延迟也大幅改善。两头都赚的情况很少见,值得留意。

试试把槽位设成 1。两种模式变得完全一样:只有一个槽位时,本来就没有填充可消除。连续批处理是一项占用率优化,而不是 kernel 优化。

亲手实现

实现

不规则 batch 是值得认真写的那一部分。不要用带填充的 [batch, seq, hidden] 张量,而是把每个请求的 token 拼成一条长序列,另带一个偏移数组。

code/s09_continuous_batching.py(节选)python
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 的

一旦 batch 变得异构,下游的一切也都异构。请求 0 可能要贪心解码,请求 1 要 temperature 1.2 配 top-p 0.9,请求 2 还要一个语法约束。把一套参数套用到整个 batch 的批量采样器,是初次实现时最常见的 bug。在有人抱怨 temperature=0 不确定之前,它一直不会露面。
在本地运行
让同一份到达轨迹上、输出长度重尾的合成负载分别经过静态批处理器和连续批处理器,并报告吞吐、GPU 利用率、填充浪费和每请求延迟分位数。
$ python code/s09_continuous_batching.py
预期输出: 连续批处理排空快约 2.9 倍、填充从约 70% 的槽位时间降到零,以及一张并排的延迟分位表,显示两个指标同时改善。

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

生产实践

在生产环境中

  • Orca(OSDI '22) —— 提出迭代级调度和选择性批处理。此后每个引擎都是它的后代。
  • vLLM —— 调度器每次 step() 跑一遍,每轮迭代产出一组全新的运行中请求。
  • TGI、TensorRT-LLM、SGLang、LMDeploy —— 都支持。TensorRT-LLM 叫它 in-flight batching,这个名字更好。
  • llama.cpp —— 在 server 程序里通过并行序列支持它。它的主战场是单用户本地推理,那里 batch 就是 1,这一切都用不上。

练习

  1. 1
    在固定到达速率下,画出吞吐和 P99 延迟随 batch 大小的变化。找到拐点:仍能满足 50 ms token 间延迟目标的最大 batch。
  2. 2
    把输出长度从重尾改成均匀,再跑一遍对比。确认连续批处理的优势基本消失,并用一句话解释原因。
  3. 3
    给批量采样器加上按请求的采样参数,并写出能抓住「整个 batch 共用一个 temperature」这个 bug 的测试。

继续学习

接下来

现在我们每一步都在准入请求,可依据是什么?准入多少?生成中途显存耗尽怎么办?来了一个巨大的 prompt 又怎么办?S10 会构建调度器。这些决策都住在那里,一个引擎大部分对外可见的行为也在那里被决定。

自测

习题

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

自测 1 题 / 共 6

真实流量的哪个性质,让连续批处理值 2–4 倍而不是几乎为零?

得分 0/6