模型推理的推理引擎切换:vLLM、TGI 与 TensorRT-LLM 对比

选推理引擎不是选最快的那个,而是选最契合你业务模式的那个。

一、场景痛点

你的团队选了一个开源 LLM 做在线推理服务,第一版用 HuggingFace Transformers 直接跑——延迟 2 秒,用户直接关页面。然后你开始做优化,发现可选方案有 vLLM、TGI(Text Generation Inference)、TensorRT-LLM,每个都说自己是"最快"。

你读了三个项目的 README、跑了十几个 benchmark、看了几十篇技术博客,最后发现最难的并不是技术选型,而是搞清楚你的业务到底是吞吐优先还是延迟优先、是单一模型还是多模型共存、是一次性部署还是需要频繁更新

没有银弹。每个引擎都是为特定场景设计的,选错的代价不是性能差 20%,而是整个推理架构需要推倒重来。

二、底层机制与原理剖析

2.1 三大引擎的核心差异

2.2 PagedAttention:vLLM 的杀手特性

KV Cache 是推理时最占内存的东西。传统实现中,每个请求的 KV Cache 是连续分配的一块显存——这意味着你需要预留 max_seq_len × num_layers × 2 × hidden_dim 的内存,哪怕实际请求只有 50 个 token。

PagedAttention 把 KV Cache 切成固定大小的"页"(类似操作系统中的内存分页),按需分配,不需要预留整块内存。这带来了两个巨大收益:

  1. 显存利用率从 30% 提升到 90%——同样的 GPU 可以服务 3 倍的并发请求
  2. Prefix Sharing:多个请求有共同前缀(如 system prompt),共享同一块 KV Cache 页
传统 KV Cache:   [####################################___________]  浪费 40%
PagedAttention:  [####][####][####][____]                               按需分配
                                    ↑
                                未使用的块可以给其他请求

2.3 In-flight Batching vs Continuous Batching

特性 Continuous Batching (vLLM) In-flight Batching (TRT-LLM)
合并时机 每次 forward 调用前检查和重组 推理进行中动态插入/移除请求
实现方式 Python 层调度,调用 GPU kernel C++ 运行时 + CUDA graph
对新请求的响应 下一个 iteration 加入 当前 iteration 就能插入
适合场景 批量离线推理 实时在线推理

三、生产级代码实现

3.1 推理引擎抽象层

"""
推理引擎抽象层:统一 vLLM、TGI、TensorRT-LLM 的调用接口
支持运行时引擎切换,让业务代码不绑定具体引擎

设计原则:
1. 引擎选择通过环境变量或配置,不硬编码
2. 统一的 generate / generate_stream 接口
3. 引擎特定的参数通过 extra_params 透传
"""
from abc import ABC, abstractmethod
from typing import AsyncIterator, Optional
from dataclasses import dataclass
from enum import Enum
import os

class InferenceEngineType(Enum):
    """推理引擎类型枚举"""
    VLLM = "vllm"
    TGI = "tgi"
    TENSORRT_LLM = "tensorrt_llm"
    HUGGINGFACE = "huggingface"

@dataclass
class GenerationConfig:
    """生成配置,与引擎无关"""
    max_tokens: int = 2048
    temperature: float = 0.7
    top_p: float = 0.95
    top_k: int = 50
    repetition_penalty: float = 1.1
    stop_sequences: Optional[list[str]] = None
    # 引擎特有参数透传
    extra_params: Optional[dict] = None

@dataclass
class GenerationResult:
    """生成结果"""
    text: str
    input_tokens: int
    output_tokens: int
    finish_reason: str            # "stop" / "length" / "error"
    latency_ms: float             # 首 token 延迟
    total_latency_ms: float       # 总延迟
    tokens_per_second: float      # 吞吐

class InferenceEngine(ABC):
    """
    推理引擎抽象基类
    所有引擎实现必须实现 generate 和 generate_stream 两个方法
    """
    
    @abstractmethod
    async def generate(self, prompt: str, config: GenerationConfig) -> GenerationResult:
        """同步生成(等待完整结果)"""
        pass
    
    @abstractmethod
    async def generate_stream(
        self, prompt: str, config: GenerationConfig
    ) -> AsyncIterator[str]:
        """流式生成(逐 token 输出)"""
        pass
    
    @abstractmethod
    async def health_check(self) -> bool:
        """健康检查"""
        pass
    
    @abstractmethod
    async def get_metrics(self) -> dict:
        """获取引擎指标"""
        pass

# ========== vLLM 引擎实现 ==========
class VLLMEngine(InferenceEngine):
    """
    vLLM 引擎封装
    使用 OpenAI 兼容 API 调用
    
    优势: PagedAttention 高吞吐、_prefix caching、continuous batching
    劣势: Python 运行时开销、冷启动慢
    """
    
    def __init__(self, base_url: str, model_name: str):
        self.base_url = base_url.rstrip('/')
        self.model_name = model_name
        self.client = None  # openai.AsyncOpenAI 实例
    
    async def generate(self, prompt: str, config: GenerationConfig) -> GenerationResult:
        import time
        start = time.perf_counter()
        
        # vLLM 的 OpenAI 兼容 API
        response = await self.client.completions.create(
            model=self.model_name,
            prompt=prompt,
            max_tokens=config.max_tokens,
            temperature=config.temperature,
            top_p=config.top_p,
            extra_body={
                "top_k": config.top_k,
                "repetition_penalty": config.repetition_penalty,
                "stop": config.stop_sequences,
                **(config.extra_params or {})
            }
        )
        
        elapsed = time.perf_counter() - start
        choice = response.choices[0]
        
        return GenerationResult(
            text=choice.text,
            input_tokens=response.usage.prompt_tokens,
            output_tokens=response.usage.completion_tokens,
            finish_reason=choice.finish_reason or "stop",
            latency_ms=elapsed * 1000,  # vLLM 暂不区分首 token
            total_latency_ms=elapsed * 1000,
            tokens_per_second=response.usage.completion_tokens / elapsed
        )
    
    async def generate_stream(
        self, prompt: str, config: GenerationConfig
    ) -> AsyncIterator[str]:
        stream = await self.client.completions.create(
            model=self.model_name,
            prompt=prompt,
            max_tokens=config.max_tokens,
            temperature=config.temperature,
            top_p=config.top_p,
            stream=True,
            extra_body={
                "top_k": config.top_k,
                "repetition_penalty": config.repetition_penalty,
                "stop": config.stop_sequences,
                **(config.extra_params or {})
            }
        )
        
        async for chunk in stream:
            if chunk.choices and chunk.choices[0].text:
                yield chunk.choices[0].text
    
    async def health_check(self) -> bool:
        try:
            resp = await self.client.models.list()
            return True
        except Exception:
            return False
    
    async def get_metrics(self) -> dict:
        """vLLM 指标通过 /metrics 端点暴露"""
        import aiohttp
        async with aiohttp.ClientSession() as session:
            async with session.get(f"{self.base_url}/metrics") as resp:
                raw_metrics = await resp.text()
                # 解析 Prometheus 格式指标
                return {"raw": raw_metrics}

# ========== TensorRT-LLM 引擎实现 ==========
class TensorRTLLMEngine(InferenceEngine):
    """
    TensorRT-LLM 引擎封装
    通过 Triton Inference Server 的 gRPC 接口调用
    
    优势: GPU kernel 级优化、FP8/INT4 硬件加速、C++ 运行时低延迟
    劣势: 模型转换复杂、不支持所有模型架构、更新慢
    """
    
    def __init__(self, triton_url: str, model_name: str):
        self.triton_url = triton_url
        self.model_name = model_name
        # 使用 tritonclient.grpc 或 HTTP
    
    async def generate(self, prompt: str, config: GenerationConfig) -> GenerationResult:
        import time
        start = time.perf_counter()
        
        # TensorRT-LLM 通过 Triton 暴露,协议为 gRPC 或 HTTP
        # 这里使用简化的 HTTP 调用示意
        import aiohttp
        
        payload = {
            "text_input": prompt,
            "max_tokens": config.max_tokens,
            "temperature": config.temperature,
            "top_p": config.top_p,
            "top_k": config.top_k,
            "stop_words": config.stop_sequences or [],
            "bad_words": [],
            "stream": False
        }
        
        async with aiohttp.ClientSession() as session:
            async with session.post(
                f"{self.triton_url}/v2/models/{self.model_name}/generate",
                json=payload
            ) as resp:
                data = await resp.json()
        
        elapsed = time.perf_counter() - start
        
        return GenerationResult(
            text=data.get("text_output", ""),
            input_tokens=data.get("input_tokens", 0),
            output_tokens=data.get("output_tokens", 0),
            finish_reason="stop",
            latency_ms=elapsed * 1000,
            total_latency_ms=elapsed * 1000,
            tokens_per_second=data.get("output_tokens", 0) / elapsed
        )
    
    async def generate_stream(
        self, prompt: str, config: GenerationConfig
    ) -> AsyncIterator[str]:
        import aiohttp
        
        payload = {
            "text_input": prompt,
            "max_tokens": config.max_tokens,
            "temperature": config.temperature,
            "top_p": config.top_p,
            "top_k": config.top_k,
            "stop_words": config.stop_sequences or [],
            "stream": True
        }
        
        async with aiohttp.ClientSession() as session:
            async with session.post(
                f"{self.triton_url}/v2/models/{self.model_name}/generate_stream",
                json=payload
            ) as resp:
                async for line in resp.content:
                    if line:
                        text = line.decode('utf-8').strip()
                        if text.startswith("data: "):
                            yield text[6:]
    
    async def health_check(self) -> bool:
        import aiohttp
        try:
            async with aiohttp.ClientSession() as session:
                async with session.get(f"{self.triton_url}/v2/health/ready") as resp:
                    return resp.status == 200
        except Exception:
            return False
    
    async def get_metrics(self) -> dict:
        import aiohttp
        async with aiohttp.ClientSession() as session:
            async with session.get(f"{self.triton_url}/v2/models/{self.model_name}/stats") as resp:
                return await resp.json()

# ========== 引擎工厂 ==========
class InferenceEngineFactory:
    """
    推理引擎工厂:根据配置创建对应的引擎实例
    
    使用方式:
    engine = InferenceEngineFactory.create()
    result = await engine.generate("Hello", GenerationConfig())
    """
    
    _registry = {
        InferenceEngineType.VLLM: VLLMEngine,
        InferenceEngineType.TENSORRT_LLM: TensorRTLLMEngine,
    }
    
    @classmethod
    def create(cls, engine_type: Optional[InferenceEngineType] = None) -> InferenceEngine:
        """
        创建推理引擎实例
        优先级: 参数指定 > 环境变量 > 默认 vLLM
        """
        if engine_type is None:
            engine_type = InferenceEngineType(
                os.getenv("INFERENCE_ENGINE", "vllm")
            )
        
        engine_cls = cls._registry.get(engine_type)
        if engine_cls is None:
            raise ValueError(f"不支持的引擎类型: {engine_type}")
        
        # 从环境变量读取引擎连接信息
        base_url = os.getenv(
            f"{engine_type.value.upper()}_URL",
            "http://localhost:8000"
        )
        model_name = os.getenv("MODEL_NAME", "default")
        
        return engine_cls(base_url, model_name)

# ========== 使用示例 ==========
async def main():
    engine = InferenceEngineFactory.create(InferenceEngineType.VLLM)
    
    # 流式生成
    config = GenerationConfig(max_tokens=512, temperature=0.7)
    async for token in engine.generate_stream("解释量子纠缠", config):
        print(token, end="", flush=True)

3.2 引擎切换的决策逻辑

"""
推理引擎选择策略:根据请求特征自动路由到最合适的引擎

路由规则:
- 实时对话 (streaming, < 100 tokens/s) → TensorRT-LLM (低延迟)
- 批量生成 (batch, > 1000 tokens) → vLLM (高吞吐)
- 离线批处理 → vLLM (PagedAttention 高并发)
- 非主流模型架构 → vLLM (兼容性好)
"""
from dataclasses import dataclass
from enum import Enum

class RequestType(Enum):
    CHAT = "chat"        # 实时对话(流式)
    BATCH = "batch"      # 批量生成
    OFFLINE = "offline"  # 离线批处理

@dataclass
class RoutingRequest:
    prompt: str
    request_type: RequestType
    max_tokens: int
    stream: bool = False

def route_engine(request: RoutingRequest) -> InferenceEngineType:
    """
    根据请求特征选择推理引擎
    """
    # 规则 1: 实时对话 → TensorRT-LLM(低延迟,in-flight batching)
    if request.request_type == RequestType.CHAT and request.stream:
        return InferenceEngineType.TENSORRT_LLM
    
    # 规则 2: 大批量生成 → vLLM(PagedAttention 高并发 + prefix caching)
    if request.request_type in (RequestType.BATCH, RequestType.OFFLINE):
        return InferenceEngineType.VLLM
    
    # 默认: vLLM(兼容性最好)
    return InferenceEngineType.VLLM

async def execute_with_engine(request: RoutingRequest) -> str:
    """自动路由到最合适的引擎执行推理"""
    engine_type = route_engine(request)
    engine = InferenceEngineFactory.create(engine_type)
    
    config = GenerationConfig(max_tokens=request.max_tokens)
    
    if request.stream:
        result_parts = []
        async for token in engine.generate_stream(request.prompt, config):
            result_parts.append(token)
        return "".join(result_parts)
    else:
        result = await engine.generate(request.prompt, config)
        return result.text

四、边界分析与架构权衡

4.1 性能不是唯一维度

维度 vLLM TGI TensorRT-LLM
单请求延迟
并发吞吐
模型兼容性 高(支持 50+ 架构) 低(需 TRT 引擎)
量化支持 GPTQ/AWQ/SqueezeLLM GPTQ/AWQ/BitsAndBytes FP8/INT4 (硬件加速)
部署复杂度 低 (pip install) 低 (docker) 高 (模型转换 + Triton)
更新频率 极高(周更) 中(月更) 低(季更)
社区生态 Python 生态完整 HuggingFace 官方 NVIDIA 企业支持

场景建议

  • 快速验证阶段 → vLLM。pip install 起一个服务,5 分钟跑通
  • 生产在线服务 → TensorRT-LLM。花一周做模型转换和 kernel 调优,换取 2-3 倍的推理性能
  • 大量模型混用 → vLLM。你不可能把 50 个模型都转成 TRT 引擎
  • HuggingFace 重度用户 → TGI。和 HF Hub 的集成最紧密

4.2 引擎迁移的隐性成本

从 vLLM 切到 TensorRT-LLM 不只是改几行代码:

  • 模型转换时间:70B 模型转 TRT 引擎大约需要 2-4 小时(A100)
  • TRT 引擎文件体积:通常是原始权重的 1.5-2 倍(因为包含了优化后的 kernel)
  • 不支持的算子:某些自定义 attention 或 activation 函数需要手写 TRT plugin
  • 运维复杂度:Triton Server 的配置比 vLLM 的 OpenAI API 复杂一个数量级

4.3 混合部署策略

我们最终采用了双引擎并存的架构:

  • TensorRT-LLM 跑流量最大的核心模型(3 个模型),部署在 8×A100 节点上
  • vLLM 跑长尾模型和实验模型(20+ 个),部署在 4×A100 节点上
  • 前端统一网关根据 model_name 路由到不同引擎

这样既享受了 TRT 的性能优势,又保留了 vLLM 的灵活性和生态兼容性。

五、总结

选推理引擎的三个原则:

  1. 先看业务需求,再看 benchmark。延迟敏感型选 TRT-LLM,吞吐敏感型选 vLLM,和 HuggingFace 生态深度绑定选 TGI。
  2. 不要追求"一个引擎统治所有场景"。双引擎甚至三引擎并存是合理的,通过统一抽象层屏蔽差异。
  3. 考虑全生命周期成本,不只是推理延迟。TRT-LLM 的 2ms 延迟优势,可能被几小时的模型转换时间、运维复杂度、和有限的模型兼容性所抵消。

最后,这三个引擎的竞争正在让整个生态快速进化——vLLM 在抄 TRT 的 in-flight batching,TRT 在兼容更多模型,TGI 在优化吞吐。一年后的格局可能完全不同。所以把引擎选择做成可替换的,比选对更重要

Logo

免费领 150 小时云算力,进群参与显卡、AI PC 幸运抽奖

更多推荐