推理服务 · 调度

连续批处理

每个解码步之后重新填充批次

玩具服务器
4 个槽位 · 12 个请求
步数
静态 56 → 连续 40
单步时间
8.0 ms LLaMA 3 8B · A100

全部步骤

  1. 为什么要批处理

    一个解码步为了给每条序列产生一个词元,要读模型的每一个权重。无论一步里是一条序列还是几十条共享,读取的时间都一样,所以批处理几乎免费地成倍提高每秒词元数,直到算术追上内存。

    单步时间 ≈ max(权重 / 带宽, FLOPs / 峰值)

  2. 静态批处理

    静态批处理时,服务器开始一批后要一直跑到其中最长的请求完成。已完成的请求占着槽位什么也不做(斜线),新请求要等整批结束。

    一批持续到其中最长的请求结束

  3. 每一步之后重新填充

    连续(迭代级)批处理出自 Orca 论文,每一步之后重新调度:完成的请求立刻离开,等待的请求在下一步接替它的槽位,其提示词的预填充也并入那一步。

    每一步之后填充空闲槽位

  4. 等待与吞吐量

    同样的请求,同样的四个槽位。每一步都重新填充,让槽位一直在做真正的工作,所以请求几乎不用等就能开始,平均完成得早得多,整个队列也用更少的步数完成。

    等待 ↓ · 延迟 ↓ · 吞吐量 ↑

  5. 长提示词分块处理

    新请求的预填充可能很长:一步跑完的话,它的 512 个提示词词元会让那一步受算力限制、慢三倍,卡住所有其他请求的下一个词元。分块预填充把它切成足够小的块,藏在受内存限制的单步时间里。

    以 128 为块预填充,与解码混合

代码

# static batching: run a batch to completion
while batch: step(batch); batch = [r for r in batch if not r.done] or next_batch()
# continuous batching: reschedule every iteration
while True:
    running = [r for r in running if not r.done]           # finished ones leave now
    while waiting and len(running) < max_batch and fits(waiting[0]):
        running.append(waiting.pop(0))                     # its prefill joins this step
    step(running)                                          # one token for each
# chunked prefill: cap the prompt tokens per step
budget = 128; chunk = prompt[done:done + budget]

延伸阅读