AI推理引擎 Continuous Batching 深度工程实战:从 PagedAttention 到 SGLang 架构

在大模型推理服务中,如何同时实现高吞吐和低延迟?这是所有推理引擎设计的核心矛盾。Continuous Batching(也叫 Iteration-Level Scheduling)正是解决这一矛盾的关键技术——它让推理引擎能够在生成过程中动态插入和移除请求,彻底改变了几静态批处理的效率瓶颈。本文将深入剖析 Continuous Batching 的底层原理,分析 PagedAttention 的内存管理创新,并以 SGLang 和 vLLM 为案例,展示生产级推理引擎的架构设计与优化实践。

一、为什么需要 Continuous Batching

1.1 静态批处理的根本性缺陷

在大模型推理服务中,请求长度不可预测是天然特性。用户 A 的问题可能只需要生成 50 个 token,用户 B 的请求需要生成 2000 个 token。在静态批处理模式下,一个批次的完成时间取决于最长的那个请求。

假设一个 batch 中有 4 个请求,分别需要生成 100、200、50、500 个 token:


时间线:
请求1: ████████████████████ (100 tokens)
请求2: ████████████████████████████████████████ (200 tokens)
请求3: ██████████ (50 tokens)
请求4: ██████████████████████████████████████████████████████████████████ (500 tokens)
                                          ↑
                         批处理完成点:500 tokens 全部完成后才能释放资源

这意味着请求 1 在第 100 个 token 生成完毕后,仍然占用着 GPU 显存和计算资源,等待请求 4 完成。在在线服务场景中,这种等待造成了严重的资源浪费和延迟累积。

1.2 Continuous Batching 的核心思路

Continuous Batching 在每一轮 decode iteration(而非整个序列生成完成后)重新调度 batch 中的请求:


Iteration-Level 调度示意:

Batch at step t:
  请求1 (gen 20/100)  ← 刚加入
  请求2 (gen 150/200) ← 已在运行
  请求3 (gen  8/50)   ← 即将完成

Batch at step t+1:     ← 请求3完成后立即释放资源,加入新请求5
  请求1 (gen 21/100)
  请求2 (gen 151/200)
  请求5 (gen  1/300)  ← 新请求立即抢占空位

关键区别:

  • 调度粒度:从 "sequence-level" 变为 "iteration-level"
  • 资源粒度:从 "batch-level fixed" 变为 "per-request dynamic"
  • 显存管理:从 "预先分配最大长度" 变为 "按需分配 / 释放"

二、PagedAttention:解决 KV Cache 内存碎片

Continuous Batching 的实现依赖于高效的 KV Cache 管理。PagedAttention 借鉴操作系统虚拟内存分页机制,将 KV Cache 从连续内存块改为非连续页表管理。

2.1 传统 KV Cache 的内存浪费

在传统方案中,每个请求的 KV Cache 需要预先分配最大序列长度(如 2048 tokens)的连续空间:


传统方案内存布局:
| 请求1: 预分配2048 slots | 请求2: 预分配2048 slots | 请求3: 预分配2048 slots |

问题:
1. 内部碎片:实际只用500个token,浪费1548个slot的空间
2. 外部碎片:释放后产生不连续空洞,分配大序列时失败
3. 预留膨胀:为安全起见,通常预留20%额外空间

实验数据显示,vLLM 论文中指出传统方案的内存浪费可达 60%~90%——这是巨大的 GPU 显存浪费。

2.2 分页实现原理

PagedAttention 将 KV Cache 切分为固定大小的页(如 16 tokens/page),通过逻辑页号到物理页帧的映射来管理:


# 核心数据结构简化示意
class PagedAttentionManager:
    def __init__(self, page_size: int = 16, num_pages: int = 100000):
        self.page_size = page_size
        self.free_pages = list(range(num_pages))  # 空闲页池
        self.page_tables: dict[int, list[int]] = {}  # 请求ID -> 物理页列表
    
    def alloc(self, request_id: int, num_tokens: int) -> list[int]:
        """为请求分配KV Cache页"""
        pages_needed = ceil(num_tokens / self.page_size)
        pages = self.free_pages[:pages_needed]
        self.free_pages = self.free_pages[pages_needed:]
        self.page_tables[request_id] = pages
        return pages
    
    def append_token(self, request_id: int) -> int:
        """追加一个token,必要时分配新页"""
        pages = self.page_tables[request_id]
        used_slots = self._used_slots(request_id)
        if used_slots % self.page_size == 0:
            # 当前页已满,申请新页
            new_page = self.free_pages.pop(0)
            pages.append(new_page)
        return pages[-1] * self.page_size + (used_slots % self.page_size)
    
    def free(self, request_id: int):
        """释放请求占用的所有页"""
        self.free_pages.extend(self.page_tables[request_id])
        del self.page_tables[request_id]

2.3 FlashAttention 与分页的协同

在 Attention 计算中,不同请求的 KV Cache 长度不一且物理地址不连续。FlashAttention 的 tiling 策略天然适配分页访问:


正常 Attention: O(N²) 显存,需要连续KV
FlashAttention: O(N) 显存,通过tiling分块计算

FlashAttention + PagedAttention:
  不需要将KV拷贝到连续buffer
  直接在页边界进行tiling
  逻辑上提供连续KV的API,物理上使用page table寻址

这种协同设计让推理引擎能够:

  1. 支持任意数量的并发请求(仅受物理显存限制)
  2. 每个请求只占用实际长度的 KV Cache
  3. 零内存碎片(页级别的内碎片 ≤ 16 tokens/page)
  4. 实验数据:vLLM 的 PagedAttention 将吞吐量提升了 2x~4x,内存浪费从 80% 降到 <4%。

    三、SGLang 的 RadixAttention:前缀缓存的工程创新

    如果说 PagedAttention 解决的是 "不同请求间的内存复用",那么 SGLang 的 RadixAttention 进一步解决了 "不同请求间相同前缀的 KV Cache 复用"。

    3.1 前缀复用的价值

    在多轮对话场景和批量请求场景中,不同请求往往共享相同的前缀:

    
    对话场景:
      System prompt: "You are a helpful assistant that..."
      User Q1: "What is Python?"
      Assistant A1: "Python is a programming language..."
      User Q2: "Is it fast?"
    
    场景2:
      System prompt: "You are a helpful assistant that..."  ← 相同前缀可以复用
      User Q3: "What is Java?"
    

    如果不复用前缀,每个请求都需要重新计算并存储 System prompt 的 KV Cache,这在长 system prompt 场景下浪费严重(数秒的计算时间和数十MB的显存)。

    3.2 Radix Tree 作为索引结构

    SGLang 使用 Radix Tree(基数树)作为前缀索引:

    
    class RadixNode:
        """基数树节点"""
        __slots__ = ['children', 'page_list', 'ref_count', 'token_len']
        
        def __init__(self):
            self.children: dict[int, 'RadixNode'] = {}  # token hash -> child
            self.page_list: list[int] = []  # 该前缀对应的KV Cache页列表
            self.ref_count: int = 0  # 引用计数
            self.token_len: int = 0  # 该节点表示的token数
    

    操作流程:

    1. 插入:将 tokenized 前缀序列插入 radix tree,每个节点记录对应的 KV Cache 页
    2. 查找:匹配请求前缀,复用已有的 KV Cache 页(ref_count++)
    3. 淘汰:当显存不足时,LRU 策略淘汰叶子节点(ref_count=0 的节点)
    4. 
      RadixAttention 示例:
      
                  root
                 /    \
               [sys]   [You] → [are] → [a] → [helpful]
                 |                                     |
          pages: [1,2,3]                        pages: [4,5,6,7]
          ref_count: 2                          ref_count: 1
          
      请求到达时:
        若前缀匹配到 Node("You are a helpful"):
          → 直接复用 pages [1,2,3,4,5,6]
          → 只需要计算新 token "assistant." 的 KV
          → ref_count += 1
      

      3.3 生产环境性能收益

      SGLang 论文数据显示,RadixAttention 在以下场景收益尤为突出:

      • Few-shot prompting:多个请求共享 few-shot examples → 1.5x~3x 加速
      • 多轮对话:每轮对话复用历史 KV Cache → 首个 token 延迟降低 80%
      • Batch 翻译:同源文本共享语言检测前缀 → 2x 吞吐提升

      四、Continuous Batching 调度策略详解

      Continuous Batching 的核心在于调度策略——决定每个 iteration 哪些请求参与计算、新请求何时插入。

      4.1 调度算法的时间轴

      
      class ContinuousBatcher:
          def __init__(self, max_batch_size: int, max_num_tokens: int):
              self.running: deque[Request] = deque()  # 运行中的请求
              self.waiting: deque[Request] = deque()  # 等待队列
              self.max_batch_size = max_batch_size
              self.max_num_tokens = max_num_tokens
          
          def step(self) -> list[Request]:
              """每轮 iteration 的调度决策"""
              
              # 阶段1: 处理已完成的请求
              for req in list(self.running):
                  if req.is_finished_or_timeout():
                      self.running.remove(req)
                      kv_manager.free(req.id)
              
              # 阶段2: 从等待队列补充新请求
              while self.waiting and len(self.running) < self.max_batch_size:
                  # 预估当前 running + 新请求的 token 预算
                  projected_tokens = self._current_tokens() + self.waiting[0].prompt_len
                  if projected_tokens <= self.max_num_tokens:
                      new_req = self.waiting.popleft()
                      self.running.append(new_req)
                  else:
                      break  # token 预算已满
              
              # 阶段3: 执行推理
              return list(self.running)
      

      4.2 三种主流调度策略

      策略一:FCFS(First-Come-First-Served)

      
      优点:公平性高,实现简单
      缺点:尾部长请求阻塞后续短请求
      适用场景:离线批处理、高利用率优先
      

      策略二:Watermark-Based Scheduling

      
      维护两个水位线:
        - 高水位 (high watermark): max_batch_size * avg_decode_len
        - 低水位 (low watermark): max_batch_size * min_decode_len
      
      当 running tokens < low watermark 时积极接纳新请求
      当 running tokens > high watermark 时保守接纳
      

      策略三:Chunked Prefill(vLLM v0.4+)

      这是目前最先进的策略,核心思想是将长 prefill 切分:

      
      class ChunkedPrefillScheduler:
          """允许 prefill 和 decode 混合执行"""
          
          def schedule(self):
              # 不再等待整个 prefill 完成
              # 将长 prefill 切分为 chunks,与 decode 交错执行
              
              chunk_budget = self.max_prefill_chunk_size  # 如 512 tokens
              
              for req in self.waiting:
                  if req.prefill_remaining > 0:
                      chunk_size = min(req.prefill_remaining, chunk_budget)
                      req.prefill_remaining -= chunk_size
                      # 将 chunk 与 running 中的请求一起调度
                      # 避免长 prefill 阻塞所有 decode 请求
      
      
      时间线对比:
      
      传统(prefill优先):
      | LONG PREFILL (1024 tokens) | req1 decode | req1 decode | ...
      
      Chunked Prefill:
      | P1 (512) | req1 decode | req2 decode | P2 (512) | req1 decode | ...
                  ↑ 不让 prefill 独占整个计算窗口
      

      收益:降低 decode 请求的 TTFT(Time To First Token)方差,提升服务质量。

      4.3 抢占与恢复(Preemption)

      当 GPU 显存不足以支撑所有 running 请求时,需要抢占(preempt)部分请求:

      
      class PreemptionPolicy:
          def need_preempt(self, running: list[Request], waiting: list[Request]) -> bool:
              """判断是否需要抢占"""
              free_pages = kv_manager.free_pages_count()
              # 确保最高优先级的新请求能进入
              return free_pages < waiting[0].pages_needed if waiting else False
          
          def preempt(self, running: list[Request]) -> Request:
              """选择被抢占的请求 (LRU策略)"""
              # 选择最后完成时间最晚的请求
              return min(running, key=lambda r: r.last_decode_time)
              # 释放其 KV Cache,将其重新加入 waiting 队列头部
      

      五、vLLM vs SGLang:生产级引擎架构对比

      5.1 架构设计哲学

      5.2 性能数据参考

      维度 vLLM SGLang
      核心创新 PagedAttention (OS级内存管理) RadixAttention + 结构化生成
      编程接口 OpenAI 兼容 API 原生 Python DSL + OpenAI API
      KV Cache 管理 纯 PagedAttention PagedAttention + Radix Tree 复用
      调度策略 FCFS + Chunked Prefill FCFS + Echo Scheduling
      特殊优势 成熟度高、生态丰富 结构化生成、复杂 LLM 编程
      劣势 前缀复用能力有限 生态系统相对年轻

      在 A100-80GB, Llama-2-70B 场景下(基于公开 benchmark):

      场景 vLLM 吞吐 (tokens/s) SGLang 吞吐 (tokens/s) 提升
      ShareGPT 数据集 ~2800 ~3200 +14%
      100% 前缀缓存命中 ~2800 ~6400 +129%
      50% 前缀缓存命中 ~2800 ~4500 +61%

      SGLang 在前缀复用场景下优势明显,这验证了 RadixAttention 的工程价值。

      5.3 生产部署考量

      选择推理引擎时,还需考虑以下维度:

      1. 模型覆盖度:vLLM 支持的模型列表更广,SGLang 对主流模型也已全面支持
      2. DeepSeek-V3/R1 的优化:两个引擎都已适配 MLA(Multi-head Latent Attention)
      3. 分布式推理:都支持 TP(Tensor Parallelism),SGLang 提供更灵活的 DP+TP 组合
      4. 可观测性:vLLM 集成 Prometheus metrics 更成熟
      5. 与框架集成:SGLang 的 DSL 更适合复杂多轮 Agent 场景
      6. 六、实战:从零实现 Mini Continuous Batching 引擎

        理解 Continuous Batching 最好的方式是动手实现。以下是一个简化版本的核心逻辑:

        
        import torch
        from dataclasses import dataclass, field
        from collections import deque
        from typing import Optional
        
        @dataclass
        class InferenceRequest:
            request_id: int
            input_ids: list[int]
            output_ids: list[int] = field(default_factory=list)
            kv_pages: list[int] = field(default_factory=list)
            status: str = "pending"  # pending | running | finished
            
            @property
            def prompt_len(self):
                return len(self.input_ids)
            
            @property
            def total_tokens(self):
                return self.prompt_len + len(self.output_ids)
        
        class MiniKVManager:
            """简化版 Paged KV Cache 管理器"""
            
            def __init__(self, num_layers: int, num_kv_heads: int, 
                         head_dim: int, page_size: int, num_pages: int, device: str):
                self.page_size = page_size
                self.num_pages = num_pages
                self.device = device
                
                # 物理存储:每层的 KV Cache 页池
                # shape: [num_pages, page_size, num_kv_heads, head_dim]
                self.kv_cache = {
                    layer: torch.zeros(num_pages, page_size, num_kv_heads, head_dim, 
                                      device=device, dtype=torch.float16)
                    for layer in range(num_layers)
                }
                
                self.free_pages = deque(range(num_pages))
                self.page_tables: dict[int, list[int]] = {}
            
            def alloc(self, request_id: int, seq_len: int) -> bool:
                pages_needed = (seq_len + self.page_size - 1) // self.page_size
                if len(self.free_pages) < pages_needed:
                    return False  # OOM
                
                pages = [self.free_pages.popleft() for _ in range(pages_needed)]
                self.page_tables[request_id] = pages
                return True
            
            def free(self, request_id: int):
                if request_id in self.page_tables:
                    self.free_pages.extend(self.page_tables[request_id])
                    del self.page_tables[request_id]
            
            def get_kv(self, request_id: int, layer: int) -> torch.Tensor:
                """获取某个请求的所有KV Cache(逻辑上是连续的)"""
                pages = self.page_tables[request_id]
                return self.kv_cache[layer][pages]  # shape: [num_pages, page_size, heads, dim]
        
        class MiniContinuousBatcher:
            """简化版 Continuous Batching 调度器"""
            
            def __init__(self, max_batch_size: int, max_tokens_per_batch: int,
                         kv_manager: MiniKVManager):
                self.max_batch_size = max_batch_size
                self.max_tokens_per_batch = max_tokens_per_batch
                self.kv_manager = kv_manager
                
                self.running: deque[InferenceRequest] = deque()
                self.waiting: deque[InferenceRequest] = deque()
            
            def add_request(self, request: InferenceRequest):
                self.waiting.append(request)
            
            def step(self) -> Optional[list[InferenceRequest]]:
                """执行一轮推理,返回当前 batch"""
                
                # 1. 清理完成的请求
                for req in list(self.running):
                    if self._is_finished(req):
                        self.running.remove(req)
                        self.kv_manager.free(req.request_id)
                
                # 2. 从等待队列填充
                while self.waiting:
                    if len(self.running) >= self.max_batch_size:
                        break
                    
                    candidate = self.waiting[0]
                    # 预估 KV Cache 是否足够
                    pages_needed = candidate.total_tokens // self.kv_manager.page_size + 1
                    if len(self.kv_manager.free_pages) >= pages_needed:
                        self.waiting.popleft()
                        if self.kv_manager.alloc(candidate.request_id, candidate.total_tokens):
                            candidate.status = "running"
                            self.running.append(candidate)
                        else:
                            break
                    else:
                        break  # 显存不足
                
                if not self.running:
                    return None
                
                return list(self.running)
            
            def finish_token(self, request_id: int, new_token: int):
                """为指定请求追加生成的 token"""
                for req in self.running:
                    if req.request_id == request_id:
                        req.output_ids.append(new_token)
                        # 当当前页写满时,追加一个新页
                        if req.total_tokens % self.kv_manager.page_size == 0:
                            if self.kv_manager.free_pages:
                                new_page = self.kv_manager.free_pages.popleft()
                                self.kv_manager.page_tables[request_id].append(new_page)
            
            def _is_finished(self, req: InferenceRequest) -> bool:
                return (req.total_tokens >= req.prompt_len + 2048  # max_new_tokens
                        or 0 in req.output_ids)  # EOS token
            
            def exec_forward(self, model, tokenizer):
                """执行一轮完整的推理步骤"""
                batch = self.step()
                if batch is None:
                    return
                
                # 收集所有请求的当前位置 token
                input_tokens = []
                kv_page_locs = []
                positions = []
                
                for req in batch:
                    # 如果是 prefill(首次进入 batch),使用所有 prompt tokens
                    if len(req.output_ids) == 0:
                        input_tokens.append(req.input_ids)
                        positions.append(list(range(req.prompt_len)))
                    else:
                        # decode 模式,只使用最新一个 token
                        new_token = req.output_ids[-1]
                        input_tokens.append([new_token])
                        positions.append([req.total_tokens - 1])
                
                # 实际推理(此处简化为注释)
                # output = model.forward(
                #     input_ids=token_batch,
                #     kv_page_tables=self.kv_manager.page_tables,
                #     positions=positions
                # )
                # 将 output.logits 采样得到 next token
                # 调用 finish_token() 追加结果
        

        以上代码展示了 Continuous Batching 的三个核心模块:分页 KV Cache 管理、迭代级调度、动态显存分配/释放。实际生产引擎还需处理 CUDA Graph 捕获、TP 并行、CUDA Stream 流水线等高级特性。

        七、前沿趋势与展望

        7.1 Disaggregated Prefill(预填充-解码分离)

        TPU 架构中的 "预填充-解码分离" 部署模式正在 GPU 场景中兴起:

        • Prefill 实例:高算力密度,优化矩阵乘法
        • Decode 实例:高内存容量,优化带宽瓶颈
        • 通过 RDMA 直接传递 KV Cache,避免重复计算

        DistServe 论文显示,在特定负载下分离部署可提升吞吐 2.5x。Mooncake 等云原生推理平台已在生产环境中实现了这种分离架构。

        7.2 Speculative Decoding 与 Continuous Batching 结合

        使用小模型(Draft Model)预测多个 token,由大模型一次验证:

        • 验证阶段天然适配 Continuous Batching:一次 forward 处理多个 speculative tokens
        • PagedAttention 在 speculative tokens 被拒绝时直接丢弃页,无需回滚开销

        7.3 Mamba / SSM 架构对 KV Cache 的挑战

        状态空间模型(如 Mamba)不使用 KV Cache,取而代之的是固定大小的隐藏状态:

        • 对 Continuous Batching 调度的影响:不再需要分页管理
        • 对显存规划的影响:并发数量由固定状态大小决定,而非动态分配
        • 工程挑战:选择性 SSM 的状态更新与 Continuous Batching 的迭代调度存在概念冲突

        7.4 vAttention:用 PagedAttention 统一 Attention

        NVIDIA 的 vAttention 项目直接将 PagedAttention 思想融入 CUDA kernel,绕过 Python/C++ 层的手动分页管理,实现更高效的显存操作。

        八、总结

        Continuous Batching 的核心价值在于:将 GPU 计算资源从 "序列级独占" 解放为 "迭代级共享"。这一理念配合 PagedAttention 的显存管理、RadixAttention 的前缀缓存,构成了现代 LLM 推理引擎的三大技术支柱。

        对于工程落地:

        • 选择引擎:通用场景选 vLLM,多轮对话/Agent 场景选 SGLang
        • 调优优先级:max_num_seqs(并发数)> chunked_prefill_size > 调度策略
        • 监控指标:throughput (tokens/s)、ITL (Inter-Token Latency)、TTFT、缓存命中率
        • 硬件感知:A100 靠 PagedAttention 释放显存;H100 靠 Flash Decoding 加速带宽

        理解这些原理不仅能帮助选择和使用推理引擎,更重要的是在模型定制(如微调、蒸馏、量化)时做出正确的架构决策。在大模型从实验室走向生产的征途中,推理效率是与模型质量同等重要的竞争力。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部