在处理海量样本时,逐条调用模型通常会遇到一个反直觉现象:模型本身并不算慢,真正拖慢整体速度的是频繁的数据搬运、张量构造、设备同步和接口调用。AI推理批量处理的价值,就是把零散请求聚合成更大的计算单元,让GPU或NPU尽可能在一次内核执行中完成更多计算,从而摊薄固定开销。对于文本分类、图像识别、向量嵌入、日志异常检测等任务,批量推理往往能把吞吐量提升数倍甚至一个数量级。

不过,批量处理并不是简单地把所有数据一次性塞进模型。批大小过大可能导致显存溢出,批大小过小又无法发挥并行能力;在线服务还要兼顾延迟,不能为了凑批让请求无限等待。本文将从批量推理的基本原理出发,讨论静态批处理、动态批处理、批大小选择、数据预处理和工程管线设计,帮助你在大规模数据场景下建立稳定的推理方案。
为什么大规模推理需要批量处理
从硬件角度看,GPU擅长的是大规模并行的矩阵运算,而不是频繁的小任务调度。一次只推理一条数据时,每次前向传播都要经历参数读取、输入拷贝、内核启动、结果回传等步骤。单看每一步可能只有几毫秒,但样本数量达到百万级时,这些固定开销会迅速放大。批量推理把多个样本合并成一个张量,让一次矩阵乘法处理更多数据,GPU的算力利用率会明显提高。
除了硬件调度,批量处理还能减少上层业务逻辑的重复成本。例如文本任务中的分词、图像任务中的解码和缩放、向量任务中的归一化,都可以按批执行。很多推理框架也会对批量输入做图优化、算子融合和内存复用,单条请求很难享受到这些收益。换句话说,批量处理不仅让计算更快,也让整个调用链路更简洁。
下面这段伪代码展示了两种常见写法的差异。逐条推理逻辑直观,但容易造成资源浪费;批量推理通过切片和合并输入,把重复操作集中处理,更适合离线任务和高吞吐场景。
# 单条推理:每次都要构造输入、搬运数据、执行模型
results = []
for text in texts:
features = tokenizer(text, return_tensors='np')
logits = model(features)
results.append(postprocess(logits))
# 批量推理:一次处理多个样本,减少重复调度
results = []
batch_size = 32
for start in range(0, len(texts), batch_size):
batch = texts[start:start + batch_size]
features = tokenizer(batch, padding=True, truncation=True, return_tensors='np')
logits = model(features)
results.extend(postprocess(logits))
需要注意的是,批量推理并不是没有代价。批次越大,单条请求等待同批数据完成的时间可能越长,首字延迟或响应延迟也会增加。因此,真正的工程问题不是要不要批处理,而是在吞吐、延迟、显存和稳定性之间找到合适边界。
静态批处理与动态批处理的差异
离线场景常用静态批处理。数据源通常是文件、数据库或对象存储,任务开始前就能确定总量。此时可以把数据按固定大小切分,例如每批64条、128条或256条,然后顺序执行。静态批处理实现简单,适合批量标注、特征抽取、向量入库、历史数据回算等任务。它的优势是可预测性强,失败后也容易按分片重跑。
在线服务场景则更适合动态批处理。请求到达时间不确定,长度和形状也可能差异很大。如果强制使用固定批次,要么请求等待太久,要么批次经常凑不满。动态批处理会在服务端维护一个短时间窗口内的请求队列,当满足最大批大小或最大等待时间时,就合并成一个批次送入模型。这样既能提升吞吐,又不会无限牺牲延迟。
一个简化的动态批处理调度器会包含三个关键参数:最大批大小、最大等待时间和请求队列。下面的示例展示了基本判断逻辑,实际系统中还会加入超时取消、优先级、背压控制和指标监控。
import time
from collections import deque
class Request:
def __init__(self, payload):
self.payload = payload
self.enqueue_time = time.monotonic()
class BatchScheduler:
def __init__(self, max_batch=16, max_wait_ms=20):
self.queue = deque()
self.max_batch = max_batch
self.max_wait_ms = max_wait_ms
def add(self, payload):
self.queue.append(Request(payload))
def ready(self):
if not self.queue:
return False
if len(self.queue) >= self.max_batch:
return True
wait_ms = (time.monotonic() - self.queue[0].enqueue_time) * 1000
return wait_ms >= self.max_wait_ms
def take_batch(self):
size = min(self.max_batch, len(self.queue))
batch = [self.queue.popleft().payload for _ in range(size)]
return batch
动态批处理的难点不在代码结构,而在参数调优。如果最大等待时间设置太短,批次经常很小,吞吐提升有限;如果设置太长,用户会感觉响应变慢。通常需要先观察请求到达率、P95延迟和GPU利用率,再逐步调整。流量波动明显时,还可以按时间段使用不同策略,例如高峰期偏向吞吐,低峰期偏向延迟。
批大小、显存和序列长度如何平衡
批大小是批量推理中最直观的参数,但它并不是越大越好。批大小增加后,吞吐通常会上升,但显存占用也会同步增长。当显存接近上限时,系统可能出现分配失败、频繁交换或推理进程崩溃。很多看似随机的推理失败,根源其实是批大小与输入长度组合后触发了显存峰值。
在文本生成或对话模型中,显存消耗不仅来自输入张量,还来自注意力计算和KV缓存。批大小相同的情况下,长序列带来的显存压力可能远高于短序列。在图像任务中,分辨率和通道数也会产生类似影响。因此,评估批大小时必须结合输入形状,而不是只看样本数量。
| 影响因素 | 对吞吐的影响 | 对延迟的影响 | 常见风险 |
|---|---|---|---|
| 批大小增大 | 通常提升吞吐 | 单请求等待时间可能增加 | 显存溢出 |
| 序列长度增加 | 单批计算量上升 | 响应时间变长 | 长尾请求拖慢整批 |
| 填充过多 | 降低有效计算比例 | 增加无效耗时 | 浪费GPU资源 |
| 混合精度 | 可能提升速度 | 通常降低显存压力 | 数值异常或算子不支持 |
为了减少无效计算,可以使用分桶策略。把长度、分辨率或形状接近的样本放进同一个批次,可以减少填充带来的浪费。下面这段代码演示了按文本长度排序后分批的思路。实际使用时需要保留样本原始编号,避免排序影响结果回写。
def make_buckets(samples, bucket_size=64):
# 按文本长度排序,让相近长度进入同一个批次
indexed = sorted(enumerate(samples), key=lambda item: len(item[1]))
buckets = []
for start in range(0, len(indexed), bucket_size):
buckets.append(indexed[start:start + bucket_size])
return buckets
# 返回结果时需要按原始索引还原顺序
另一个常用手段是混合精度推理。使用fp16或bf16可以降低显存占用并提升部分硬件上的计算速度,但需要先验证数值稳定性。对于量化模型,还要关注精度损失、校准数据和算子兼容性。批大小优化不是一次性配置,而是需要结合压测、监控和线上反馈持续调整。
构建可落地的批量推理管线
大规模推理很少只是调用一个模型接口这么简单。完整管线通常包括数据读取、清洗、分片、预处理、批组装、模型执行、结果序列化、写入存储和失败重试。任何一个环节过慢,都会让GPU等待数据,造成算力空转。设计管线时,首先要明确瓶颈是在CPU预处理、磁盘IO、网络传输,还是在模型计算本身。
对于离线任务,建议使用流式读取而不是把全部数据一次性载入内存。生成器配合批处理函数,可以在控制内存占用的同时保持较高吞吐。如果单进程预处理不足,可以引入多进程数据加载,或者把预处理和推理解耦成独立队列。下面是一个简单的批量处理骨架。
def read_records(path):
with open(path, 'r', encoding='utf-8') as f:
for line in f:
line = line.strip()
if line:
yield line
def batch_iter(iterable, size):
batch = []
for item in iterable:
batch.append(item)
if len(batch) == size:
yield batch
batch = []
if batch:
yield batch
for batch in batch_iter(read_records('input.txt'), 64):
inputs = preprocess(batch)
outputs = model.predict(inputs)
write_results(outputs)
工程上还要特别重视幂等和断点续跑。大规模任务运行时间长,网络抖动、节点重启和模型服务限流都可能出现。如果每条结果都有唯一键,并把已完成分片记录到状态表或对象存储,就可以避免重复计算。对于重要业务,建议先小批量试跑,再逐步扩大并发;同时记录每批耗时、失败率和输入长度分布,方便定位异常。
常见性能瓶颈与优化策略
第一个常见瓶颈是预处理慢于推理。表现为GPU利用率低,但CPU占用很高。解决思路包括提前缓存预处理结果、使用更快的分词器、减少重复编解码、把图像缩放放到专用库中执行,或者增加预处理工作进程。如果预处理必须在线完成,可以考虑把轻量预处理下沉到客户端或网关层。
第二个瓶颈是批内长尾。一个批次里如果混入特别长的文本或特别大的图像,整批都要等待最长样本完成。除了分桶,还可以设置最大序列长度、动态拆分超长样本,或在服务端实现连续批处理,让已完成请求及时返回,新请求随时插入。对于生成式模型,连续批处理能显著减少空闲等待。
第三个问题是显存不足导致的偶发失败。直接捕获异常并重试并不够,因为同样的批次可能再次触发相同问题。更稳妥的做法是记录触发失败的输入长度和批大小,自动降级批次,或者把异常样本转入单独队列。下面的代码展示了一种显存不足时自动折半重试的思路。
def is_oom(error):
return 'out of memory' in str(error).lower()
def safe_predict(batch):
try:
return model.predict(batch)
except RuntimeError as error:
if is_oom(error) and len(batch) > 1:
# 显存不足时自动折半重试,避免整批失败
mid = len(batch) // 2
left = safe_predict(batch[:mid])
right = safe_predict(batch[mid:])
return left + right
raise
此外,模型编译、算子融合、缓存热点结果、限制最大并发、使用向量数据库批量写入等方法,也能在不同场景下带来收益。批量推理优化的核心不是追求某个极限数字,而是让数据供给、模型计算和结果落盘形成稳定节奏。只要能够持续监控吞吐、延迟、显存和失败率,就能把大规模AI推理任务从偶尔跑通变成稳定可复用的生产流程。