推理服务 · 调度
- 玩具服务器
- 4 个槽位 · 12 个请求
- 步数
- 静态 56 → 连续 40
- 单步时间
- 8.0 ms LLaMA 3 8B · A100
全部步骤
为什么要批处理
一个解码步为了给每条序列产生一个词元,要读模型的每一个权重。无论一步里是一条序列还是几十条共享,读取的时间都一样,所以批处理几乎免费地成倍提高每秒词元数,直到算术追上内存。
单步时间 ≈ max(权重 / 带宽, FLOPs / 峰值)静态批处理
静态批处理时,服务器开始一批后要一直跑到其中最长的请求完成。已完成的请求占着槽位什么也不做(斜线),新请求要等整批结束。
一批持续到其中最长的请求结束每一步之后重新填充
连续(迭代级)批处理出自 Orca 论文,每一步之后重新调度:完成的请求立刻离开,等待的请求在下一步接替它的槽位,其提示词的预填充也并入那一步。
每一步之后填充空闲槽位等待与吞吐量
同样的请求,同样的四个槽位。每一步都重新填充,让槽位一直在做真正的工作,所以请求几乎不用等就能开始,平均完成得早得多,整个队列也用更少的步数完成。
等待 ↓ · 延迟 ↓ · 吞吐量 ↑长提示词分块处理
新请求的预填充可能很长:一步跑完的话,它的 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]