本文已于 2026.02.23 发表于公众号知乎

styled-头图2-style5

1. 背景

mini-sglang 不仅实现了大模型推理的核心功能,更在架构设计上体现出工业级推理引擎的关键特征 —— 多进程架构支撑、功能模块高内聚拆分、关键节点可扩展设计。基于这一架构叠加新功能时,效率和稳定性优势将非常显著。因此,mini-sglang 不仅仅是用于大模型推理关键功能的入门学习材料,更是一套兼具实践价值与参考意义的工业化推理引擎架构学习案例。

本文可视为《基于 nano-vLLM 学习大模型推理关键功能》的进阶,重点介绍 nano-vllm 没有涉及到的推理加速的核心功能。

2. 简介

mini-sglang 是和 nano-vllm 类似的小型推理引擎,但 mini-sglang 的功能更全面,代码量也达到了 5000 多行。它的关键功能包括:

  • 连续批处理(Continuous Batching)
  • KV 缓存(Prefix KV Cache / Paged KV Cache)
  • 高性能编译与执行优化(Torch Compilation、Triton、CUDA Graph)
  • 张量并行(Tensor Parallelism)
  • 分块预填充(Chunked Prefill)
  • 重叠调度(Overlap Scheduling)
  • 不同阶段支持不同的注意力后端

上述功能中,文字加粗部分是 nano-vllm 没有支持的。除此之外,mini-sglang 还支持了一些 nano-vllm 没有的周边能力:

  • 兼容 OpenAI 协议的在线服务
  • 交互式 Shell 模式
  • 多进程

该框架极具入门学习价值,本文将先介绍 mini-sglang 的整体架构,再对 mini-sglang 特有的核心技术要点深入解析。

3. 系统架构

3.1. 进程视角架构

styled-图1-style5

图 1 (图来源:https://lmsys.org/blog/2025-12-17-minisgl/)

mini-sglang 包括 4 类进程:

  • 入口进程,负责对外接口的 HTTP 服务、交互式 Shell
  • Tokenizer 进程,负责对 prompt 做分词,获取 token ID
  • Detokenizer 进程,负责将 token ID 转成文本
  • Scheduler 进程(也就是 TP worker 进程)

3.2. 类层面架构

styled-图2-style5
图 2

上图中六种颜色代表系统的六个组成部分

  • 浅黄色,HTTP 入口
  • 浅蓝色,请求调度
  • 浅橙色,文本和 token id 互转
  • 浅红色,KV cache 管理
  • 浅绿色,模型推理
  • 浅紫色,权重加载和矩阵计算的封装

3.3. 可扩展设计点

相较于 nano-vLLM,mini-sglang 在类层面架构更为复杂,同时在一些关键节点上做了可扩展设计:

(1)CacheManager(KV Cache 管理模块)

以 BaseCacheManager 为基类,提供两类实现:一是可用的 RadixCacheManager;二是空实现的 NaiveCacheManager,可直接用于对比有无 prefix KV Cache的性能差异,为性能优化提供基准参照。

(2)MHAKVCache(KV Cache 存储模块)

以 BaseKVCache 为基类,对 KV Cache 的存储层做了可扩展设计,当前已实现 MHAKVCache 作为核心存储方案,预留了后续接入其他存储逻辑的扩展空间。

(3)Attention Backend(注意力后端模块)

以 BaseAttnBackend 为基类,当前已支持两类主流后端:FlashInferBackend、FlashAttentionBackend,支持根据部署环境和性能需求灵活选型,并通过 HybridBackend 实现 prefill 和 decode 阶段使用不同后端的混合策略,充分发挥各后端在不同阶段的性能优势。

(4)模型实现模块

以 BaseLLMModel 为基类,Qwen3ForCausalLM 作为具体实现类,后续可快速扩展适配其他 LLM 模型。

除了关键节点的可扩展设计外,在代码封装复用上也有较多考虑,譬如:提炼出注意力子层 RopeAttn 类,可以在多个模型中复用。

3.4. 源码层面划分

目录结构如下:

minisgl/
├── attention   # 注意力
├── benchmark   # 性能评测
├── distributed # 通信
├── engine      # 推理引擎入口
├── kernel      # 内核(CU/CPP)
├── kvcache     # kvcache
├── layers      # 模型推理的层
├── llm         # 离线推理接口
├── message     # 进程间通信的消息
├── models      # 模型实现
├── scheduler   # 请求调度
├── server      # HTTP 服务
├── tokenizer   # 文本和 token id 转换
└── utils       # 工具函数和类

源码结构拆分很清晰,每个目录职责单一、边界明确。结合上文类架构、可扩展设计的分析,mini-sglang 不仅实现了基础的推理功能,更是在架构设计、代码内聚、接口抽象上都做了工业级的考量,这是工业化推理引擎的特征:既要考虑功能和性能,也要考虑迭代效率和可维护性。

下面的章节将对 mini-sglang 特有的重点功能展开介绍。

4. shell 模式

(1)启动

python -m minisgl --model /data/modelscope/Qwen3-0.6B --shell

效果示例:

styled-图3-style5
图 3

(2)实现方法

shell 模式通过 prompt_toolkit 包实现,下面给一份示例代码:

from prompt_toolkit import PromptSession
from prompt_toolkit.history import InMemoryHistory
from prompt_toolkit.auto_suggest import AutoSuggestFromHistory

async def main():
    session = PromptSession(history=InMemoryHistory())
    print("--- Welcome to demo Interactive Shell ---")
    while True:
        try:
            text = await session.prompt_async(
                'sglang >>> ', 
                auto_suggest=AutoSuggestFromHistory()
            )
            
            if text.lower() in ("exit", "quit"):
                break
                
            print(f"Model Response: (Simulated) Hello for '{text}'")
            
        except KeyboardInterrupt:
            continue  # Ctrl+C 不退出,模拟真实 Shell
        except EOFError:
            break     # Ctrl+D 退出

if __name__ == "__main__":
    import asyncio
    asyncio.run(main())

效果示例:

styled-图4-style5
图 4

5. 分块预填充 Chunked Prefill

5.1. 概念理解

Chunked Prefill 是大语言模型(LLM)针对长序列输入场景设计的优化策略,核心目标是解决长序列 prefill 阶段的显存溢出(OOM) 和请求调度阻塞问题。通过将超长输入序列切分为多个小 chunk 逐块处理,既降低了单次计算的显存峰值,又让系统能及时响应新请求。

(1)解决请求调度阻塞,保持系统及时响应

常规 prefill 需一次性处理整个长序列,计算耗时会随序列长度增长而增加。在此期间,新请求需要排队等待当前长序列 prefill 完成才能被调度,导致用户感知到响应延迟、系统卡顿。

而 Chunked Prefill 把长任务拆成多个短任务,调度器可以让新请求在前后批次 chunk 之间或者是和 chunk 一起执行,从而提升系统的并发响应能力。

(2)降低显存峰值,避免 OOM

我们很容易有这样的疑问:整个 prefill 完成后,KV Cache 的总占用量与不分块时一样,为什么能解决 OOM?

实际上这里降低的是中间激活值显存,不是 KV Cache 显存。模型推理的显存占用分为三部分:模型权重(固定值)、KV Cache(与总序列长度正相关)、中间激活值(与单次计算的序列长度强相关)。

Chunked Prefill 优化的是中间激活值显存:注意力计算中的 Q×Kᵀ 操作产出中间矩阵,其元素个数为序列长度 n 的平方,是中间激活值的主要来源。分 chunk 后,n 大幅度降低,相应中间矩阵随之急剧降低,并且每个 chunk 计算完成后,该 chunk 的中间激活值会被立即释放,然后再处理下一个 chunk。最终实现显存峰值的显著下降,从而避免 OOM。

5.2. 核心思路

styled-图5-style5
图5

如上图所示,一个包含 3 个 chunk 的请求,分三次迭代执行完 prefill,每次处理一个 chunk,第 2、3 个 chunk 处理时,需要使用到前序 chunk 的 KV Cache。

5.3. 执行流程

styled-图6-style5
图 6

在 Scheduler 调度请求时,会进行 Chunked Prefill 的处理,其核心流程如上图所示:判断请求是否可以一次推理完,如果不行则创建 ChunkedReq,然后和常规请求执行一样的推理流程。其核心机制包括以下 3 点:

(1)KV Cache 复用

Chunked Prefill 可以视为多个串行执行的、有相同前缀的请求。后继 chunk 需要复用前序 chunk 的 KV Cache。

  • 不经过 Radix Cache:与 SGLang 的实现不同,mini-sglang 通过 TableManager 记录的请求信息实现 KV Cache indices 追踪和复用
  • 共享机制:所有 chunk 共享同一个 table_idx,都访问 page_table(即 KV Cache indices) 的同一行
  • 累积写入:每个 chunk 都将新生成的 KV Cache 物理页索引追加到 page_table 中,后续 chunk 可以访问到完整的 KV Cache

(2)Meta 信息复用

分 chunk 也可以理解为请求的续跑,续跑时会复用首次运行记录的信息,避免重复获取。

  • 检测续传:遍历队列时检测 pending_req.chunked_req 是否存在,如果存在,则说明这是非首 chunk
  • 复用资源:续传请求直接复用前序 chunk 的 table_idx 和 cache_handle(可复用的前缀 KV Cache 信息)
  • 状态追踪:cached_len 始终表示 Radix Tree 中的缓存长度(不变),device_len 表示已处理长度(累积增长)

(3)所有 Chunk 完成后才能执行 Decode

  • ChunkedReq 标记:chunk 请求使用 ChunkedReq 类型(继承自 Req 类型),其 can_decode() 返回 False
  • 末尾 chunk 的处理:当剩余 tokens 可以一次处理完时,创建普通 Req 对象,允许进入 decode 阶段

6. 重叠调度 Overlap Scheduling

6.1. 概念理解

定义:Overlap Scheduling 是一种将 CPU 调度开销与 GPU 计算重叠执行的优化技术,mini-sglang 通过双 CUDA Stream 机制实现 CPU 和 GPU 并行工作,有效隐藏 CPU 延迟,提高 GPU 利用率和系统吞吐。

核心价值:Overlap Scheduling 的核心目标是让 CPU 的调度操作与 GPU 的推理计算同时进行。由于 CPU 资源相对廉价且通常不是瓶颈,而 GPU 算力昂贵且稀缺,因此该技术的核心价值在于:让瓶颈硬件(GPU)保持持续的满负荷计算状态,避免其因等待 CPU 指令而 “空转”。

技术手段:能够并行工作,本质上是让 GPU 计算可以异步执行,多流是实现 GPU 异步计算的关键手段。对于熟悉传统 CPU 后台开发的工程师来说,这类似于经典的 “生产者 - 消费者” 模型:外部消息响应线程在遇到计算密集型任务时,通常会将计算任务丢到工作线程里,待计算完成后再返回给消息响应线程,而消息响应线程则通过轮询回包队列或者事件触发的方式,将计算结果回给外部请求者。mini-sglang 就是采用这类思路实现的,CUDA 的工作流就是 CPU 里的工作线程。

扩展:本章介绍的是调度的重叠,大模型推理中还有其他重叠,譬如:数据传输和计算的重叠、核函数启动和计算的重叠。无论哪种重叠,其本质都是为了掩盖非瓶颈环节的延迟,从而让昂贵的瓶颈硬件(GPU)得到最充分的利用。

6.2. 核心技术点

如前所述,Overlap Scheduling 的本质技术是异步执行,而异步执行的核心有两点:

创建异步环境,让推理计算可以在单独的流里执行

置位同步信号,推理计算开始前需要确保数据已经准备好,调度线程处理结果前需要确保推理已经完成

mini-sglang 的实现很简洁,关键代码如下:

(1)创建推理流

# Engine 对象的 __init__:
self.stream = torch.cuda.Stream()
...
# Scheduler 对象的 __init__:
self.engine_stream_ctx = torch.cuda.stream(self.engine.stream)

(2)Stream 同步:

with self.engine_stream_ctx:  # 切换到推理流
    self.engine.stream.wait_stream(self.stream)  # 等待调度流完成元数据准备
    ongoing_data = (forward_input, self._forward(forward_input)) 

(3)Event 同步:

def _process_last_data(self, last_data, ongoing_data):
    if last_data is None:
        return
    batch, (_, next_tokens_cpu, copy_done) = last_data[0].batch, last_data[1]
    copy_done.synchronize()  # 等待 GPU→CPU 数据拷贝完成
    # 处理采样结果... 

6.3. 推理时序图

styled-图7-style5
图 7

6.4. 原型代码

基于 mini-sglang 的实现,我们可以自行实现一份 demo 代码:

"""Overlap Scheduling 原型"""

import torch
import time
import queue

# 模拟各 CPU 阶段耗时(秒)
PRE_PROCESS_TIME = 0.03  # 前处理(CPU: 接收请求、准备 Metadata, Tokenization)
POST_PROCESS_TIME = 0.02  # 后处理(CPU: Sample, De-tokenization)
TOTAL_REQUESTS = 10  # 测试总请求数


class OverlapEngine:
    def __init__(self):
        self.engine_stream = torch.cuda.Stream()
        self.request_queue = queue.Queue()
        self.results_queue = queue.Queue()

    def _pre_process(self, req_id):
        """模拟前处理"""
        time.sleep(PRE_PROCESS_TIME)

    def _post_process_normal(self, req_id):
        """模拟 normal 模式的后处理"""
        time.sleep(POST_PROCESS_TIME)

    def _post_process_overlap(self, last_data):
        """模拟 overlap 模式的后处理"""
        if last_data is not None:
            _, sync_event, _ = last_data
            sync_event.synchronize()
            time.sleep(POST_PROCESS_TIME)

    def _mock_forward(self):
        """模拟 GPU 推理"""
        dummy_tensor = torch.randn(1000, 1000, device="cuda")
        for _ in range(1000):
            dummy_tensor = torch.matmul(dummy_tensor, dummy_tensor)

        return dummy_tensor

    def _gpu_inference_async(self, req_id):
        """模拟异步 GPU 推理"""
        with torch.cuda.stream(self.engine_stream):
            result = self._mock_forward()
            sync_event = torch.cuda.Event()
            sync_event.record(self.engine_stream)
            return req_id, sync_event, result

    def _gpu_inference_blocking(self):
        """模拟同步 GPU 推理"""
        result = self._mock_forward()
        torch.cuda.synchronize()
        return result

    def run_normal(self):
        """模拟运行 normal 模式"""
        start_time = time.time()
        for i in range(TOTAL_REQUESTS):
            self._pre_process(i)
            self._gpu_inference_blocking()
            self._post_process_normal(i)
        return time.time() - start_time

    def run_overlap(self):
        """模拟运行 overlap 模式"""
        start_time = time.time()

        last_data = None

        for i in range(TOTAL_REQUESTS + 1):
            ongoing_data = None
            if i < TOTAL_REQUESTS:
                self._pre_process(i)
                ongoing_data = self._gpu_inference_async(i)

            self._post_process_overlap(last_data)
            last_data = ongoing_data

        return time.time() - start_time


# 预热
dummy_tensor = torch.randn(1000, 1000, device="cuda")
torch.matmul(dummy_tensor, dummy_tensor)

engine = OverlapEngine()
print(f"开始测试 {TOTAL_REQUESTS} 个请求...")

duration_normal = engine.run_normal()
throughput_normal = TOTAL_REQUESTS / duration_normal

duration_overlap = engine.run_overlap()
throughput_overlap = TOTAL_REQUESTS / duration_overlap

print("-" * 30)
print(f"Normal  模式耗时: {duration_normal:.4f}s | 吞吐: {throughput_normal:.2f} req/s")
print(
    f"Overlap 模式耗时: {duration_overlap:.4f}s | 吞吐: {throughput_overlap:.2f} req/s"
)
print(f"提升效率: {((throughput_overlap/throughput_normal)-1)*100:.2f}%")

效果如下:

开始测试 10 个请求...
------------------------------
Normal  模式耗时: 1.2408s | 吞吐: 8.06 req/s
Overlap 模式耗时: 0.8476s | 吞吐: 11.80 req/s
提升效率: 46.39%

从原型代码的运行结果看到有 46.39% 的吞吐提升收益。

6.5. 深度解析

前文介绍了 Overlap Scheduling 的基础知识,在业务实际应用中,还需要解答两个问题:

  • Overlap Scheduling 为什么会导致 TTFT 上涨?有什么优化手段?
  • Overlap Scheduling 消除 GPU 空闲气泡,系统吞吐和 TTFT 就一定会更好吗?

解答见文章:《大模型推理加速:Overlap Scheduling 的深入剖析与性能权衡艺术》

7. 张量并行(TP)

mini-sglang 的 TP 和 nano-vllm 大同小异,并且都只执行单机内的 TP,不同点在于控制消息的通信,nano-vllm 采用共享内存,而 mini-sglang 使用 CPU 通道的 pytorch dist 和 ZMQ 混合。下面重点讲解 mini-sglang 的请求在 TP rank 之间是如何同步的。

7.1. mini-sglang 通信视角

(1)mini-sglang 请求通信和数据张量通信示意图

styled-图8-style5
图 8

(2)三种推理引擎的请求通信对比:

特性 mini-sglang nano-vllm SGLang
请求元数据同步 PyTorch Dist (个数) + ZMQ (详情) 共享内存 (Shared Memory) PyTorch Dist (CPU Group)
通信开销 中等(涉及序列化与 Socket 交互) 极低(零拷贝,直接读内存) 较高(依赖 NCCL/Gloo CPU 进程间通信)

mini-sglang 使用 CPU 通道的 pytorch dist 传递请求个数,然后使用 ZMQ 传递请求详情。

7.2. 请求通信示例代码

下面代码展示 TP 模式的请求通信,涉及 CPU 通道的 pytorch dist 和 ZMQ 组件。

#!/usr/bin/env python3
"""
演示:使用 PyTorch Dist + ZMQ 进行请求同步
============================================

本演示展示如何在多个 rank 之间同步请求:
1. 使用 CPU 通道的 PyTorch Dist 广播请求个数
2. 使用 ZMQ (PUB/SUB) 广播请求详情

架构说明:
- Rank 0 (主节点): 接收请求,通过 PyTorch Dist 广播请求个数,通过 ZMQ 广播请求详情
- Rank 1+ (工作节点): 通过 PyTorch Dist 接收请求个数,通过 ZMQ 接收请求详情

使用方法:
    # 默认 TP size=3
    python demo_req_sync.py
    
    # 指定 TP size
    python demo_req_sync.py --tp-size 4
"""

import argparse
import os
import pickle
import time
from dataclasses import dataclass
from typing import List
from multiprocessing import Process

import torch
import torch.distributed as dist
import zmq


@dataclass
class FakeRequest:
    request_id: str
    prompt: str
    max_tokens: int
    temperature: float
    
    def __repr__(self):
        return f"FakeRequest(id={self.request_id}, prompt='{self.prompt[:20]}...', max_tokens={self.max_tokens})"


class RequestSyncDemo:
    def __init__(self, rank: int, world_size: int):
        self.rank = rank
        self.world_size = world_size
        self.is_primary = (rank == 0)
        
        self._init_pytorch_dist()
        self._init_zmq()
        
        print(f"[Rank {self.rank}] Initialized successfully")
    
    def _init_pytorch_dist(self):
        """Initialize PyTorch Distributed (CPU backend for control messages)"""
        # Use gloo backend for CPU communication
        dist.init_process_group(
            backend='gloo',
            init_method='tcp://127.0.0.1:12345',
            rank=self.rank,
            world_size=self.world_size
        )
        
        print(f"[Rank {self.rank}] PyTorch Dist initialized (backend=gloo)")
    
    def _init_zmq(self):
        """Initialize ZMQ for request details broadcasting"""

        zmq_broadcast_addr = "tcp://127.0.0.1:23456"
        self.zmq_context = zmq.Context()
        
        if self.is_primary:
            # Rank 0: Create PUB socket to broadcast requests
            self.zmq_pub_socket = self.zmq_context.socket(zmq.PUB)
            self.zmq_pub_socket.bind(zmq_broadcast_addr)
            print(f"[Rank {self.rank}] ZMQ PUB socket bound to {zmq_broadcast_addr}")
        else:
            # Rank 1, 2: Create SUB socket to receive requests
            self.zmq_sub_socket = self.zmq_context.socket(zmq.SUB)
            self.zmq_sub_socket.connect(zmq_broadcast_addr)
            self.zmq_sub_socket.setsockopt_string(zmq.SUBSCRIBE, "")  # Subscribe to all messages
            print(f"[Rank {self.rank}] ZMQ SUB socket connected to {zmq_broadcast_addr}")
    
    def generate_fake_requests(self, num_requests: int) -> List[FakeRequest]:
        """Generate fake requests (only for Rank 0)"""
        requests = []
        for i in range(num_requests):
            req = FakeRequest(
                request_id=f"req_{int(time.time())}_{i}",
                prompt=f"This is a test prompt number {i} for demonstration purposes",
                max_tokens=100 + i * 10,
                temperature=0.7 + i * 0.1
            )
            requests.append(req)
        return requests
    
    def broadcast_requests_rank0(self, requests: List[FakeRequest]):
        """Rank 0: Broadcast requests to all other ranks"""
        num_requests = len(requests)
        
        count_tensor = torch.tensor(num_requests, dtype=torch.long)
        dist.broadcast(count_tensor, src=0)
        
        for req in requests:
            serialized_req = pickle.dumps(req)
            self.zmq_pub_socket.send(serialized_req)
        
        print(f"[Rank {self.rank}] ✓ Finished broadcasting {num_requests} requests")
        return requests
    
    def receive_requests_rank_n(self) -> List[FakeRequest]:
        """Rank 1: Receive requests from Rank 0"""
        count_tensor = torch.tensor(0, dtype=torch.long)
        dist.broadcast(count_tensor, src=0)
        num_requests = int(count_tensor.item())
        
        print(f"[Rank {self.rank}] Received request count: {num_requests}")
        
        requests = []
        for i in range(num_requests):
            serialized_req = self.zmq_sub_socket.recv()
            req = pickle.loads(serialized_req)
            requests.append(req)
            print(f"[Rank {self.rank}] Received via ZMQ [{i+1}/{num_requests}]: {req}")
        
        print(f"[Rank {self.rank}] ✓ Finished receiving {num_requests} requests")
        return requests
    
    def run_demo(self):
        print(f"[Rank {self.rank}] Starting Request Sync Demo")
        
        if self.is_primary:
            requests = self.generate_fake_requests(3)
            self.broadcast_requests_rank0(requests)
            
        else:
            requests = self.receive_requests_rank_n()
        
        print(f"[Rank {self.rank}] Demo completed successfully! ✓")
    
    def cleanup(self):
        if self.is_primary:
            self.zmq_pub_socket.close()
        else:
            self.zmq_sub_socket.close()
        self.zmq_context.term()
        dist.destroy_process_group()


def run_rank_process(rank: int, world_size: int):
    """Run a single rank process"""
    demo = RequestSyncDemo(rank=rank, world_size=world_size)
    time.sleep(1)
    
    try:
        demo.run_demo()
    except KeyboardInterrupt:
        print(f"\n[Rank {rank}] Interrupted by user")
    except Exception as e:
        print(f"\n[Rank {rank}] Error: {e}")
        import traceback
        traceback.print_exc()
    finally:
        demo.cleanup()


def main():
    parser = argparse.ArgumentParser(description="Request Synchronization Demo")
    parser.add_argument("--tp-size", type=int, default=3, help="Tensor parallel size (default=3)")
    args = parser.parse_args()
    
    world_size = args.tp_size
    print(f"\n{'='*60}")
    print(f"Starting {world_size} processes for request synchronization demo")
    print(f"{'='*60}\n")
    
    processes = []
    for rank in range(world_size):
        p = Process(target=run_rank_process, args=(rank, world_size))
        p.start()
        processes.append(p)
        print(f"[Main] Started process for Rank {rank} (PID: {p.pid})")
    
    print(f"\n[Main] All {world_size} processes started. Waiting for completion...\n")
    
    try:
        for rank, p in enumerate(processes):
            p.join()
    except KeyboardInterrupt:
        print(f"\n[Main] Interrupted by user. Terminating all processes...")
        for p in processes:
            p.terminate()
        for p in processes:
            p.join()
    

if __name__ == "__main__":
    main()

关键技术点:

  • CPU pytorch dist 使用 gloo 后端,并且需要占用一个网络端口
  • ZMQ 需要占用一个网络端口,发送者使用 PUB socket,接收者使用 SUB socket

8. 注意力后端的多态封装简介

8.1. 概念理解

mini-sglang 针对注意力后端构建了标准化的多态封装层,对不同注意力后端的底层接口进行抽象与统一封装,实现了调用层与实现层的解耦。上层调用方仅需基于统一的标准接口完成集成调用,无需感知和适配各类注意力后端的内部实现逻辑。当前业界主流的注意力计算后端为 FlashAttention 与 FlashInfer,在 Hopper 架构下采用「FA3 负责 Prefill Attention 计算、FlashInfer 负责 Decode Attention 计算」的混合模式,可获得最佳的性能收益,该混合注意力后端的使用模式已被 mini-sglang 支持,可以由 mini-sglang 根据 NVIDIA 硬件做自动配置,也可以手动配置。

8.2. 关键技术点

(1)抽象注意力后端接口,让不同的后端都继承自 BaseAttnBackend,包括混合注意力后端

(2)基于装饰器的工厂注册模式,通过泛型实现工厂容器(工厂的工厂)

8.3. 多态类图

styled-图9-style5
图 9

Python 语言有「鸭子类型」的特点,该概念源于经典定义:"如果它走起来像鸭子,叫起来像鸭子,那它就是鸭子"。这也是 Python 作为动态语言的核心特征 —— 不关注对象的类型名,而是根据对象拥有的方法和属性来判定其类型。正因如此,Python 里的基类,其价值仅在于描述统一接口、完成代码协作、提升可读性,基类并非实现多态的必要选项

9. KV Cache 前缀索引

9.1. KV Cache 管理核心类

mini-sglang 的 prefix KV Cache 管理和 nano-vllm 用的链式 hash 不同,这里采用 radix tree 实现,核心类如下图。

styled-图10-style5
图 10

(1)TableManager

TableManager 记录 Scheduler 要使用的 KVCache 元数据信息,有三重点成员:

  • _free_slots 空闲的槽位。进程启动时会指定系统可以同时跑的最大请求数,空闲槽位就是当前可继续接收请求的槽位。
  • page_table KVCache 索引,二维数组,一行代表一个请求,行里的内容为 KVCache 的物理地址索引值。
  • token_pool token ID 列表,二维数组,一行代表一个请求,行里的内容为一个请求的 token ID 序列,包括 prompt 和新生成的 token ID。

举例说明:

# 槽位状态
_free_slots = [2, 3]  # 槽位 0、1 被占用,2、3 空闲

# page_table(物理页号)
page_table = [
    [5, 12, 7, 23, 0, 0, 0, 0],   # 槽位 0:请求 A 的 KV cache 位置
    [8, 15, 19, 0, 0, 0, 0, 0],   # 槽位 1:请求 B 的 KV cache 位置
    [0, 0, 0, 0, 0, 0, 0, 0],     # 槽位 2:空闲
    [0, 0, 0, 0, 0, 0, 0, 0],     # 槽位 3:空闲
]

# token_pool(Token ID)
token_pool = [
    [101, 2023, 2003, 1037, 0, 0, 0, 0],  # 槽位 0:请求 A 的 token 序列
    [101, 7592, 2088, 0, 0, 0, 0, 0],     # 槽位 1:请求 B 的 token 序列
    [0, 0, 0, 0, 0, 0, 0, 0],             # 槽位 2:空闲
    [0, 0, 0, 0, 0, 0, 0, 0],             # 槽位 3:空闲
]

(2)CacheManager

CacheManager 是 KVCache 管理的入口类,包括两部分核心内容:

  • _free_slots 空闲 KVCache 物理页列表,page size 当前都是 1,因此等价于 KVCache 物理存储的位置索引信息
  • manager 又是一个 cache manager,但这个是做前缀管理的。

(3)RadixCacheManager

mini-sglang 的 prefix KV Cache 管理和 nano-vllm 用的链式 hash 不同,这里采用 radix tree 实现。

9.2. KV Cache 管理类的多态设计

styled-图11-style5
图 11

和注意力后端的封装类似,也采用工厂模式实现 KV Cache 管理类的多态设计,工厂和插件化让代码更内聚。

9.3. KV Cache 管理核心技术点

(1)KV Cache 前缀复用

通过 Radix Tree 查找最长匹配前缀,实现 KVCache 的高效复用:

  • 匹配过程:在 RadixCacheManager 中查找与新请求 token 序列匹配的最长前缀路径
  • 结果存储:
    将匹配的 token 序列存储到 TableManager 的 token_pool 中
    将对应的 KVCache 物理页索引存储到 page_table 中
  • 引用保护:匹配路径上的所有 RadixTreeNode 增加引用计数(ref_count++),防止被驱逐释放
  • 核心价值:避免重复计算相同前缀的 KVCache,显著提升推理效率。

(2)KV Cache 加入前缀索引

在请求完成时,将整个请求的 KVCache 写入前缀索引:

  • 写入时机:请求完成后,在 Scheduler 处理 finished_reqs 时触发
  • 写入内容:包括 prompt 和所有已生成的 tokens
  • 增量策略:通过 insert_prefix 方法,只写入 Radix Tree 中尚不存在的新增部分
  • 去重处理:如果新写入的序列与已有路径重复,会自动合并到现有节点
  • 资源释放:写入完成后,释放重复部分的物理页,并解锁请求持有的引用计数
  • 核心价值:请求完成后统一写入,构建全局共享的前缀索引,为后续请求提供复用基础。

(3)KV Cache 的空间属性和按需驱逐

styled-图12-style5
图 12

KV Cache 的空间划分为三种核心属性,明确界定不同状态下的资源使用规则:

  • Protected(受保护):有请求正在使用,不可被驱逐释放
  • Evictable(可驱逐):已存入 Radix Tree 但无请求使用,必要时可释放回收
  • Free(空闲):未分配的可用空间

如图 12 所示,已分配(Protected + Evictable)的空间通过 Radix Tree 统一维护,空间回收采用按需驱逐(Lazy Eviction)策略,具体规则如下:

  • 触发条件:当分配新 KVCache 空间时,若空闲空间不足则触发驱逐
  • 驱逐策略:
    候选筛选:只驱逐 ref_count == 0(未被使用)的叶子节点
    优先级排序:基于节点创建时间(FIFO),优先驱逐最早创建的节点
    级联回收:驱逐叶子节点后,若父节点也变为未使用的叶子,继续向上驱逐
    实现机制:使用最小堆(Min-Heap)高效选择驱逐目标

核心价值:在缓存保留和空间回收之间取得平衡——最大化前缀复用率,同时保证空间可用性。

(4)KV Cache 物理空间预留

KV Cache 空间耗尽会导致推理服务陷入死锁:新请求因无空间无法进入,而正在运行的请求因无法分配后续 Decode 所需空间而阻塞,导致系统挂起。mini-sglang 与 nano-vllm 采用了不同的策略来规避这一风险:

  • nano-vllm(回滚重试策略):当空间不足时,将运行队列尾部的请求弹出,释放其占用的 KV Cache 空间,并将其作为新请求重新加入等待队列头部。代价: 被回滚的请求需要重新竞争资源,导致请求延迟(Latency)出现毛刺,且可能降低解码效率。

  • mini-sglang(资源预留策略):在请求进入 Decode 阶段前,强制预留其后续生成所需的最大 KV Cache 空间。代价: 虽然彻底避免了死锁,但激进的预留会导致并发度(Concurrency)受限。在高并发场景下,这可能会牺牲部分吞吐量(Throughput),导致部分请求在等待队列中产生不必要的等待。

核心价值: mini-sglang 选择了以 “吞吐量换稳定性” 的设计,确保在任何负载下服务都能持续推进。

10. TVM FFI

10.1. 概念理解

TVM FFI 是 2025 年推出的面向机器学习领域的高性能跨语言调用组件。相比传统 FFI,它在张量传递和算子调用上具有显著的性能优势。TVM FFI 主要支持两种部署模式:

(1)AOT(提前编译):预编译 C++ 代码为.so 动态库供 Python 调用,运行性能极致,无编译开销,但需维护编译产物,工程整合有一定成本。

(2)JIT(即时编译):基于 tvm_ffi.cpp 的 load_inline\load 接口实现运行时编译,搭配缓存规避重复编译,开发灵活,无需管理编译产物,缓存生效后性能与 AOT 持平。

注:mini-sglang 使用 JIT,没有使用 AOT,部分函数命名上叫做 load_aot 是不准确的。

tvm_ffi.cpp 模块的使用比较简单,可以参考官方文档:https://tvm.apache.org/ffi/reference/python/cpp/generated/tvm_ffi.cpp.load_inline.html#

10.2. 示例代码

下面举一个使用 load_inline demo:

"""
TVM_FFI 演示:使用 'functions' 或者 TVM_FFI_DLL_EXPORT_TYPED_FUNC 宏都可以导出 CPP 函数
"""

from __future__ import annotations

from functools import lru_cache
from typing import TYPE_CHECKING

from tvm_ffi.cpp import load_inline

if TYPE_CHECKING:
    from tvm_ffi import Module

# 共享的 C++ 模板代码
CPP_POWER_TEMPLATE = '''
#include <tvm/ffi/function.h>
#include <cmath>

// 基于模板的幂函数
template<int EXPONENT>
struct PowerCalculator {
    static auto compute(float base) -> float {
        float result = 1.0f;
        for (int i = 0; i < EXPONENT; ++i) {
            result *= base;
        }
        return result;
    }
    
    static auto run(float base) -> float {
        return compute(base);
    }
};
'''


@lru_cache(maxsize=None)
def load_power_module_macro_export(exponent: int) -> Module:
    """宏导出: Using TVM_FFI_DLL_EXPORT_TYPED_FUNC macro"""
    cpp_sources = [
        CPP_POWER_TEMPLATE,
        f'TVM_FFI_DLL_EXPORT_TYPED_FUNC(power_{exponent}, (PowerCalculator<{exponent}>::run));'
    ]
    
    return load_inline(
        f"demo_power_old_{exponent}",
        cpp_sources=cpp_sources,
    )


@lru_cache(maxsize=None)
def load_power_module_param_export(exponent: int) -> Module:
    """参数导出: Using 'functions' parameter (cleaner!)"""
    cpp_sources = [
        CPP_POWER_TEMPLATE,
        f'auto power_{exponent}(float base) {{ return PowerCalculator<{exponent}>::run(base); }}'
    ]
    
    return load_inline(
        f"demo_power_new_{exponent}",
        cpp_sources=cpp_sources,
        functions=f"power_{exponent}",
    )


@lru_cache(maxsize=None)
def load_power_module_param_export_with_docs(exponent: int) -> Module:
    """参数导出(带文档): Using dict for functions parameter"""
    cpp_sources = [
        CPP_POWER_TEMPLATE,
        f'auto power_{exponent}(float base) {{ return PowerCalculator<{exponent}>::run(base); }}'
    ]
    
    return load_inline(
        f"demo_power_docs_{exponent}",
        cpp_sources=cpp_sources,
        functions={f"power_{exponent}": f"Calculate base^{exponent}"},
    )


def demo_two_export_way():
    """对比宏导出和参数导出两种方式。"""
    print("=" * 60)
    print("对比:宏导出 vs 参数导出")
    print("=" * 60)
    
    base, exponent = 2.0, 3
    
    # 宏导出
    print(f"宏导出: TVM_FFI_DLL_EXPORT_TYPED_FUNC(power_{exponent}, ...)")
    module_macro = load_power_module_macro_export(exponent)
    result_macro = module_macro[f"power_{exponent}"](base)
    
    # 参数导出
    print(f"参数导出: functions='power_{exponent}'")
    module_param = load_power_module_param_export(exponent)
    result_param = module_param[f"power_{exponent}"](base)
    
    assert abs(result_macro - result_param) < 1e-6


def demo_functions_parameter_variants():
    """展示 'functions' 参数的不同用法。"""
    print("\n" + "=" * 60)
    print("参数导出的不同用法")
    print("=" * 60)
    
    print("\n1. 字符串: functions='power_2'")
    module = load_power_module_param_export(2)
    print(f"   调用: module.power_2(3.0) = {module.power_2(3.0)}")
    
    print("\n2. 列表: functions=['power_2', 'power_3']")
    demo_multiple_functions()
    
    print("\n3. 字典: functions={'power_4': 'docstring'}")
    module_docs = load_power_module_param_export_with_docs(4)
    print(f"   调用: module_docs.power_4(2.0) = {module_docs.power_4(2.0)}")


def demo_multiple_functions():
    """展示如何一次导出多个函数。"""
    cpp_sources = [
        CPP_POWER_TEMPLATE,
        'auto power_2(float base) { return PowerCalculator<2>::run(base); }',
        'auto power_3(float base) { return PowerCalculator<3>::run(base); }',
        'auto power_4(float base) { return PowerCalculator<4>::run(base); }',
    ]
    
    module = load_inline(
        "demo_multi_power",
        cpp_sources=cpp_sources,
        functions=["power_2", "power_3", "power_4"],
    )
    
    print("   调用示例:")
    print(f"  module.power_2(2.0) = {module.power_2(2.0)}")
    print(f"  module.power_3(2.0) = {module.power_3(2.0)}")
    print(f"  module.power_4(2.0) = {module.power_4(2.0)}")
    

def main():
    demo_two_export_way()
    demo_functions_parameter_variants()
    

if __name__ == "__main__":
    main()

10.3. 性能测试

在官方博客里提到使用 TVM FFI 绑定自定义的内核函数,性能比默认的 Pytorch 接口更好。这里也测试一下 fast_compare_key 函数,这个函数用于在 Radix Node 里获取节点 key 和输入的 input id 的最长相同前缀。

#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include <tvm/ffi/dtype.h>
#include <tvm/ffi/extra/c_env_api.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/object.h>
#include <algorithm>
#include <stdexcept>

namespace {

auto _is_1d_cpu_int_tensor(const tvm::ffi::TensorView tensor) -> bool {
  return tensor.ndim() == 1 && tensor.is_contiguous() &&
         tensor.device().device_type == kDLCPU &&
         (tensor.dtype().code == kDLInt) &&
         (tensor.dtype().bits == 32 || tensor.dtype().bits == 64);
}

auto fast_compare_key(const tvm::ffi::TensorView a,
                      const tvm::ffi::TensorView b) -> size_t {
  if (!(_is_1d_cpu_int_tensor(a) && _is_1d_cpu_int_tensor(b))) {
    throw std::runtime_error("Both tensors must be 1D CPU int tensors.");
  }
  if (a.dtype() != b.dtype()) {
    throw std::runtime_error("Tensors must have the same dtype.");
  }
  const auto a_ptr = a.data_ptr();
  const auto b_ptr = b.data_ptr();
  const auto common_len = std::min(a.size(0), b.size(0));
  if (a.dtype().bits == 64) {
    const auto a_ptr_64 = static_cast<const int64_t *>(a_ptr);
    const auto b_ptr_64 = static_cast<const int64_t *>(b_ptr);
    const auto diff_pos =
        std::mismatch(a_ptr_64, a_ptr_64 + common_len, b_ptr_64);
    return static_cast<size_t>(diff_pos.first - a_ptr_64);
  } else {
    const auto a_ptr_32 = static_cast<const int32_t *>(a_ptr);
    const auto b_ptr_32 = static_cast<const int32_t *>(b_ptr);
    const auto diff_pos =
        std::mismatch(a_ptr_32, a_ptr_32 + common_len, b_ptr_32);
    return static_cast<size_t>(diff_pos.first - a_ptr_32);
  }
}

} // namespace

TVM_FFI_DLL_EXPORT_TYPED_FUNC(fast_compare_key, fast_compare_key);
"""
性能对比测试:C++ 版本 vs Python 版本的 compare_key

测试场景:
1. 不同长度的数组(短、中、长)
2. 不同的匹配位置(开头不匹配、中间不匹配、完全匹配)
3. 不同的数据类型(int32、int64)
"""

import os
import time
import tracemalloc
from functools import lru_cache
from typing import Callable

import torch
import tvm_ffi.cpp


def python_compare_key(a: torch.Tensor, b: torch.Tensor) -> int:
    """纯 Python 实现的 compare_key"""
    a_np = a.numpy()
    b_np = b.numpy()
    common_len = min(len(a_np), len(b_np))

    for i in range(common_len):
        if a_np[i] != b_np[i]:
            return i
    return common_len


@lru_cache(maxsize=None)
def _load_local_fast_compare_key():
    """加载 cpp fast_compare_key 函数(带缓存)"""
    try:
        current_dir = os.path.dirname(os.path.abspath(__file__))
        cpp_file = os.path.join(current_dir, "compare_key.cpp")

        if not os.path.exists(cpp_file):
            raise FileNotFoundError(f"C++ 源文件不存在: {cpp_file}")

        module = tvm_ffi.cpp.load(name="compare_key_benchmark", cpp_files=cpp_file)
        return module.fast_compare_key
    except Exception as e:
        raise ImportError(f"无法编译本地 C++ 代码: {e}")


def cpp_compare_key(a: torch.Tensor, b: torch.Tensor) -> int:
    """C++ 实现的 compare_key(通过 TVM FFI)"""
    return _load_local_fast_compare_key()(a, b)


def benchmark_function(
    func: Callable[[torch.Tensor, torch.Tensor], int],
    a: torch.Tensor,
    b: torch.Tensor,
    warmup: int = 100,
    iterations: int = 10000,
) -> tuple[float, int]:
    """对函数进行性能测试"""
    for _ in range(warmup):
        result = func(a, b)

    start = time.perf_counter()
    for _ in range(iterations):
        result = func(a, b)
    end = time.perf_counter()

    avg_time_us = (end - start) / iterations * 1_000_000
    return avg_time_us, result


def create_test_case(
    length: int, diff_pos: int | None, dtype: torch.dtype = torch.int32
) -> tuple[torch.Tensor, torch.Tensor]:
    a = torch.randint(0, 1000, (length,), dtype=dtype)
    b = a.clone()

    if diff_pos is not None and diff_pos < length:
        b[diff_pos] = (b[diff_pos] + 1) % 1000

    return a, b


def run_benchmark():
    """运行完整的性能对比测试"""
    print("=" * 80)
    print("性能对比测试:C++ vs Python 的 compare_key 实现")
    print("=" * 80)
    print()

    # 测试配置
    test_configs = [
        # (长度, 差异位置, 描述)
        (10, 0, "短数组-开头不匹配"),
        (10, 5, "短数组-中间不匹配"),
        (10, None, "短数组-完全匹配"),
        (100, 0, "中等数组-开头不匹配"),
        (100, 50, "中等数组-中间不匹配"),
        (100, None, "中等数组-完全匹配"),
        (1000, 0, "长数组-开头不匹配"),
        (1000, 500, "长数组-中间不匹配"),
        (1000, 999, "长数组-末尾不匹配"),
        (1000, None, "长数组-完全匹配"),
        (10000, 0, "超长数组-开头不匹配"),
        (10000, 5000, "超长数组-中间不匹配"),
        (10000, None, "超长数组-完全匹配"),
    ]

    # 测试不同数据类型
    dtypes = [
        (torch.int32, "int32"),
        (torch.int64, "int64"),
    ]

    for dtype, dtype_name in dtypes:
        print(f"\n{'=' * 80}")
        print(f"数据类型: {dtype_name}")
        print(f"{'=' * 80}")
        print()

        print(f"{'测试场景':<25} {'Python(μs)':<15} {'C++(μs)':<15} {'加速比':<10} {'结果验证'}")
        print("-" * 80)

        for length, diff_pos, desc in test_configs:
            a, b = create_test_case(length, diff_pos, dtype)
            py_time, py_result = benchmark_function(python_compare_key, a, b)

            cpp_time, cpp_result = benchmark_function(cpp_compare_key, a, b)
            speedup = py_time / cpp_time if cpp_time > 0 else float("inf")
            result_match = "✓" if py_result == cpp_result else "✗"
            print(
                f"{desc:<25} {py_time:>12.3f} {cpp_time:>12.3f}   {speedup:>8.2f}x      {result_match}"
            )


def run_memory_overhead_test():
    """测试内存开销"""
    print("\n\n")
    print("=" * 80)
    print("内存开销测试")
    print("=" * 80)
    print()

    length = 10000
    a, b = create_test_case(length, 5000, torch.int32)

    # 测试 Python 版本
    tracemalloc.start()
    for _ in range(1000):
        python_compare_key(a, b)
    py_current, py_peak = tracemalloc.get_traced_memory()
    tracemalloc.stop()

    # 测试 C++ 版本
    tracemalloc.start()
    for _ in range(1000):
        cpp_compare_key(a, b)
    cpp_current, cpp_peak = tracemalloc.get_traced_memory()
    tracemalloc.stop()

    print(f"Python 版本:")
    print(f"  当前内存: {py_current / 1024:.2f} KB")
    print(f"  峰值内存: {py_peak / 1024:.2f} KB")
    print()
    print(f"C++ 版本:")
    print(f"  当前内存: {cpp_current / 1024:.2f} KB")
    print(f"  峰值内存: {cpp_peak / 1024:.2f} KB")
    print()
    print(f"内存节省: {(py_peak - cpp_peak) / 1024:.2f} KB ({(1 - cpp_peak/py_peak)*100:.1f}%)")


if __name__ == "__main__":
    run_benchmark()
    run_memory_overhead_test()

对比效果:

================================================================================
性能对比测试:C++ vs Python 的 compare_key 实现
================================================================================
测试场景                      Python(μs)      C++(μs)         加速比        结果验证
--------------------------------------------------------------------------------
短数组-开头不匹配                        2.401        0.312       7.70x      ✓
短数组-中间不匹配                        2.923        0.305       9.58x      ✓
短数组-完全匹配                         3.347        0.310      10.81x      ✓
中等数组-开头不匹配                       2.401        0.301       7.98x      ✓
中等数组-中间不匹配                       7.656        0.327      23.39x      ✓
中等数组-完全匹配                       12.688        0.342      37.11x      ✓
================================================================================
内存开销测试
================================================================================
Python 版本:
  当前内存: 0.03 KB
  峰值内存: 0.56 KB
C++ 版本:
  当前内存: 0.03 KB
  峰值内存: 0.12 KB
内存节省: 0.44 KB (79.0%)

可以看到 C++ 版本的性能更好,内存峰值更低。

11. 模型算子

11.1. 极简的算子基类

mini-sglang 通过自研的 BaseOP 类替代 PyTorch 的 nn.Module,实现了模型权重加载与推理计算的解耦。

核心特性:

  • 基于字典树的反射机制,通过点号分隔的字符串键(如 layer.0.weight)直接映射到成员变量
  • 从 Hugging Face 模型文件读取的 key-value 对可无缝加载
  • 提供 state_dict() 和 load_state_dict() 接口,兼容标准权重管理流程

设计优势:
相比传统的 nn.Module 方案,BaseOP 代码更简洁、易于理解,同时实现了权重管理的完全自主可控,便于深入学习和定制化开发。

11.2. 分拆模型计算和权重加载

权重读取采用了集中处理模式,而非像 nano-vllm 一样的内聚处理。这可以将模型的算子计算和权重加载拆分到两处,更有利于未来扩充新的分布式模式或者新模型结构复用旧模型结构的权重加载。同时,多个算子在 TP 模式下,有相似的权重处理逻辑,统一处理后,可以减少重复代码。但这也意味着阅读一个模型的实现需要跨越多个文件,权重加载和权重的使用没有放在一起,显得不内聚。对于大型项目来说,这种拆分是更划算的。nano-vllm 更易于阅读理解、入门学习,而 mini-sglang 更适合于大型项目:更高效的添加新的模型结构,更高效的增加新的并行模式。

TP 模式的权重加载和计算拆分设计:

模块 职责 TP 相关
weight.py 权重加载 + TP 分片 集中处理所有 TP 分片逻辑
base.py 通用加载逻辑 ❌ 不关心 TP,只负责赋值
embedding.py Embedding 特殊逻辑 ❌ 不关心 TP,只处理 tied embedding
linear.py Linear 层实现 ❌ 不关心 TP,只负责前向计算
qwen3.py 模型结构定义 ❌ 不关心 TP,只定义结构

12. 通信插件

mini-sglang 将通信相关的操作封装到 DistributedCommunicator 类中,可以在多个需要分布式通信的地方复用。同时,DistributedCommunicator 也通过一些小技巧让新通信组件可以最小改动的引入。

class DistributedCommunicator:
    plugins: List[DistributedImpl] = [TorchDistributedImpl()]

    def all_reduce(self, x: torch.Tensor) -> torch.Tensor:
        return self.plugins[-1].all_reduce(x)

    def all_gather(self, x: torch.Tensor) -> torch.Tensor:
        return self.plugins[-1].all_gather(x)
      
......

def enable_pynccl_distributed(
    tp_info: DistributedInfo, tp_cpu_group: torch.distributed.ProcessGroup, max_bytes: int
) -> None:
    """
    Enable PyNCCL-based distributed communication for tensor parallelism.
    """
    if tp_info.size == 1:
        return
    from minisgl.kernel import init_pynccl

    comm = init_pynccl(
        tp_rank=tp_info.rank,
        tp_size=tp_info.size,
        tp_cpu_group=tp_cpu_group,
        max_size_bytes=max_bytes,
    )

    DistributedCommunicator.plugins.append(PyNCCLDistributedImpl(comm))

这里的 plugins 是通信插件,定义为 list,当有新的通信插件加入时,只需要把新插件对象 append 到 plugins 里,无需修改其他代码。

13. 进程参数和模型参数

(1)模型结构配置

ModelConfig 类用于定义模型结构的配置信息。

(2)进程参数的继承层次

进程参数采用了三级继承结构:ServerArg → SchedulerConfig → EngineConfig,即 ServerArg(SchedulerConfig(EngineConfig))。

存在的问题:

  • 继承层次过深,可读性差
  • 子类覆盖父类属性的现象时有发生,容易引发错误
  • 虽然设计初衷是区分服务层、调度层、引擎层的参数,但实际效果不佳

改进建议:使用组合模式替代继承,将各层参数作为独立的配置对象,通过组合方式关联,可提升代码的清晰度和可维护性。

(3)模型参数配置的冗余转换

在创建模型对象时,模型参数配置经历了不必要的转换过程:

使用 transformers.AutoConfig.from_pretrained() 加载 config.json 文件,得到 PretrainedConfig 对象

将 PretrainedConfig 对象转换为自定义的 ModelConfig 对象

存在的问题:

  • PretrainedConfig 支持根据 config.json 中的键值对自动扩展属性
  • ModelConfig 不支持自动扩展属性,每当模型引入新的配置参数时,需要手动在 ModelConfig 类中添加对应属性,这种设计增加了维护成本,降低了扩展性,不如直接使用 PretrainedConfig 对象。

14. 参考资料

https://github.com/sgl-project/mini-sglang

https://deepwiki.com/sgl-project/mini-sglang

https://lmsys.org/blog/2025-12-17-minisgl/

本文所在:https://www.cnblogs.com/cswuyg/p/19852266

posted on 2026-04-11 17:08  -银光-  阅读(233)  评论(0)    收藏  举报