前言
上线三个月的 RAG 服务,Embedding 阶段每天凌晨批量处理 10 万条文档,跑了 6 小时,API 账单快到上限。
排查之后发现:代码里每条文本独立发一次 /embeddings 请求,10 万条 = 10 万次 HTTP roundtrip。把请求合批之后,同样的数据量跑了 22 分钟,API 调用次数从 10 万次降到 800 次。
但问题来了——合批之后在线服务的 P99 延迟从 80ms 涨到 340ms,用户体验直接崩了。
这就是 Adaptive Batching 要解决的核心矛盾:吞吐和延迟不是一个旋钮可以同时拧好的。这篇文章讲的是应用层如何做动态合批,让两个目标在可接受范围内共存。
一、先把概念理清楚
"Batching"在 LLM 语境里至少有三层含义,混淆了会走很多弯路。
1.1 服务端 Continuous Batching(不是本文主题)
vLLM、TGI 等推理框架做的事情:在 token 生成的 iteration 级别动态调度多个请求,共享 KV cache,GPU 利用率接近 100%。这是推理框架内部的优化,应用层感知不到,也控制不了。
1.2 Batch API(离线批处理)
主流大模型 Batch API(各家均有提供):提交一个 JSONL 文件,异步返回,价格打 5 折。适合不需要实时响应的大批量任务:数据标注、离线评估、报告生成。
这个用法很多人知道,本文不重点讲。
1.3 应用层 Adaptive Batching(本文主题)
在你自己的应用代码里,把短时间内积累的多个请求动态合并成一个 API 调用发出去。适合:
- 实时 Embedding 服务(RAG 检索前处理)
- 多用户并发请求同一个分类/标注任务
- 需要实时返回但又想控制 API 调用频率的场景
关键词是**"动态"**——不是固定等 100 个凑满再发,而是根据当前负载实时决定等多久、批多大。
二、为什么 Static Batching 不够用
最朴素的批处理方案是这样的:
# 攒够 N 个请求再发
batch = []
for text in texts:
batch.append(text)
if len(batch) >= 100:
results = embed(batch)
batch = []
if batch:
results = embed(batch) # 处理剩余
这在离线批处理里完全够用,但在在线服务里有几个致命问题:
问题一:低峰期永远在等
流量低的时候,batch 凑不满 100 个,要么一直等(延迟无限高),要么设个超时 T——那每个请求的延迟至少是 T。
问题二:高峰期 batch 堆积
流量突然来了 500 个请求,按顺序每批 100 个发出去,第 5 批请求等待时间 = 前 4 批的处理时间总和。
问题三:一刀切的 batch size 不合适
Embedding 模型对 batch 大小有不同的最优点,单批 token 上限取决于模型配置,但实际文本长短不一,固定 100 条可能超限也可能浪费。
Static batching 的本质问题是:参数是离线拍的,运行时不感知负载变化。
三、Adaptive Batching 的核心算法
Adaptive Batching 的核心思路:等待时间随队列深度动态调整。
队列深(请求堆积)→ 说明请求速率高 → 多等一会儿能凑更大的 batch → 吞吐收益高。
队列浅(请求稀疏)→ 等下去也凑不了多少 → 尽快发出,减少等待延迟。
3.1 基础公式
actual_wait_ms = max_wait_ms × (1 - 1 / (1 + queue_depth / scale_factor))
当 queue_depth = 0:actual_wait = 0,立即发出。
当 queue_depth = scale_factor:actual_wait = max_wait * 0.5,等一半时间。
当 queue_depth >> scale_factor:actual_wait → max_wait,等满最长时间。
这是一个平滑的自适应曲线,不需要分段 if-else。
3.2 双触发条件(实践中最常用)
无论 adaptive wait 结果如何,任意一个条件满足就立即 flush:
- Size trigger:
queue_depth >= max_batch_size - Time trigger: 队列中最老的请求等待时间 >=
actual_wait_ms
should_flush = (
len(self.queue) >= self.max_batch_size
or (self.queue and time.monotonic() - self.queue[0].enqueue_time >= actual_wait)
)
四、完整工程实现
4.1 Python asyncio 版(适合 FastAPI/aiohttp 服务)
import asyncio
import time
from dataclasses import dataclass, field
from typing import Any, List, Optional, Callable, Awaitable
from collections import deque
@dataclass
class BatchItem:
payload: Any
future: asyncio.Future
enqueue_time: float = field(default_factory=time.monotonic)
class AdaptiveBatcher:
"""
应用层 Adaptive Batching 实现。
参数说明:
max_batch_size: 单批最大数量,超过立即 flush
max_wait_ms: 最长等待时间(ms),P99 延迟上界
min_batch_size: 最小批量,低于此不触发 size-based flush
scale_factor: 自适应曲线参数,queue_depth=scale_factor 时等待约 0.5*max_wait
process_fn: 接收 List[payload],返回 List[result](顺序必须对应)
"""
def __init__(
self,
process_fn: Callable[[List[Any]], Awaitable[List[Any]]],
max_batch_size: int = 100,
max_wait_ms: float = 30.0,
min_batch_size: int = 1,
scale_factor: int = 10,
):
self.process_fn = process_fn
self.max_batch_size = max_batch_size
self.max_wait_ms = max_wait_ms
self.min_batch_size = min_batch_size
self.scale_factor = scale_factor
self._queue: deque[BatchItem] = deque()
self._lock = asyncio.Lock()
self._flush_task: Optional[asyncio.Task] = None
# 监控计数器
self._stats = {
"total_items": 0,
"total_batches": 0,
"size_triggered": 0,
"time_triggered": 0,
}
def _adaptive_wait_ms(self) -> float:
"""根据当前队列深度计算自适应等待时间"""
depth = len(self._queue)
if depth == 0:
return 0.0
return self.max_wait_ms * (1 - 1 / (1 + depth / self.scale_factor))
async def add(self, payload: Any) -> Any:
"""添加一个请求到批处理队列,返回处理结果"""
loop = asyncio.get_event_loop()
future: asyncio.Future = loop.create_future()
item = BatchItem(payload=payload, future=future)
async with self._lock:
self._queue.append(item)
self._stats["total_items"] += 1
if len(self._queue) >= self.max_batch_size:
# size trigger:立即 flush,取消已有的定时 flush
if self._flush_task and not self._flush_task.done():
self._flush_task.cancel()
self._flush_task = None
asyncio.create_task(self._do_flush("size"))
elif self._flush_task is None or self._flush_task.done():
# 没有等待中的 flush task,创建一个
self._flush_task = asyncio.create_task(self._schedule_flush())
return await future
async def _schedule_flush(self):
"""等待 adaptive wait 时间后 flush"""
wait_ms = self._adaptive_wait_ms()
if wait_ms > 0:
await asyncio.sleep(wait_ms / 1000)
async with self._lock:
if self._queue:
await self._do_flush("time")
self._flush_task = None
async def _do_flush(self, reason: str):
"""执行一次批处理"""
if not self._queue:
return
# 取出当前队列的全部(或 max_batch_size 个)
batch_items: List[BatchItem] = []
while self._queue and len(batch_items) < self.max_batch_size:
batch_items.append(self._queue.popleft())
self._stats["total_batches"] += 1
self._stats[f"{reason}_triggered"] += 1
payloads = [item.payload for item in batch_items]
try:
results = await self.process_fn(payloads)
if len(results) != len(batch_items):
raise ValueError(
f"process_fn returned {len(results)} results for {len(batch_items)} items"
)
for item, result in zip(batch_items, results):
if not item.future.done():
item.future.set_result(result)
except Exception as e:
for item in batch_items:
if not item.future.done():
item.future.set_exception(e)
def get_stats(self) -> dict:
return {
**self._stats,
"queue_depth": len(self._queue),
"avg_batch_size": (
self._stats["total_items"] / self._stats["total_batches"]
if self._stats["total_batches"] > 0
else 0
),
}
4.2 接入大模型 Embedding API
# 以国产大模型 API(兼容接口)为例
from openai import AsyncOpenAI
client = AsyncOpenAI(
base_url="https://api.therouter.io/v1", # 统一网关
api_key="your_api_key"
)
async def embed_batch(texts: List[str]) -> List[List[float]]:
resp = await client.embeddings.create(
model="qwen/text-embedding-v3", # 国产 Embedding 模型
input=texts,
)
# 保证顺序与输入一致
return [item.embedding for item in sorted(resp.data, key=lambda x: x.index)]
# 创建 batcher
embedding_batcher = AdaptiveBatcher(
process_fn=embed_batch,
max_batch_size=200,
max_wait_ms=25.0, # P99 额外延迟上界 25ms
scale_factor=15,
)
# 在请求处理函数里使用
async def get_embedding(text: str) -> List[float]:
return await embedding_batcher.add(text)
4.3 Node.js 版:DataLoader 模式
Facebook DataLoader 是应用层 batching 的经典实现,GraphQL 社区广泛使用:
import DataLoader from 'dataloader';
import OpenAI from 'openai';
const client = new OpenAI({
baseURL: 'https://api.therouter.io/v1',
apiKey: process.env.API_KEY,
});
const embeddingLoader = new DataLoader<string, number[]>(
async (texts: readonly string[]) => {
const response = await client.embeddings.create({
model: 'qwen/text-embedding-v3',
input: texts as string[],
});
// DataLoader 要求返回数组与输入一一对应
return response.data
.sort((a, b) => a.index - b.index)
.map(item => item.embedding);
},
{
maxBatchSize: 100,
// batchScheduleFn 控制等待时间:20ms 后发出,不等满
batchScheduleFn: (callback) => setTimeout(callback, 20),
// 相同 text 自动去重(cacheKeyFn 可自定义)
cache: true,
}
);
// 使用:每次调用 load,DataLoader 自动合批
const embedding = await embeddingLoader.load(text);
// 多个并发调用自动合批:
const [emb1, emb2, emb3] = await Promise.all([
embeddingLoader.load(text1),
embeddingLoader.load(text2),
embeddingLoader.load(text3),
]);
DataLoader 的 batchScheduleFn 对应 adaptive batching 的时间触发器,maxBatchSize 对应 size 触发器。
五、5 个生产踩坑
坑 1:批内一个请求超时,整批受牵连
现象:某条特别长的文本(3000 字)和 99 条短文本合批,Embedding API 响应超时,100 条请求全部失败重试。
根因:process_fn 整体超时,没有批内独立的 per-item deadline。
解法:
async def embed_batch_with_timeout(texts: List[str]) -> List[List[float]]:
try:
resp = await asyncio.wait_for(
client.embeddings.create(model="qwen/text-embedding-v3", input=texts),
timeout=10.0 # 整批最多等 10s
)
return [item.embedding for item in sorted(resp.data, key=lambda x: x.index)]
except asyncio.TimeoutError:
# 批级超时:把批拆小,独立重试
if len(texts) == 1:
raise # 单条还超时,真的有问题
mid = len(texts) // 2
left, right = await asyncio.gather(
embed_batch_with_timeout(texts[:mid]),
embed_batch_with_timeout(texts[mid:]),
)
return left + right
坑 2:重试时整批重发,浪费成功的结果
现象:100 条请求里有 3 条因为 token 超限失败(文本太长),重试时把 100 条全部重发,97 条被重复计费。
解法:在 _do_flush 里记录每条 item 的结果,只对 future.done() == False 的 item 重试:
# 批内精细重试
failed_items = [
item for item, result in zip(batch_items, results)
if isinstance(result, Exception)
]
if failed_items:
retry_payloads = [item.payload for item in failed_items]
retry_results = await self.process_fn(retry_payloads)
for item, result in zip(failed_items, retry_results):
item.future.set_result(result)
坑 3:高优先级请求被低优先级大批堵住
现象:用户实时搜索请求(高优先级)被后台 indexing 任务(低优先级,大批量)堵在队列里,响应时间 P99 从 100ms 涨到 800ms。
解法:分优先级维护独立队列,flush 时优先取高优先级:
from enum import IntEnum
class Priority(IntEnum):
HIGH = 0 # 用户实时请求
LOW = 1 # 后台任务
class PriorityAdaptiveBatcher:
def __init__(self, ...):
self._queues = {
Priority.HIGH: deque(),
Priority.LOW: deque(),
}
def _next_batch(self) -> List[BatchItem]:
batch = []
# 先取高优先级
for priority in sorted(Priority):
q = self._queues[priority]
while q and len(batch) < self.max_batch_size:
batch.append(q.popleft())
if len(batch) >= self.max_batch_size:
break
return batch
坑 4:队列深度监控缺失,背压传递失效
现象:下游大模型 API 触发 rate limit,处理变慢,队列持续增长,内存 OOM。应用层没有任何报警,直到服务崩溃。
解法:在 add() 入口做队列深度检查,超过阈值直接返回 429:
async def add(self, payload: Any, priority: Priority = Priority.HIGH) -> Any:
queue_depth = sum(len(q) for q in self._queues.values())
if queue_depth >= self.max_queue_depth:
raise BatcherOverloadError(
f"Queue depth {queue_depth} exceeds limit {self.max_queue_depth}"
)
# ... 正常入队逻辑
配合 Prometheus 指标:
QUEUE_DEPTH = Gauge('batcher_queue_depth', 'Current queue depth', ['batcher_name'])
BATCH_SIZE = Histogram('batcher_batch_size', 'Batch sizes', buckets=[1,5,10,20,50,100,200])
WAIT_TIME_MS = Histogram('batcher_wait_ms', 'Item wait time in ms', buckets=[1,5,10,25,50,100,200,500])
坑 5:不同模型请求被错误合并
现象:系统同时使用两种 Embedding 模型,混合进同一个 batcher,model 参数用了第一个请求的,后续请求拿到的是错误模型的 embedding。
解法:按 (model, encoding_format) 分桶,每个桶独立 batcher:
from functools import lru_cache
@lru_cache(maxsize=16)
def get_batcher(model: str, encoding_format: str = "float") -> AdaptiveBatcher:
async def _process(texts):
resp = await client.embeddings.create(
model=model, input=texts, encoding_format=encoding_format
)
return [item.embedding for item in sorted(resp.data, key=lambda x: x.index)]
return AdaptiveBatcher(process_fn=_process, max_batch_size=200, max_wait_ms=25)
# 调用时指定 model
embedding = await get_batcher("qwen/text-embedding-v3").add(text)
六、性能数据:实测对比
测试场景:1000 条文本(平均 120 字),并发度 50,目标 P99 < 150ms。
| 策略 | 吞吐(items/s) | P50 延迟 | P99 延迟 | API 调用次数 |
|---|---|---|---|---|
| 无 batching(串行) | 220 | 45ms | 95ms | 1000 |
| 固定 batch=50,wait=50ms | 580 | 52ms | 108ms | 20 |
| 固定 batch=100,wait=100ms | 820 | 102ms | 215ms | 10 |
| Adaptive(max_wait=25ms,max_batch=200) | 1650 | 30ms | 128ms | 6-8 |
结论:
- 固定 batch=100 吞吐提升明显,但 P99 超过 200ms,不适合实时场景
- Adaptive batching 在保持 P99 < 150ms 的前提下,吞吐是无 batching 的 7.5 倍,是最优固定策略(wait=50ms)的 2.8 倍
- API 调用次数从 1000 次降到 6-8 次,成本下降 99%
参数调优建议:
max_wait_ms = target_p99_latency_ms × 0.15 # 等待时间不超过 P99 目标的 15%
max_batch_size = rate_limit_per_minute / 60 / expected_batches_per_second
scale_factor = max_batch_size / 5 # 队列到 max 的 20% 时开始明显延迟
七、监控与告警设计
7.1 核心指标
# 四个必须跟踪的指标
metrics = {
# 1. 队列深度(leading indicator,比延迟先报警)
"batcher_queue_depth": Gauge,
# 2. 批量大小分布(判断参数是否合理)
"batcher_batch_size": Histogram, # buckets: [1,5,10,25,50,100,200]
# 3. 每个 item 的等待时间(用户感知延迟的组成部分)
"batcher_item_wait_ms": Histogram, # buckets: [1,5,10,25,50,100,200,500]
# 4. flush 触发原因(size vs time,诊断策略有效性)
"batcher_flush_total": Counter, # labels: reason=[size,time]
}
7.2 告警规则
# Prometheus alerting rules
groups:
- name: adaptive_batcher
rules:
- alert: BatcherQueueDepthHigh
expr: batcher_queue_depth > 500
for: 30s
annotations:
summary: "Batcher queue 堆积超过 500,可能下游限速或崩溃"
- alert: BatcherP99WaitHigh
expr: histogram_quantile(0.99, batcher_item_wait_ms) > 200
for: 1m
annotations:
summary: "Batcher P99 等待时间超过 200ms,检查 max_wait_ms 参数"
- alert: BatcherTimeTriggerRatioHigh
expr: |
rate(batcher_flush_total{reason="time"}[5m]) /
rate(batcher_flush_total[5m]) > 0.8
for: 5m
annotations:
summary: "80% 的 flush 由超时触发,考虑减小 max_batch_size 或增大 max_wait_ms"
八、什么时候不该用 Adaptive Batching
Adaptive Batching 不是万能药,以下场景不适合:
1. 请求之间有强依赖:A 的结果是 B 的输入,没法合批,合了也没意义。
2. 每条请求的 payload 差异极大:有的 1000 token,有的 10 token,合批后某条必然超限,分裂重试成本反而更高。
3. 结果顺序无法保证的 process_fn:如果你的批处理函数不能保证返回结果与输入顺序一一对应,会导致结果错位(这是个常见 bug)。
4. 已经在用服务端 Continuous Batching:vLLM/TGI 在推理层已经做了 iteration-level 调度,应用层再加一层 adaptive wait 反而引入不必要的延迟。应用层 adaptive batching 主要针对的是你调用外部 API 的场景,而不是你自己部署的推理服务。
总结
回到开头的问题:吞吐和延迟的矛盾,核心解法是让 batch 策略感知负载变化。
Adaptive Batching 的工程要点:
- 自适应等待时间:队列深度决定等多久,不是固定参数拍脑袋
- 双触发条件:size OR time,先到先发,避免两个极端
- 批内独立 deadline:超时不能把整批拖死
- 优先级分桶:高低优先级隔离,防止后台任务堵塞实时请求
- 按 (model, format) 分 batcher 实例:不同模型的请求不能混批
- 背压监控:queue_depth 是最早的危险信号,P99 wait 是最终表现
代码量不大(Python 版核心逻辑约 80 行),但每个细节都对应一个生产踩坑。DataLoader 模式适合 Node.js 生态,asyncio 版适合 Python FastAPI 服务,两者设计思路一致,选适合自己技术栈的就行。