大模型推理框架SGLang的源码分析

基于 SGLang v0.5.x 系列源码分析,GitHub: sgl-project/sglang


一、整体架构与源码结构

1.1 进程级架构

SGLang 的 SRT (SGLang Runtime) Server 采用多进程架构,核心由三个组件构成:

arduino 复制代码
┌─────────────────────────────────────────────────────────────────┐
│                     SRT Server 进程架构                           │
│                                                                  │
│  ┌──────────────────┐                                            │
│  │  HTTP Server     │  FastAPI, 接收 OpenAI 兼容 API 请求          │
│  │  (entrypoints/)  │                                            │
│  └────────┬─────────┘                                            │
│           │                                                      │
│  ┌────────▼─────────┐                                            │
│  │  TokenizerManager│ 进程1: tokenize → 发送 token_id 到 Scheduler│
│  │  (python/sglang/ │  ZMQ 通信                                  │
│  │   runtime/)      │                                           │
│  └────────┬─────────┘                                           │
│           │ ZMQ                                                 │
│  ┌────────▼──────────────────────────────────────┐              │
│  │  Scheduler (进程2: 核心调度引擎)                 │             │
│  │  ├── waiting_queue    (等待调度的请求)           │             │
│  │  ├── grammar_queue    (等待语法解析的请求)       │              │
│  │  ├── running_batch    (正在执行的 batch)        │              │
│  │  ├── RadixTree        (KV Cache 基数树)        │              │
│  │  └── TokenToKVPool    (KV Cache 物理内存池)     │              │
│  └────────┬──────────────────────────────────────┘              │
│           │                                                     │
│  ┌────────▼──────────────────────────────────────┐              │
│  │  TPModelWorker (进程3+: TP 并行的工作进程)       │              │
│  │  ├── ModelRunner                              │              │
│  │  │   ├── Model (transformers 模型)             │              │
│  │  │   └── AttentionBackend (FlashInfer/FlashAttn)             │
│  │  └── ForwardMode: EXTEND / DECODE             │              │
│  └───────────────────────────────────────────────┘              │
│           │                                                     │
│  ┌────────▼─────────┐                                           │
│  │DetokenizerManager│ 进程N: token_id → 文本, 流式输出            │
│  └──────────────────┘                                           │
└─────────────────────────────────────────────────────────────────┘

1.2 源码目录结构

bash 复制代码
sglang/
├── python/sglang/
│   ├── runtime/                    # 核心运行时
│   │   ├── scheduler.py            # ★ 调度器核心(最复杂)
│   │   ├── radix_attention.py      # ★ 基数树 KV Cache 管理
│   │   ├── model_runner.py         # 模型执行器
│   │   └── server_args.py          # 启动参数
│   ├── srt/                        # SRT 引擎
│   │   ├── layers/attention/       # ★ 注意力后端
│   │   │   ├── flashinfer_backend.py   # FlashInfer 实现
│   │   │   ├── flashattention_backend.py
│   │   │   └── triton_backend.py
│   │   ├── layers/radix_attention.py # ★ RadixAttention 层
│   │   ├── managers/               # 各管理器
│   │   │   ├── schedule_batch.py   # ★ Batch 数据结构
│   │   │   ├── tokenizer_manager.py
│   │   │   └── detokenizer_manager.py
│   │   ├── constrained/            # ★ 结构化约束解码
│   │   │   ├── fsm.py              # 有限状态机
│   │   │   └── grammar.py
│   │   └── server.py               # 入口
│   ├── backend/                    # 前端语言运行时
│   └── entrypoints/                # HTTP 入口
└── test/

二、RadixAttention:基数树 KV Cache 管理(核心创新)

2.1 数据结构设计

源码位置:python/sglang/srt/layers/radix_attention.py + python/sglang/runtime/radix_attention.py

python 复制代码
class RadixTreeNode:
    """基数树节点"""
    def __init__(self):
        self.children: Dict[int, RadixTreeNode] = {}  
        # key = token_id, value = 子节点
        self.parent: RadixTreeNode = None
        self.lock_ref: int = 0          # 引用计数(被多少请求使用)
        self.last_access_time: float = 0 # LRU 时间戳
        self.value: int = 0             # 关联的 KV Cache 物理块索引
        self.key: Tuple[int] = ()       # 节点对应的 token 序列

class RadixTree:
    """基数树主体"""
    def __init__(self):
        self.root = RadixTreeNode()
        self.token_to_kv_pool = TokenToKVPool()  # KV Cache 物理内存池
        self.evictable_size_ = 0  # 可淘汰的 KV Cache 大小

关键设计细节

  • 节点粒度 :不是每个 token 一个节点,而是每个唯一 token 子序列一个节点,这是基数树(压缩前缀树)的核心特征
  • 引用计数lock_ref 记录当前有多少活跃请求正在使用该节点的 KV Cache,引用计数 > 0 的节点不可淘汰
  • 物理内存池TokenToKVPool 是一个预分配的 GPU 显存池,所有 KV Cache 物理存储统一管理

2.2 前缀匹配与插入流程

ini 复制代码
请求到达,token 序列: [system_prompt, user_query, ...]

Step 1: 前缀匹配(最长前缀匹配)
┌─────────────────────────────────────────────────────┐
│  输入: [t1, t2, t3, t4, t5, t6, t7, t8]             │
│                                                     │
│  基数树中已有:                                       │
│    root → [t1,t2,t3] → [t4,t5] → [t6]              │
│                                                     │
│  匹配过程:                                          │
│    [t1,t2,t3] ✓ 完全匹配 → 继续到子节点               │
│    [t4,t5]   ✓ 完全匹配 → 继续到子节点               │
│    [t6]      ✓ 匹配,但输入还有 [t7,t8]              │
│    [t7,t8]   ✗ 无匹配 → 停止                        │
│                                                     │
│  结果: prefix_len = 6, 不需要计算前6个token的KV      │
│        新增部分: [t7, t8] 需要计算 KV                │
└─────────────────────────────────────────────────────┘

源码核心逻辑(简化):

python 复制代码
def match_prefix(self, token_ids: List[int]) -> Tuple[List[int], List[RadixTreeNode]]:
    """在基数树中做最长前缀匹配,返回匹配的 token 和对应的节点链"""
    node = self.root
    matched_tokens = []
    matched_nodes = []
    
    remaining = list(token_ids)
    while remaining:
        # 尝试在当前节点的子节点中匹配
        matched = False
        for child_key, child_node in node.children.items():
            if self._prefix_match(remaining, child_key):
                # 完全匹配子节点
                matched_tokens.extend(child_key)
                matched_nodes.append(child_node)
                node = child_node
                remaining = remaining[len(child_key):]
                child_node.lock_ref += 1  # 增加引用计数
                child_node.last_access_time = time.time()
                matched = True
                break
            elif self._partial_prefix_match(remaining, child_key):
                # 部分匹配 → 需要分裂节点
                # ...
                pass
        
        if not matched:
            break
    
    return matched_tokens, matched_nodes, remaining  # remaining = 未匹配部分

2.3 节点分裂机制

当新请求与已有节点部分匹配 时,需要分裂节点

ini 复制代码
分裂前:  root → [t1, t2, t3, t4] → [t5, t6]
                         ↑
新请求: [t1, t2, t3, t7, t8]
         完全匹配前3个,但第4个不同

分裂后:  root → [t1, t2, t3] → [t4] → [t5, t6]    (原有路径保留)
                              └→ [t7, t8]            (新路径)

关键实现细节

  • 分裂时,KV Cache 的物理块也需要相应拆分
  • [t1,t2,t3,t4] 的 KV Cache 被拆分为 [t1,t2,t3][t4] 两部分
  • [t4] 的 KV Cache 物理块不需要重新计算,只是索引关系变更
  • 分裂操作是O(1) 的,因为只是修改指针和引用计数

2.4 LRU 淘汰策略

python 复制代码
def evict(self, need_size: int) -> int:
    """淘汰最少使用的 KV Cache,释放至少 need_size 的空间"""
    # 使用最小堆按 last_access_time 排序
    # 只淘汰 lock_ref == 0 的节点(无活跃请求引用)
    
    evicted = 0
    candidates = self._collect_evictable_nodes()  # 按LRU排序
    
    for node in candidates:
        if evicted >= need_size:
            break
        if node.lock_ref > 0:
            continue  # 被活跃请求引用,不可淘汰
        
        # 从父节点断开
        parent = node.parent
        del parent.children[node.key]
        
        # 释放 KV Cache 物理内存
        self.token_to_kv_pool.free(node.value)
        evicted += len(node.key)
        
        # 如果父节点只剩一个子节点,可以合并(压缩)
        self._try_merge(parent)
    
    return evicted

淘汰后合并 :当淘汰导致父节点只剩一个子节点时,基数树会自动合并节点,保持压缩特性:

ini 复制代码
淘汰前: [t1,t2] → [t3] → [t5,t6]   (t4 已被淘汰)
                                    (t3 只有一个子节点)
合并后: [t1,t2] → [t3,t5,t6]        (自动压缩)

2.5 与 vLLM PagedAttention 的本质区别

维度 vLLM PagedAttention SGLang RadixAttention
数据结构 块表 (Block Table) 基数树 (Radix Tree)
管理粒度 单请求内分页 跨请求前缀复用
复用方式 Copy-on-Write(有限) 最长前缀匹配(原生)
淘汰策略 请求结束释放 LRU 按需淘汰
缓存生命周期 请求级 跨请求持久化
树结构变更 分裂/合并/压缩

三、调度器:零开销 CPU 调度 + Overlap 机制

3.1 调度器核心循环

源码位置:python/sglang/runtime/scheduler.py

SGLang 的 Scheduler 是整个推理引擎最复杂的类,核心入口是一个事件循环:

python 复制代码
class Scheduler:
    def event_loop_normal(self):
        """Normal 模式的事件循环"""
        while True:
            # 1. 从 ZMQ 队列接收新请求
            recv_requests()
            
            # 2. 挑选待执行的请求,组成 batch
            batch = self.prepare_batch()
            
            # 3. 发起 GPU 推理 (run_batch)
            result = self.run_batch(batch)
            
            # 4. 后处理:采样结果、更新 KV Cache 等
            self.process_batch_result(batch, result)

3.2 Overlap Scheduler:CPU-GPU 流水线重叠

源码关键:event_loop_overlap vs event_loop_normal

Normal 模式(无重叠):

ini 复制代码
时间轴: ──────────────────────────────────────────►
CPU:  [组batch] [等待GPU] [后处理] [组batch] [等待GPU] [后处理]
GPU:            [推理]              [推理]
                          ↑ GPU空闲

Overlap 模式(流水线重叠):

ini 复制代码
时间轴: ──────────────────────────────────────────►
CPU:  [组batch_N+1] [组batch_N+2] [后处理N] [后处理N+1]
GPU:  [推理batch_N] [推理batch_N+1] [推理batch_N+2]
                    ↑ CPU调度与GPU计算并行

实现细节

python 复制代码
def event_loop_overlap(self):
    """Overlap 模式:CPU调度与GPU计算并行"""
    # 使用双缓冲策略
    batch_in_flight = None    # GPU 正在执行的 batch
    batch_prepared = None     # CPU 预先准备好的 batch
    
    while True:
        # Phase 1: GPU 推理上一个 batch (异步)
        if batch_in_flight is not None:
            gpu_future = self.model_worker.run_batch_async(batch_in_flight)
        
        # Phase 2: CPU 并行准备下一个 batch
        #   - 从 ZMQ 接收新请求
        #   - 调度决策(哪些请求加入batch)
        #   - 准备输入张量
        #   - RadixTree 前缀匹配
        #   - 内存分配
        recv_requests()
        batch_prepared = self.prepare_next_batch()
        
        # Phase 3: 等待 GPU 完成,获取结果
        if batch_in_flight is not None:
            result = gpu_future.wait()
            self.process_batch_result(batch_in_flight, result)
        
        # Phase 4: 交换
        batch_in_flight = batch_prepared
        batch_prepared = None

为什么是"零开销"

  • 所有 CPU 调度逻辑(组 batch、前缀匹配、内存分配)都在 GPU 计算期间并行完成
  • GPU 的空闲时间趋近于零,从"等待 CPU 调度"变为"GPU 算完立即执行下一个 batch"
  • 调度器纯 Python 实现,不需要 C++ 扩展,因为调度开销被完全隐藏

3.3 Prefill 优先调度策略

SGLang 的调度以 Prefill 为主导

python 复制代码
def prepare_batch(self):
    """准备一个 batch,Prefill 优先"""
    
    # 1. 如果有待执行的 prefill,优先执行
    if self.has_waiting_prefill():
        # 可能会中断正在运行的 decode batch
        # 这是为了控制 TTFT(首字延迟)
        return self.prepare_prefill_batch()
    
    # 2. 如果没有 prefill,执行 decode
    if self.running_decode_requests:
        return self.prepare_decode_batch()
    
    # 3. 空闲时做自检和状态重置
    self.on_idle()

Prefill 优先的权衡

  • 优点:新请求的 TTFT 不会因等待 decode 而过高
  • 缺点:可能中断正在运行的 decode batch,导致 decode 请求的 TPOT(每token延迟)波动
  • SGLang 的策略 :通过 --schedule-policy 参数可配置,默认是 prefill 优先

3.4 请求调度策略:LSPF(Longest-Shared-Prefix-First)

python 复制代码
def get_priority(self, req: Req) -> float:
    """计算请求的调度优先级"""
    # LSPF: 最长共享前缀优先
    # 前缀越长 → 与基数树中已有 KV Cache 匹配越多
    # → 需要重新计算的部分越少 → 调度效率越高
    prefix_len = self.radix_tree.match_prefix_length(req.token_ids)
    return prefix_len  # 优先调度前缀匹配最长的请求

LSPF 的意义

  • 前缀匹配长的请求几乎不需要重新计算 KV Cache,执行速度快
  • 优先执行这些请求可以快速释放资源,提高整体吞吐
  • 类似于"最短作业优先"的调度思想,但更关注 KV Cache 复用效率

四、注意力后端:FlashInfer + Extend/Decode 双模式

4.1 ForwardMode:EXTEND vs DECODE

源码位置:python/sglang/srt/layers/attention/

SGLang 将推理的注意力计算分为两种模式:

ini 复制代码
class ForwardMode(Enum):
    EXTEND = 1    # 处理新 token(类似 Prefill,但可复用已有 KV Cache)
    DECODE = 2    # 逐 token 生成(每次只生成1个新 token)

EXTEND vs 传统 Prefill 的区别

markdown 复制代码
传统 Prefill:  整个 prompt 的所有 token 都需要计算 KV Cache
                → 完全并行计算

EXTEND:        只有前缀匹配之后的新 token 需要计算 KV Cache
                → 已有 KV Cache 的 token 直接复用
                → 只对新增部分做并行计算
                → 计算 Q 对已复用 K,V 的注意力

4.2 FlashInfer 后端实现

SGLang 主要使用 FlashInfer 作为注意力后端,而非原生的 FlashAttention:

ini 复制代码
class FlashInferBackend:
    def __init__(self):
        self.prefill_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper()
        self.decode_wrapper = flashinfer.BatchDecodeWithPagedKVCacheWrapper()
    
    def forward_extend(self, batch: ScheduleBatch):
        """EXTEND 模式:处理新 token 的注意力计算"""
        # 1. 构建输入张量
        qo_indptr = batch.extend_indptr          # query/output 的索引指针
        kv_indptr = batch.kv_indptr              # KV Cache 的索引指针
        kv_indices = batch.kv_indices            # KV Cache 物理块索引
        kv_last_page_len = batch.kv_last_page_len
        
        # 2. 调用 FlashInfer 的 prefill kernel
        #    支持分页 KV Cache(非连续内存)
        output = self.prefill_wrapper.run(
            q=batch.query,
            qo_indptr=qo_indptr,
            kv_data=batch.kv_data,               # 从 TokenToKVPool 取
            kv_indptr=kv_indptr,
            kv_indices=kv_indices,
            kv_last_page_len=kv_last_page_len,
            sm_scale=1.0 / math.sqrt(head_dim),
        )
        return output
    
    def forward_decode(self, batch: ScheduleBatch):
        """DECODE 模式:每个请求只生成1个新 token"""
        # Decode 时 Q 只有 1 个 token,K,V 是全部历史
        # 使用 FlashInfer 的 decode kernel,高度优化访存
        output = self.decode_wrapper.run(
            q=batch.query,                        # [batch_size, 1, head_dim]
            kv_data=batch.kv_data,
            kv_indptr=batch.kv_indptr,
            kv_indices=batch.kv_indices,
            kv_last_page_len=batch.kv_last_page_len,
        )
        return output

4.3 KV Cache 物理内存管理:TokenToKVPool

python 复制代码
class TokenToKVPool:
    """KV Cache 物理内存池,统一管理所有请求的 KV Cache 存储"""
    
    def __init__(self, max_total_tokens: int, head_dim: int, num_layers: int, ...):
        # 预分配一块大的 GPU 显存
        self.k_buffer = torch.empty(...)  # [max_total_tokens, num_kv_heads, head_dim]
        self.v_buffer = torch.empty(...)  # [max_total_tokens, num_kv_heads, head_dim]
        
        # 空闲 token 槽位的管理
        self.free_slots = list(range(max_total_tokens))  # 可用槽位列表
        self.is_free = torch.ones(max_total_tokens, dtype=torch.bool)
    
    def alloc(self, num_tokens: int) -> List[int]:
        """分配 num_tokens 个槽位,返回索引列表"""
        if len(self.free_slots) < num_tokens:
            return None  # 需要淘汰
        indices = self.free_slots[:num_tokens]
        self.free_slots = self.free_slots[num_tokens:]
        self.is_free[indices] = False
        return indices
    
    def free(self, indices: List[int]):
        """释放槽位"""
        self.free_slots.extend(indices)
        self.is_free[indices] = True
    
    def get_kv_data(self, indices: List[int]):
        """根据索引获取 KV Cache 数据"""
        return self.k_buffer[indices], self.v_buffer[indices]

关键设计

  • 统一物理池:所有请求的 KV Cache 共享同一个 GPU 显存池,消除碎片化
  • 索引式访问 :通过 indices 列表访问,不需要物理连续
  • 与 RadixTree 的关系:RadixTree 管理逻辑到物理的映射,TokenToKVPool 管理物理内存的分配/释放

五、结构化约束解码:压缩有限状态机

5.1 问题背景

结构化输出(如 JSON、Regex)需要约束解码 :每个生成的 token 必须满足预定义的语法。传统方案使用拒绝采样(生成 token 后检查是否合法),效率极低。

5.2 压缩有限状态机(Compressed FSM)

源码位置:python/sglang/srt/constrained/

python 复制代码
class CompressedFSM:
    """压缩有限状态机,用于约束解码"""
    
    def __init__(self, grammar: Grammar):
        # 从 JSON Schema / Regex 构建原始 DFA
        self.dfa = self._build_dfa(grammar)
        
        # 压缩:合并具有相同转移函数的状态
        self.compressed_dfa = self._compress(self.dfa)
        
        # 为每个状态预计算允许的 token 集合
        self.state_allowed_tokens = self._precompute_allowed_tokens()

压缩原理

arduino 复制代码
原始 DFA:
  State 0 ──'{'──→ State 1 ──'"'──→ State 2 ──'k'──→ State 3
  State 0 ──'['──→ State 4 ──'"'──→ State 5
  State 1 和 State 4 的后续转移函数相同(都需要 '"')
  
压缩后:
  State 0 ──'{'──→ State 1+4(merged) ──'"'──→ State 2+5(merged)
  State 0 ──'['──→ State 1+4(merged)

关键优化 :压缩后状态数大幅减少,状态转移查找从 O(V) 降到 O(1)

5.3 约束解码的执行流程

python 复制代码
def constrained_decode(self, logits, fsm_state, req: Req):
    """在采样阶段应用约束"""
    
    # 1. 获取当前 FSM 状态允许的 token 集合
    allowed_tokens = self.state_allowed_tokens[fsm_state]
    
    # 2. 对 logits 做 mask:不允许的 token 设为 -inf
    logits_mask = torch.full_like(logits, float('-inf'))
    logits_mask[allowed_tokens] = 0
    constrained_logits = logits + logits_mask
    
    # 3. 正常采样(top-k, top-p, temperature)
    token_id = sample(constrained_logits, req.sampling_params)
    
    # 4. 更新 FSM 状态
    new_fsm_state = self.dfa.transition(fsm_state, token_id)
    
    return token_id, new_fsm_state

与 Outlines 的对比

维度 Outlines (vLLM 集成) SGLang Compressed FSM
FSM 构建时间 较慢 快(压缩后状态少)
状态数 原始 DFA 压缩后减少 50-80%
每步开销 O(V) logit mask O(
多请求支持 每请求独立 FSM 每请求维护 FSM 状态
与 RadixTree 协同 有(共享前缀的 FSM 状态可复用)

六、PD 分离架构

6.1 架构设计

源码位置:PD 分离模式是 SGLang v0.4+ 的核心特性

scss 复制代码
┌─────────────────────────────────────────────────────────────┐
│                 PD 分离架构                                    │
│                                                             │
│  ┌──────────────┐         ┌──────────────┐                  │
│  │  Prefiller    │  KV     │  Decoder     │                  │
│  │  (计算密集型)  │  Cache  │  (访存密集型) │                  │
│  │              │  传输    │              │                  │
│  │  高算力 GPU   │────────→│  高显存 GPU   │                  │
│  │  (H100 SXM)  │  RDMA   │  (H100 80GB) │                  │
│  └──────────────┘         └──────────────┘                  │
│        ↑                        ↑                           │
│        │                        │                           │
│  ┌─────┴──────────────────────┴─────┐                       │
│  │         Router (路由器)            │                       │
│  │  - 请求路由:新请求 → Prefiller    │                       │
│  │  - KV Cache 传输:Prefiller → Decoder                   │
│  │  - 负载均衡:Prefiller/Decoder 独立扩缩容                │
│  └─────────────────────────────────┘                       │
└─────────────────────────────────────────────────────────────┘

6.2 KV Cache 传输机制

python 复制代码
# Prefiller 端
class PrefillWorker:
    def prefill_request(self, req: Req):
        """执行 prefill,计算 KV Cache"""
        # 1. 前缀匹配(RadixTree)
        matched_len, matched_nodes = self.radix_tree.match_prefix(req.token_ids)
        
        # 2. 只计算新增部分的 KV Cache
        new_tokens = req.token_ids[matched_len:]
        new_kv = self.forward_extend(new_tokens, matched_nodes)
        
        # 3. 通过 RDMA/NCCL 将 KV Cache 传输到 Decoder
        self.transfer_kv_cache(req.req_id, new_kv, matched_kv)
        
        # 4. 通知 Decoder 开始 decode
        self.notify_decoder(req.req_id)

# Decoder 端
class DecodeWorker:
    def on_kv_cache_received(self, req_id: str, kv_cache):
        """接收 KV Cache,开始 decode"""
        # 1. 将 KV Cache 写入本地 TokenToKVPool
        indices = self.kv_pool.alloc(kv_cache.num_tokens)
        self.kv_pool.write(indices, kv_cache)
        
        # 2. 将请求加入 decode 队列
        self.decode_queue.add(req_id, indices)

6.3 三层分离:EPD(Extended Prefill-Decode)

SGLang 进一步提出了EPD 三层分离,专为多模态和超大模型设计:

scss 复制代码
┌──────────┐    ┌──────────┐    ┌──────────┐
│ Encoder   │    │ Prefiller │    │ Decoder  │
│ (多模态)   │───→│ (文本预填充)│───→│ (逐token) │
│           │    │          │    │          │
│ 视觉编码器 │    │ 文本处理   │    │ 自回归生成 │
│ GPU: V100 │    │ GPU: H100 │    │ GPU: H100│
└──────────┘    └──────────┘    └──────────┘

七、投机解码(Speculative Decoding)

7.1 架构实现

SGLang 支持 Eagle、Medusa 等投机解码方案:

python 复制代码
class SpeculativeDecoder:
    """投机解码实现"""
    
    def __init__(self, draft_model, target_model):
        self.draft_model = draft_model    # 小模型(草稿模型)
        self.target_model = target_model  # 大模型(目标模型)
        self.spec_len = 5                 # 每次投机生成5个token
    
    def forward(self, batch: ScheduleBatch):
        # Phase 1: 草稿模型快速生成 spec_len 个 token
        draft_tokens = self.draft_model.generate(batch, num_tokens=self.spec_len)
        
        # Phase 2: 目标模型一次性验证所有 draft tokens
        # 关键:利用 RadixTree 的前缀复用
        # 草稿 token 的 KV Cache 可以被目标模型复用
        target_logits = self.target_model.forward_with_draft(batch, draft_tokens)
        
        # Phase 3: 从左到右验证,接受匹配的 token
        accepted = self._verify(draft_tokens, target_logits)
        
        # Phase 4: 被拒绝的 token 的 KV Cache 需要淘汰
        # RadixTree 自动处理:回滚到接受点,淘汰后续 KV Cache
        self.radix_tree.rollback(batch, accepted_len=accepted)
        
        return accepted

7.2 与 RadixAttention 的协同

投机解码的拒绝回滚是核心难点,RadixAttention 提供了天然的回滚机制:

markdown 复制代码
投机: [t1, t2, t3, t4, t5]  (草稿模型生成)
验证: [t1, t2, t3, ✗, ...]  (目标模型验证,t4 被拒绝)

回滚:
  1. RadixTree 回滚到 t3 的状态
  2. t4, t5 的 KV Cache 从基数树中删除
  3. TokenToKVPool 释放 t4, t5 对应的物理槽位
  4. 从 t3 开始用目标模型重新生成

八、Batch 数据结构与内存管理

8.1 ScheduleBatch 数据结构

源码位置:python/sglang/srt/managers/schedule_batch.py

python 复制代码
class ScheduleBatch:
    """最上层的 batch 结构,与 scheduler 交互"""
    
    # ---- 请求级信息 ----
    reqs: List[Req]                    # batch 中的请求列表
    forward_mode: ForwardMode          # EXTEND or DECODE
    
    # ---- 输入数据 ----
    input_ids: torch.Tensor            # [total_tokens] 输入 token id
    positions: torch.Tensor            # [total_tokens] 位置编码
    seq_lens: torch.Tensor             # [batch_size] 每个请求的序列长度
    
    # ---- KV Cache 索引 ----
    prefix_indices: List[List[int]]    # 每个请求的前缀 KV Cache 索引
    extend_num_tokens: List[int]       # 每个请求需要扩展的 token 数
    
    # ---- 采样参数 ----
    temperatures: torch.Tensor         # [batch_size]
    top_ps: torch.Tensor               # [batch_size]
    top_ks: torch.Tensor               # [batch_size]
    
    # ---- 约束解码 ----
    grammar_states: List[int]          # 每个请求的 FSM 状态

8.2 Req 对象生命周期

python 复制代码
class Req:
    """单个请求的完整生命周期"""
    
    # ---- 输入 ----
    rid: str                           # 请求唯一 ID
    token_ids: List[int]               # 完整的输入 token 序列
    origin_input_text: str             # 原始输入文本
    
    # ---- KV Cache 状态 ----
    prefix_indices: List[int]          # 前缀匹配的 KV Cache 索引
    extend_lens: int                   # 需要扩展的 token 数
    kv_indices: List[int]              # 当前所有 KV Cache 的物理索引
    
    # ---- 输出 ----
    output_ids: List[int]              # 已生成的 token id
    output_text: str                   # 已生成的文本
    
    # ---- 约束解码 ----
    grammar: Optional[Grammar]         # 约束语法
    fsm_state: int                     # 当前 FSM 状态
    
    # ---- 采样参数 ----
    sampling_params: SamplingParams

8.3 采样在 GPU 上完成

python 复制代码
# 采样在 GPU 上完成,而非 CPU
# 原因:logits[batch_size, vocab_size] 的 D2H 传输开销太大
# SGLang 的采样 Kernel 直接在 GPU 上执行 top-k, top-p, temperature

def sample(logits, sampling_params):
    """GPU 采样"""
    # 1. 应用 temperature
    logits = logits / sampling_params.temperature
    
    # 2. 应用 top-k
    top_k_logits, top_k_indices = torch.topk(logits, sampling_params.top_k)
    
    # 3. 应用 top-p (nucleus sampling)
    sorted_logits = torch.sort(top_k_logits, descending=True)
    cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
    ...
    
    # 4. 采样
    probs = torch.softmax(top_k_logits, dim=-1)
    token_id = torch.multinomial(probs, num_samples=1)
    
    return token_id

九、张量并行(TP)实现

9.1 多进程 TP 架构

scss 复制代码
┌──────────────────────────────────────────────────────┐
│  TP=2 的进程架构                                       │
│                                                      │
│  进程0 (Rank 0)          进程1 (Rank 1)               │
│  ┌────────────────┐      ┌────────────────┐          │
│  │ Scheduler      │      │ (无 Scheduler) │          │
│  │ + RadixTree    │      │                │          │
│  │ + TokenToKVPool│      │                │          │
│  └───────┬────────┘      └───────┬────────┘          │
│          │                       │                    │
│  ┌───────▼────────┐      ┌───────▼────────┐          │
│  │ TPModelWorker  │      │ TPModelWorker  │          │
│  │ (GPU 0)        │NCCL  │ (GPU 1)        │          │
│  │ 权重: 前半部分  │◄────►│ 权重: 后半部分  │          │
│  │ KV: 前半heads  │      │ KV: 后半heads  │          │
│  └────────────────┘      └────────────────┘          │
└──────────────────────────────────────────────────────┘

关键实现

  • 只有 Rank 0 运行 Scheduler,其他 Rank 只执行模型推理
  • 通信通过 NCCL(AllReduce/AllGather)
  • 权重按行/列拆分到不同 GPU
  • KV Cache 的每个 head 完整存储在一个 GPU 上(不拆分)

十、关键技术总结

10.1 SGLang 的核心创新点与对应源码位置

创新点 源码位置 核心原理
RadixAttention srt/layers/radix_attention.py 基数树管理 KV Cache,跨请求前缀复用
零开销调度 runtime/scheduler.py (overlap) CPU 调度与 GPU 计算流水线重叠
压缩 FSM srt/constrained/ 合并等价 DFA 状态,高效约束解码
LSPF 调度 runtime/scheduler.py 最长共享前缀优先调度
PD 分离 runtime/scheduler.py (disagg) Prefill/Decode 分离到不同 GPU
Extend 模式 srt/layers/attention/ 复用已有 KV Cache 的增量 prefill
TokenToKVPool srt/layers/radix_attention.py 统一物理内存池,索引式访问

10.2 一次完整请求的执行流程

less 复制代码
1. HTTP 请求到达 → FastAPI 路由
2. TokenizerManager: 文本 → token_ids
3. token_ids 通过 ZMQ 发送到 Scheduler
4. Scheduler:
   a. RadixTree.match_prefix(token_ids) → 获取匹配前缀和剩余部分
   b. TokenToKVPool.alloc(remaining_len) → 为新 KV 分配物理槽位
   c. 将请求加入 waiting_queue
5. Scheduler 组 batch:
   a. LSPF 排序 waiting_queue
   b. 选择请求组成 ScheduleBatch
   c. 构建 input_ids, positions, kv_indices 等
6. TPModelWorker.run_batch():
   a. ForwardMode.EXTEND: 只计算新增 token 的 KV + 注意力
   b. FlashInfer kernel 执行注意力计算
7. 采样:
   a. GPU 上执行 top-k/top-p/temperature
   b. 如果有约束: CompressedFSM mask logits
8. 后处理:
   a. 检查是否生成了 EOS
   b. 更新 RadixTree(插入新计算的 KV Cache)
   c. 如果是 decode,继续下一轮
   d. 如果生成了 EOS,释放资源
9. DetokenizerManager: token_ids → 文本,流式返回

10.3 SGLang 的工程哲学

SGLang 的设计哲学可以总结为**"让 KV Cache 跨越请求边界"**:

  • vLLM 的 PagedAttention 解决了单请求内的 KV Cache 碎片化
  • SGLang 的 RadixAttention 解决了跨请求的 KV Cache 复用
  • 从"每次请求都从头计算"到"自动识别并复用已有计算结果"
  • 这不仅仅是内存管理优化,更是改变了推理的语义模型

一句话总结 :SGLang 的核心洞察是------在 LLM 推理中,KV Cache 是最宝贵的资源 ,而传统框架将其视为"一次性消耗品"。RadixAttention 将 KV Cache 提升为可持久化、可复用、可淘汰的一等公民,从而在 Agent、多轮对话、RAG 等真实场景中实现了数量级的性能提升。

相关推荐
码途潇潇1 小时前
Codex Rules 与 Skills:项目级和全局级配置一览
人工智能
饼干哥哥1 小时前
字节Seedance2.5终于上线,这次在收割谁?
人工智能·设计模式·前端框架
小渔村的拉线工1 小时前
1.HPM6E80 解析芯片整体工作原理和存储架构
单片机·嵌入式硬件·mcu·架构·hpm6e80
小程故事多_801 小时前
Claude Code与Codex六大组件拆解,Coding Agent怎么把任务做稳
人工智能
董员外2 小时前
RAG 系统进化论(六):GraphRAG(基于知识图谱的 RAG),从相似文本走向实体关系
人工智能·后端·设计模式
Summer-Bright2 小时前
AI 软件简报 07.29-08.02:OpenAI 降价 80%、DeepSeek 全开源、欧盟动刀
人工智能·ai·开源·ai软件
武子康2 小时前
模型发布可以按天看,生产默认模型不能按天切:一套可回退的 30 天验收流程
人工智能·llm·agent
天天鸭2 小时前
5 万处中文的老项目实现国际化,如何用架构思维完成改造?
前端·javascript·架构
神经蛙19962 小时前
我用 TRAE Work 做周报:从 4 小时到 45 分钟的完整实操,建议收藏
人工智能