欢迎光临

Continuous Batching 动态批处理深度解析:如何突破 LLM 推理服务的吞吐量瓶颈

引言:LLM 推理服务的吞吐量困局

当我们将一个大语言模型部署到生产环境时,最核心的挑战往往不是模型本身的精度,而是如何高效地利用 GPU 算力服务海量并发请求。一个 7B 参数的模型在单张 A100 上串行推理,每个请求可能需要数百毫秒到数秒;但如果能将多个请求合并处理,吞吐量可以提升一个数量级以上。这就是批处理(Batching)技术的价值所在。

然而,传统的静态批处理在面对 LLM 这种自回归生成任务时存在严重的效率问题。请求的序列长度差异巨大、生成时间不可预测,导致 GPU 利用率常常低于 30%。Continuous Batching(连续批处理)作为现代 LLM 推理框架(vLLM、TGI、TensorRT-LLM)的核心调度技术,通过迭代级别的动态调度彻底解决了这一瓶颈。本文将从原理到实现,深度解析这一关键技术。

传统静态批处理的局限性

静态批处理的工作方式

在传统的深度学习推理中,批处理的方式很简单:将多个输入样本组成一个 batch,一次性送入模型前向传播,然后收集所有输出。这种方式在图像分类、目标检测等固定长度输出的场景中工作良好,但在 LLM 推理中却遇到了根本性困难。

以下是一个静态批处理的典型流程:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
# 传统静态批处理伪代码
def static_batch_serve(model, requests):
    # 等待凑够 batch_size 或超时
    batch = wait_for_batch(requests, max_batch=8, timeout=50ms)

    # 统一 padding 到最长序列
    padded_input = pad_sequences(batch, max_len=max(len(r) for r in batch))

    # 整体前向传播,直到所有序列生成完毕
    while not all_finished(batch):
        logits = model.forward(padded_input)
        next_tokens = sample(logits)
        # 已经完成的序列继续空转,浪费算力
        for r in batch:
            if not r.finished:
                r.append(next_tokens)

    return [r.output for r in batch]

“护航问题”(Convoy Problem)

静态批处理的核心痛点被称为护航问题。想象一支车队行驶在公路上——车队的速度取决于最慢的车辆。在 LLM 推理中,如果一个 batch 中有一个请求需要生成 500 个 token,而其他请求只需要 50 个 token,那么已经完成的请求必须”空转”等待最慢的请求完成,这段时间的 GPU 计算完全是浪费。

具体来说,静态批处理存在以下三大问题:

  • Padding 浪费:不同请求的输入长度差异可能达到 10 倍以上,统一 padding 到最长序列会导致大量计算浪费在 padding 位置上
  • 队头阻塞:一个长生成请求会拖慢整个 batch 的响应时间,短请求的用户体验被严重损害
  • GPU 空转:已完成的请求仍在 batch 中占据位置,GPU 在这些位置上进行无效计算
  • 延迟与吞吐无法兼得:为了提高吞吐量需要增大 batch size,但这会增加单个请求的排队延迟

实测数据表明,在请求长度分布不均匀的真实场景下,静态批处理的 GPU 利用率往往只有 20%-40%。这意味着你花了大量金钱购买的 GPU,有超过一半的时间在”空转”。

Continuous Batching 核心原理

迭代级别调度的核心思想

Continuous Batching 的核心创新在于将调度的粒度从请求级别降低到迭代(iteration)级别。在静态批处理中,一个请求从加入 batch 到离开 batch 是一个完整的生成过程;而在 Continuous Batching 中,调度器在每一次 token 生成迭代之后都重新评估 batch 的组成。

具体来说,Continuous Batching 的工作流程如下:

  1. 调度器维护一个请求队列和一个正在运行的 batch
  2. 每次迭代(生成一个 token)完成后,检查哪些请求已经生成完毕
  3. 将已完成的请求从 batch 中移除,立即返回结果给客户端
  4. 从队列中取出新请求加入 batch,在下一个迭代中开始生成
  5. 重复上述过程,直到队列为空且 batch 中所有请求完成

这意味着在任何时刻,batch 中的请求都可以动态变化——旧的请求完成后腾出位置,新的请求立即填补进来,GPU 始终保持高效运转。

与静态批处理的对比

特性 静态批处理 Continuous Batching
调度粒度 请求级别(整个生成过程) 迭代级别(每个 token)
请求加入时机 仅 batch 开始时 任意迭代间隙
请求离开时机 整个 batch 完成后 该请求完成后立即
GPU 利用率 20%-40% 70%-90%
P99 延迟 受最慢请求拖累 接近单请求延迟
Padding 开销 高(需统一 padding) 低(按需分配)

Continuous Batching 的工程实现

核心数据结构设计

要实现一个高效的 Continuous Batching 调度器,需要精心设计几个核心数据结构。以下是一个简化但完整的 Python 实现框架:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
import torch
from dataclasses import dataclass, field
from typing import List, Optional
from collections import deque

@dataclass
class RequestState:
    """单个推理请求的状态"""
    request_id: str
    input_ids: List[int]          # 输入 token IDs
    output_ids: List[int] = field(default_factory=list)  # 已生成的 token IDs
    max_tokens: int = 512         # 最大生成长度
    finished: bool = False        # 是否已完成(遇到 EOS 或达到最大长度)

    @property
    def all_ids(self):
        return self.input_ids + self.output_ids

    @property
    def current_len(self):
        return len(self.all_ids)

class ContinuousBatchScheduler:
    """Continuous Batching 调度器核心实现"""

    def __init__(self, model, max_batch_size=32, max_seq_len=2048):
        self.model = model
        self.max_batch_size = max_batch_size
        self.max_seq_len = max_seq_len
        self.waiting_queue = deque()   # 等待加入的请求
        self.running_batch = []        # 当前正在运行的请求

    def add_request(self, request: RequestState):
        """添加新请求到等待队列"""
        self.waiting_queue.append(request)

    def _can_admit(self) -> bool:
        """检查是否有空闲 slot 可以接纳新请求"""
        return (len(self.running_batch) < self.max_batch_size and
                len(self.waiting_queue) > 0)

    def _admit_requests(self):
        """从等待队列中接纳新请求到运行 batch"""
        while self._can_admit():
            # 检查新请求加入后是否会超出最大序列长度限制
            req = self.waiting_queue[0]
            max_current = max(
                (r.current_len for r in self.running_batch),
                default=0
            )
            if max(max_current, req.current_len) > self.max_seq_len:
                break  # 超出限制,停止接纳

            self.waiting_queue.popleft()
            self.running_batch.append(req)

    def _evict_finished(self):
        """移除已完成的请求"""
        finished = [r for r in self.running_batch if r.finished]
        self.running_batch = [r for r in self.running_batch if not r.finished]
        return finished

    def step(self) -> List[RequestState]:
        """执行一次迭代:生成一个 token 并管理 batch"""
        if not self.running_batch and not self.waiting_queue:
            return []

        # 1. 接纳新请求
        self._admit_requests()

        if not self.running_batch:
            return []

        # 2. 准备输入(动态组装当前 batch 的所有序列)
        seqs = [r.all_ids for r in self.running_batch]
        max_len = max(len(s) for s in seqs)

        # 左 padding 对齐(生成任务通常用左 padding)
        input_ids = torch.full(
            (len(seqs), max_len),
            fill_value=self.model.pad_token_id,
            dtype=torch.long
        )
        attention_mask = torch.zeros(len(seqs), max_len)

        for i, seq in enumerate(seqs):
            offset = max_len - len(seq)
            input_ids[i, offset:] = torch.tensor(seq)
            attention_mask[i, offset:] = 1.0

        # 3. 前向传播,获取下一个 token
        with torch.no_grad():
            logits = self.model(input_ids.cuda(), attention_mask.cuda())
            next_token_logits = logits[:, -1, :]  # 取最后一个位置
            next_tokens = torch.argmax(next_token_logits, dim=-1)

        # 4. 将生成的 token 分配给对应的请求
        for i, req in enumerate(self.running_batch):
            token_id = next_tokens[i].item()
            req.output_ids.append(token_id)

            # 检查是否完成
            if (token_id == self.model.eos_token_id or
                len(req.output_ids) >= req.max_tokens):
                req.finished = True

        # 5. 移除已完成请求并返回
        finished = self._evict_finished()
        return finished

    def has_pending(self) -> bool:
        return len(self.running_batch) > 0 or len(self.waiting_queue) > 0

关键实现细节解析

上述实现中有几个关键的工程细节需要特别注意:

1. 左 Padding 策略。在生成任务中,模型需要从序列的末尾预测下一个 token。如果使用右 padding,padding token 会出现在序列末尾,模型将尝试在 padding 位置之后进行预测,导致错误。左 padding 将 padding 放在序列开头,保持有效 token 在末尾,这样模型可以直接取最后一个位置的 logits 进行预测。

2. 动态序列组装。每次迭代时,batch 中的请求可能已经变化(有新加入的、有完成的),因此不能复用上一次的输入张量。每个迭代都需要重新组装输入,这是 Continuous Batching 的一个固有开销。现代框架通过预分配内存池来减少频繁分配的开销。

3. 完成条件判断。一个请求在两种情况下被认为完成:生成了 EOS(End of Sequence)token,或者达到了用户指定的最大生成长度。调度器必须在每次迭代后检查所有请求的完成状态。

与 PagedAttention 的协同工作机制

Continuous Batching 解决的是调度层面的问题——何时将哪些请求放入 batch。但在实际运行中还有一个内存层面的挑战:LLM 推理的 KV Cache 会随序列长度线性增长,而不同请求的 KV Cache 大小各不相同,传统的连续内存分配方式会导致严重的碎片化。

这就是 PagedAttention 登场的地方。PagedAttention 借鉴了操作系统的虚拟内存分页机制,将 KV Cache 划分为固定大小的块(block),每个块可以独立分配和回收。这两项技术的协同工作方式如下:

  • Continuous Batching负责动态管理 batch 中的请求——完成后立即移除,新请求随时加入
  • PagedAttention负责动态管理 KV Cache 的内存——每个请求的 KV Cache 按需分块分配,请求完成后立即回收
  • 两者配合实现了零碎片化的内存管理和零空转的 GPU 调度

在 vLLM 的实现中,这两者的结合使得同一张 GPU 上可以并发处理的请求数量提升了 2-4 倍,整体吞吐量相比传统方式提升了 10-20 倍。以下是在不同负载下的性能对比:

场景 静态批处理 (tokens/s) Continuous Batching (tokens/s) 提升倍数
均匀短序列 (avg 50 tokens) 1,200 4,800 4.0x
不均匀混合 (50-500 tokens) 680 5,200 7.6x
长序列为主 (500-2000 tokens) 420 3,900 9.3x
高并发突发流量 380 4,600 12.1x

可以看到,负载越不均匀、并发越高,Continuous Batching 的优势越明显。这正是因为它消除了护航问题和 padding 浪费,而这些问题在不均匀负载下最为严重。

生产环境部署实践

主流框架对比与选型

目前支持 Continuous Batching 的主流推理框架有三个,它们在实现方式和适用场景上各有侧重:

框架 Continuous Batching PagedAttention 量化支持 适用场景
vLLM 原生支持 原生支持 AWQ/GPTQ/FP8 通用 LLM 服务部署
TGI (HuggingFace) 原生支持 支持 (v0.9+) bitsandbytes/EETQ HuggingFace 生态集成
TensorRT-LLM 原生支持 支持 INT8/FP8 NVIDIA 极致性能

vLLM 部署配置详解

以 vLLM 为例,以下是一个生产环境的部署配置,关键参数直接影响 Continuous Batching 的效果:


1
2
3
4
5
6
7
8
9
10
11
12
# vLLM 服务启动配置
python -m vllm.entrypoints.openai.api_server \
    --model meta-llama/Llama-2-7b-chat-hf \
    --tensor-parallel-size 1 \
    --gpu-memory-utilization 0.90 \
    --max-num-batched-tokens 8192 \
    --max-num-seqs 256 \
    --max-model-len 4096 \
    --enable-chunked-prefill \
    --max-num-partial-prefills 4 \
    --max-long-partial-prefills 1 \
    --port 8000

各关键参数的含义和调优建议:

  • gpu-memory-utilization (0.90):预留 10% 显存给 PyTorch 运行时开销。设置过高可能导致 OOM,过低则浪费显存。建议从 0.85 开始逐步上调
  • max-num-batched-tokens (8192):单次迭代处理的最大 token 总数。这是 Continuous Batching 的核心参数——它决定了每次迭代可以同时处理多少 token。增大此值可提升吞吐量但增加延迟
  • max-num-seqs (256):同时运行的最大序列数。这是 batch size 的上限,取决于 KV Cache 的内存预算
  • enable-chunked-prefill:启用分块预填充。长 prompt 的 prefill 阶段会被拆分成多个 chunk,与 decode 请求混在一起调度,避免长 prompt 阻塞短请求
  • max-num-partial-prefills (4):同时进行的分块预填充数量。控制 prefill 和 decode 之间的资源分配比例

性能监控与调优

部署后,需要持续监控以下关键指标来验证 Continuous Batching 是否正常工作:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
import time
import requests
from collections import defaultdict

class InferenceMetrics:
    """推理服务性能监控"""

    def __init__(self):
        self.request_latencies = defaultdict(list)
        self.token_throughput = []
        self.batch_sizes = []

    def record_request(self, endpoint, latency, tokens_generated):
        self.request_latencies[endpoint].append(latency)
        self.token_throughput.append(tokens_generated / latency)

    def get_stats(self):
        all_latencies = []
        for lats in self.request_latencies.values():
            all_latencies.extend(lats)

        all_latencies.sort()
        return {
            'p50_latency_ms': all_latencies[len(all_latencies)//2] * 1000,
            'p99_latency_ms': all_latencies[int(len(all_latencies)*0.99)] * 1000,
            'avg_throughput_tokens_per_s': sum(self.token_throughput) / len(self.token_throughput),
            'total_requests': len(all_latencies),
        }

# 使用示例
metrics = InferenceMetrics()

# 发送并发请求并记录指标
def benchmark_concurrent_load(url, num_requests=100):
    import concurrent.futures

    def send_request(i):
        start = time.time()
        response = requests.post(url, json={
            "model": "llama-2-7b-chat",
            "messages": [{"role": "user", "content": f"Explain concept {i}"}],
            "max_tokens": 256
        })
        latency = time.time() - start

        # 从响应中获取生成的 token 数
        result = response.json()
        tokens = result.get('usage', {}).get('completion_tokens', 256)
        metrics.record_request('/v1/chat/completions', latency, tokens)
        return latency

    with concurrent.futures.ThreadPoolExecutor(max_workers=50) as executor:
        futures = [executor.submit(send_request, i) for i in range(num_requests)]
        concurrent.futures.wait(futures)

    stats = metrics.get_stats()
    print(f"=== Continuous Batching 性能报告 ===")
    print(f"总请求数: {stats['total_requests']}")
    print(f"P50 延迟: {stats['p50_latency_ms']:.1f}ms")
    print(f"P99 延迟: {stats['p99_latency_ms']:.1f}ms")
    print(f"平均吞吐: {stats['avg_throughput_tokens_per_s']:.1f} tokens/s")

如果 P99 延远高于 P50(比如 10 倍以上),说明可能存在调度瓶颈,需要调整

1
max-num-batched-tokens

1
max-num-seqs

参数。理想状态下,P99/P50 比值应该控制在 3 倍以内。

Chunked Prefill:Continuous Batching 的进阶优化

在标准的 Continuous Batching 中,存在一个尚未解决的矛盾:Prefill 阶段和 Decode 阶段的计算特性完全不同。Prefill 需要处理整个 prompt(可能数千个 token),是一次计算密集型操作;Decode 每次只生成一个 token,是访存密集型操作。如果一个长 prompt 的 prefill 正在进行,它会占用大量 GPU 计算资源,导致正在 decode 的请求被”饿死”。

Chunked Prefill 技术将这个问题优雅地解决了。它的核心思想是:将长 prompt 的 prefill 拆分成多个固定大小的 chunk,每个 chunk 与 decode 请求一起在一个迭代中处理。这样长 prompt 不会独占 GPU,decode 请求的延迟也能得到保障。


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
# Chunked Prefill 调度逻辑(简化版)
def schedule_with_chunked_prefill(running_batch, waiting_queue,
                                   max_batched_tokens=8192,
                                   chunk_size=512):
    """带分块预填充的调度器"""
    scheduled = []
    remaining_tokens = max_batched_tokens

    # 优先调度 decode 请求(延迟敏感)
    for req in running_batch:
        if not req.prefill_done:
            continue
        scheduled.append(('decode', req, 1))  # decode 每次 1 token
        remaining_tokens -= 1

    # 然后调度 prefill 的 chunk
    for req in waiting_queue + running_batch:
        if req.prefill_done:
            continue

        # 计算本次迭代能处理的 chunk 大小
        chunk = min(chunk_size, len(req.input_ids) - req.prefill_offset,
                    remaining_tokens)
        if chunk <= 0:
            continue

        scheduled.append(('prefill_chunk', req, chunk))
        req.prefill_offset += chunk
        remaining_tokens -= chunk

        if req.prefill_offset >= len(req.input_ids):
            req.prefill_done = True

        if remaining_tokens <= 0:
            break

    return scheduled

这种策略确保了 prefill 和 decode 请求在每次迭代中都能获得公平的计算资源分配。在实际测试中,启用 Chunked Prefill 后,Decode 请求的 P99 延迟可以降低 40%-60%,而整体吞吐量几乎不受影响。

总结与展望

Continuous Batching 是现代 LLM 推理服务的基石技术。它通过将调度粒度从请求级别降低到迭代级别,彻底消除了传统静态批处理的护航问题和 GPU 空转问题,在实际场景中可带来 5-15 倍的吞吐量提升。

回顾本文的核心要点:

  • 传统静态批处理在 LLM 推理中存在严重的 GPU 利用率问题(20%-40%),核心原因是护航问题和 padding 浪费
  • Continuous Batching通过迭代级别的动态调度,让请求在任意 token 生成完成后立即离开 batch,新请求随时加入,将 GPU 利用率提升到 70%-90%
  • 与 PagedAttention 协同,从调度和内存两个层面消除浪费,实现 10-20 倍的整体性能提升
  • Chunked Prefill进一步解决了长 prompt 阻塞 decode 请求的问题,在保持吞吐量的同时大幅降低尾延迟
  • 在生产部署中,
    1
    max-num-batched-tokens

    1
    max-num-seqs

    是最关键的两个调优参数

展望未来,Continuous Batching 技术仍在持续演进。Speculative Decoding(投机解码)与 Continuous Batching 的结合、跨 GPU 的 Distributed Continuous Batching、以及针对多模态模型的混合调度,都是当前前沿的研究方向。对于任何需要大规模部署 LLM 的团队来说,深入理解 Continuous Batching 的原理和调优方法,都是不可或缺的核心能力。

【本站文章皆为原创,未经允许不得转载】:汤不热吧 » Continuous Batching 动态批处理深度解析:如何突破 LLM 推理服务的吞吐量瓶颈
分享到: 更多 (0)