前言

💡 痛点: 部署大模型推理太慢?KV Cache 爆炸?并发上不去?量化后精度掉太多?多卡推理效率低下?

🎯 解决方案: 从 vLLM(PagedAttention)到 TensorRT-LLM(Kernel 融合),再到 SGLang(激进调度),构建生产级 LLM 推理服务全链路。

高级优化

量化

KV Cache 管理

推理引擎

vLLM
PagedAttention
Continuous Batching

TensorRT-LLM
Kernel Fusion
FP8/INT8

SGLang
RadixAttention
激进调度

LightLLM
分层调度
Token 级别

PagedAttention
页式管理

RadixAttention
共享前缀

Prefix Caching
系统提示复用

AWQ
激活感知

GPTQ
后训练量化

FP8
H100 专用

Smooth Quant
通道平衡

Speculative Decoding
投机采样

Continuous Batching
连续批处理

Chunked Prefill
分块预填充

Data Parallel
数据并行

Tensor Parallel
张量并行

Pipeline Parallel
流水线并行

推理引擎选型矩阵:

引擎 吞吐量 延迟 易用性 量化支持 适用场景
vLLM ★★★★★ ★★★★ ★★★★★ AWQ/GPTQ 通用部署
TensorRT-LLM ★★★★★ ★★★★★ ★★★ FP8/INT8/INT4 NVIDIA GPU 极致性能
SGLang ★★★★ ★★★★ ★★★★ AWQ/GPTQ RadixAttention 场景
LightLLM ★★★★ ★★★★ ★★★ AWQ/GPTQ 大规模部署
llama.cpp ★★★ ★★★ ★★★★★ GGUF 多量化 CPU/边缘设备
MLC-LLM ★★★★ ★★★★ ★★★★ 编译优化 跨平台部署

一、vLLM 深度实战

1.1 安装

# 方法 1: pip 安装(推荐)
pip install vllm

# 方法 2: 从源码编译(最新特性)
git clone https://github.com/vllm-project/vllm.git
cd vllm
pip install -e .

# 验证安装
python -c "import vllm; print(vllm.__version__)"

# 检查 CUDA 可用性
python -c "import torch; print(torch.cuda.is_available())"

1.2 基础推理

# vllm_basic.py

from vllm import LLM, SamplingParams

# ===== 初始化 LLM =====
llm = LLM(
    model="meta-llama/Llama-3.1-8B-Instruct",
    tensor_parallel_size=2,            # 2 卡张量并行
    dtype="bfloat16",                   # 计算精度
    gpu_memory_utilization=0.9,         # GPU 显存利用率
    max_model_len=8192,                 # 最大序列长度
    trust_remote_code=True,
    # 高级选项
    enable_prefix_caching=True,          # 启用 Prefix Caching
    block_size=16,                       # PagedAttention block 大小
    swap_space=4,                        # CPU-GPU swap 空间(GB)
    cpu_offload_gb=0,                    # CPU offload 大小
    # 分布式
    # pipeline_parallel_size=2,          # 流水线并行
    # data_parallel_size=2,              # 数据并行
)

# ===== 采样参数 =====
sampling_params = SamplingParams(
    temperature=0.7,
    top_p=0.95,
    top_k=50,
    max_tokens=1024,
    stop=["<|im_end|>"],               # 停止词
    repetition_penalty=1.1,
    frequency_penalty=0.0,
    presence_penalty=0.0,
)

# ===== 批量推理 =====
prompts = [
    "解释一下量子计算的基本原理",
    "用 Python 写一个快速排序算法",
    "什么是 Transformer 架构?",
]

outputs = llm.generate(prompts, sampling_params)

# 输出结果
for output in outputs:
    prompt = output.prompt
    generated_text = output.outputs[0].text
    print(f"Prompt: {prompt}")
    print(f"Generated: {generated_text}")
    print(f"Tokens: {len(output.outputs[0].token_ids)}")
    print("-" * 80)

1.3 OpenAI 兼容 API 服务

# vllm_server.py

from vllm import LLM, SamplingParams
from vllm.entrypoints.openai.api_server import start_server

# 启动 OpenAI 兼容服务
if __name__ == "__main__":
    import argparse
    
    parser = argparse.ArgumentParser()
    parser.add_argument("--model", type=str, default="meta-llama/Llama-3.1-8B-Instruct")
    parser.add_argument("--tensor-parallel-size", type=int, default=2)
    parser.add_argument("--dtype", type=str, default="bfloat16")
    parser.add_argument("--gpu-memory-utilization", type=float, default=0.9)
    parser.add_argument("--max-model-len", type=int, default=8192)
    parser.add_argument("--enable-prefix-caching", action="store_true")
    parser.add_argument("--port", type=int, default=8000)
    parser.add_argument("--host", type=str, default="0.0.0.0")
    args = parser.parse_args()
    
    # 启动服务
    start_server(
        model=args.model,
        tensor_parallel_size=args.tensor_parallel_size,
        dtype=args.dtype,
        gpu_memory_utilization=args.gpu_memory_utilization,
        max_model_len=args.max_model_len,
        enable_prefix_caching=args.enable_prefix_caching,
        port=args.port,
        host=args.host,
    )
# 启动命令
python vllm_server.py \
  --model meta-llama/Llama-3.1-8B-Instruct \
  --tensor-parallel-size 2 \
  --dtype bfloat16 \
  --gpu-memory-utilization 0.9 \
  --max-model-len 8192 \
  --enable-prefix-caching \
  --port 8000

# 测试 API
curl http://localhost:8000/v1/completions \
  -H "Content-Type: application/json" \
  -d '{
    "model": "meta-llama/Llama-3.1-8B-Instruct",
    "prompt": "Hello, how are you?",
    "max_tokens": 100
  }'

1.4 PagedAttention 原理

# paged_attention_explained.py

"""
PagedAttention 核心思想:
将 KV Cache 分成固定大小的「块」(block),类似操作系统的虚拟内存管理。

优势:
1. 消除内存碎片
2. 支持动态批处理(请求可以随时加入/离开)
3. 共享前缀(不同请求可共享相同的 block)

vLLM 实现细节:
- Block 大小:16 tokens(可配置)
- 逻辑 block ID → 物理 block ID(通过映射表)
- 当物理 block 不足时,将最近最少使用的 block swap 到 CPU

内存节省对比:
- 传统 KV Cache:每个请求预分配 max_seq_len 的显存
- PagedAttention:按需分配,节省 50-90% 显存
"""

# 简化版 PagedAttention 实现
import torch
import torch.nn.functional as F

class PagedAttention:
    """简化版 PagedAttention 实现"""
    
    def __init__(self, num_heads: int, head_dim: int, block_size: int = 16):
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.block_size = block_size
        
        # 物理 KV Cache(全局共享)
        self.physical_cache_k = None   # [num_blocks, block_size, num_heads, head_dim]
        self.physical_cache_v = None
        
        # 逻辑 block → 物理 block 映射(每个请求一个)
        self.logical_to_physical = {}  # request_id → List[physical_block_id]
    
    def allocate_physical_cache(self, num_blocks: int, device: torch.device):
        """分配物理 KV Cache"""
        self.physical_cache_k = torch.zeros(
            num_blocks, self.block_size, self.num_heads, self.head_dim,
            device=device
        )
        self.physical_cache_v = torch.zeros_like(self.physical_cache_k)
    
    def compute_attention(
        self,
        query: torch.Tensor,           # [batch, num_heads, head_dim]
        request_ids: list[str],
        logical_lengths: list[int],    # 每个请求的逻辑长度
    ) -> torch.Tensor:
        """计算 PagedAttention"""
        
        outputs = []
        
        for i, req_id in enumerate(request_ids):
            # 获取该请求的物理 block 列表
            physical_blocks = self.logical_to_physical[req_id]
            logical_len = logical_lengths[i]
            
            # 收集该请求的 KV(跨越多个物理 block)
            kv_k = []
            kv_v = []
            
            for block_id in physical_blocks:
                start = 0
                end = min(self.block_size, logical_len)
                kv_k.append(self.physical_cache_k[block_id, start:end])
                kv_v.append(self.physical_cache_v[block_id, start:end])
                logical_len -= (end - start)
            
            # 拼接 KV
            k = torch.cat(kv_k, dim=0)  # [seq_len, num_heads, head_dim]
            v = torch.cat(kv_v, dim=0)
            
            # 标准 Attention 计算
            scores = torch.einsum("hd,shd->hs", query[i], k)
            scores = scores / (self.head_dim ** 0.5)
            attn_weights = F.softmax(scores, dim=-1)
            output = torch.einsum("hs,shd->hd", attn_weights, v)
            
            outputs.append(output)
        
        return torch.stack(outputs, dim=0)

二、TensorRT-LLM 深度优化

2.1 安装

# 方法 1: pip 安装(推荐)
pip install tensorrt-llm

# 方法 2: 从源码编译(最新特性)
git clone https://github.com/NVIDIA/TensorRT-LLM.git
cd TensorRT-LLM
pip install -r requirements.txt
python setup.py install

# 安装 TensorRT(必需)
# 从 NVIDIA 官网下载 TensorRT 10.x
# 或:pip install nvidia-tensorrt

2.2 模型编译与量化

# trtllm_compile.py

import tensorrt_llm as trtllm
from tensorrt_llm.quantization import QuantMode
from tensorrt_llm.models import LLaMAForCausalLM

# ===== 1. 加载模型 =====
model = LLaMAForCausalLM.from_huggingface(
    "meta-llama/Llama-3.1-8B-Instruct",
    dtype="float16",
)

# ===== 2. 量化配置 =====
# 选项 1: INT8 Smooth Quant
quant_mode = QuantMode.from_description(
    quantize_weights=True,
    quantize_activations=True,
    per_token=True,
    per_channel=False,
    use_int8=True,
    use_int4=False,
)

# 选项 2: INT4 AWQ(更激进)
# quant_mode = QuantMode.from_description(
#     quantize_weights=True,
#     quantize_activations=False,
#     per_token=False,
#     per_channel=True,
#     use_int8=False,
#     use_int4=True,
# )

# 选项 3: FP8(H100/A100-80G 推荐)
# quant_mode = QuantMode.from_description(
#     quantize_weights=True,
#     quantize_activations=True,
#     per_token=True,
#     per_channel=False,
#     use_fp8=True,
# )

# ===== 3. 编译模型 =====
engine = model.to_trt(
    engine_name="llama-3.1-8b-trt",
    quantization=quant_mode,
    max_batch_size=128,
    max_input_len=2048,
    max_output_len=1024,
    max_beam_width=1,
    use_refit=False,
    use_inflight_batching=True,   # 连续批处理
    use_paged_kv_cache=True,       # Paged KV Cache
    paged_kv_cache_max_token_num=4096,
)

# ===== 4. 保存引擎 =====
engine.save("./trt_engines/llama-3.1-8b-int8")

print("✅ TensorRT-LLM 引擎编译完成")

2.3 推理服务

# trtllm_server.py

import tensorrt_llm
from tensorrt_llm.runtime import ModelRunner

class TRTLLMServer:
    def __init__(self, engine_dir: str, tokenizer_dir: str):
        self.runner = ModelRunner.from_dir(
            engine_dir=engine_dir,
            tokenizer_dir=tokenizer_dir,
        )
    
    def generate(
        self,
        prompts: list[str],
        max_new_tokens: int = 1024,
        temperature: float = 0.7,
        top_p: float = 0.95,
    ) -> list[str]:
        """批量生成"""
        
        # 构建采样参数
        sampling_config = tensorrt_llm.SamplingConfig(
            temperature=temperature,
            top_p=top_p,
            top_k=50,
            max_new_tokens=max_new_tokens,
            end_id=self.runner.tokenizer.eos_token_id,
        )
        
        # 推理
        outputs = self.runner.generate(
            prompts,
            sampling_config=sampling_config,
            streaming=False,
        )
        
        # 解码
        results = []
        for output in outputs:
            text = self.runner.tokenizer.decode(output.outputs[0].token_ids)
            results.append(text)
        
        return results

# 启动服务
if __name__ == "__main__":
    server = TRTLLMServer(
        engine_dir="./trt_engines/llama-3.1-8b-int8",
        tokenizer_dir="meta-llama/Llama-3.1-8B-Instruct",
    )
    
    prompts = [
        "解释一下量子计算的基本原理",
        "用 Python 写一个快速排序算法",
    ]
    
    results = server.generate(prompts)
    for prompt, result in zip(prompts, results):
        print(f"Prompt: {prompt}")
        print(f"Result: {result}\n")

三、SGLang 激进调度

3.1 安装与使用

# 安装
pip install "sglang[all]"

# 启动服务(命令行)
python -m sglang.launch_server \
  --model-path meta-llama/Llama-3.1-8B-Instruct \
  --port 30000 \
  --tp 2 \
  --enable-radix-cache \
  --disable-cuda-graph \
  --schedule-policy lpm \
  --max-total-tokens 4096

3.2 RadixAttention 原理

# sglang_radix.py

"""
RadixAttention:SGLang 的核心创新

问题:许多请求共享相同的前缀(如系统提示、few-shot 示例)
传统方法:每个请求独立存储 KV Cache,浪费大量显存

RadixAttention 方案:
- 用 Radix Tree 组织 KV Cache
- 相同前缀的请求共享 KV Cache 节点
- 命中时直接复用,无需重新计算

效果:
- 系统提示复用:节省 90%+ 显存
- 多轮对话:历史消息复用
- 代码生成:共享上下文复用
"""

import torch
from dataclasses import dataclass
from typing import Dict, Optional

@dataclass
class RadixNode:
    """Radix Tree 节点"""
    token_ids: list[int]           # 该节点存储的 token IDs
    parent: Optional['RadixNode']
    children: Dict[str, 'RadixNode']  # 子节点(key: 下一个 token 的 hash)
    kv_cache_ptr: Optional[torch.Tensor]  # KV Cache 指针
    ref_count: int = 0             # 引用计数

class RadixAttentionManager:
    """RadixAttention KV Cache 管理器"""
    
    def __init__(self, max_cache_size: int = 10000):
        self.root = RadixNode(token_ids=[], parent=None, children={}, kv_cache_ptr=None)
        self.max_cache_size = max_cache_size
        self.lru_cache = []        # LRU 淘汰队列
    
    def find_longest_prefix(
        self, 
        token_ids: list[int]
    ) -> tuple[RadixNode, int]:
        """查找最长匹配前缀"""
        node = self.root
        matched_len = 0
        
        for i, token_id in enumerate(token_ids):
            key = str(token_id)
            if key in node.children:
                child = node.children[key]
                # 检查该 child 的 token_ids 是否匹配
                if (i + len(child.token_ids) <= len(token_ids) and
                    token_ids[i:i+len(child.token_ids)] == child.token_ids):
                    node = child
                    matched_len += len(child.token_ids)
                else:
                    break
            else:
                    break
        
        return node, matched_len
    
    def insert(
        self, 
        token_ids: list[int], 
        kv_cache: torch.Tensor
    ) -> RadixNode:
        """插入新的 KV Cache"""
        node, prefix_len = self.find_longest_prefix(token_ids)
        
        if prefix_len == len(token_ids):
            # 完全命中
            node.ref_count += 1
            return node
        
        # 需要创建新节点
        remaining = token_ids[prefix_len:]
        
        # 分裂节点(如果有部分匹配)
        if prefix_len > 0 and node != self.root:
            # ... 实现 Radix Tree 分裂逻辑
            pass
        
        new_node = RadixNode(
            token_ids=remaining,
            parent=node,
            children={},
            kv_cache_ptr=kv_cache,
            ref_count=1,
        )
        
        # 添加到父节点的 children
        if remaining:
            node.children[str(remaining[0])] = new_node
        
        # LRU 更新
        self.lru_cache.append(new_node)
        self._evict_if_needed()
        
        return new_node
    
    def _evict_if_needed(self):
        """LRU 淘汰"""
        while len(self.lru_cache) > self.max_cache_size:
            victim = self.lru_cache.pop(0)
            if victim.ref_count == 0:
                # 释放 KV Cache
                victim.kv_cache_ptr = None
                # 从父节点移除
                if victim.parent:
                    key = str(victim.token_ids[0])
                    victim.parent.children.pop(key, None)

3.3 SGLang Python API

# sglang_api.py

import sglang as sgl

# ===== 定义生成函数 =====
@sgl.function
def code_generation(s: sgl.SglState, question: str):
    s += sgl.system("你是一个专业的 Python 开发者。")
    s += sgl.user(question)
    s += sgl.assistant(sgl.gen("answer", max_tokens=1024, temperature=0.7))

@sgl.function
def multi_turn_chat(s: sgl.SglState, history: list[dict]):
    for msg in history:
        if msg["role"] == "user":
            s += sgl.user(msg["content"])
        else:
            s += sgl.assistant(msg["content"])
    
    s += sgl.assistant(sgl.gen("response", max_tokens=512))

# ===== 批量执行 =====
if __name__ == "__main__":
    # 初始化后端
    sgl.set_default_backend(sgl.RuntimeEndpoint("http://localhost:30000"))
    
    # 问题列表
    questions = [
        "用 Python 实现快速排序",
        "解释一下装饰器模式",
        "如何优化 SQL 查询性能?",
    ]
    
    # 批量执行(自动利用 RadixAttention)
    for question in questions:
        state = sgl.SglState()
        code_generation(state, question=question)
        print(f"Q: {question}")
        print(f"A: {state['answer']}\n")

四、量化技术详解

4.1 AWQ(Activation-aware Weight Quantization)

# awq_quantize.py

from awq import AutoAWQForCausalLM
from transformers import AutoTokenizer

def quantize_with_awq(
    model_path: str,
    output_path: str,
    quant_config: dict = None,
):
    """使用 AWQ 量化模型"""
    
    if quant_config is None:
        quant_config = {
            "zero_point": True,       # 是否使用零点量化
            "q_group_size": 128,      # 量化组大小
            "w_bit": 4,               # 权重量化位数
            "version": "GEMM",        # 实现版本(GEMM / GEMV)
        }
    
    # 加载模型
    model = AutoAWQForCausalLM.from_pretrained(
        model_path,
        trust_remote_code=True,
    )
    tokenizer = AutoTokenizer.from_pretrained(
        model_path,
        trust_remote_code=True,
    )
    
    # 量化
    model.quantize(
        tokenizer,
        quant_config=quant_config,
        calib_data="pile",            # 校准数据集
        num_calib_data=512,           # 校准样本数
        max_calib_samples=512,
        max_seq_len=512,
    )
    
    # 保存量化模型
    model.save_quantized(output_path)
    tokenizer.save_pretrained(output_path)
    
    print(f"✅ AWQ 量化完成: {output_path}")
    print(f"   量化配置: w_bit={quant_config['w_bit']}, group_size={quant_config['q_group_size']}")

# 使用量化模型
def run_awq_model(model_path: str):
    """运行 AWQ 量化模型"""
    from awq import AutoAWQForCausalLM
    from transformers import AutoTokenizer
    import torch
    
    model = AutoAWQForCausalLM.from_quantized(
        model_path,
        device_map="auto",
        trust_remote_code=True,
    )
    tokenizer = AutoTokenizer.from_pretrained(model_path)
    
    prompt = "解释一下量子计算"
    inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
    
    with torch.no_grad():
        outputs = model.generate(
            **inputs,
            max_new_tokens=512,
            do_sample=True,
            temperature=0.7,
        )
    
    print(tokenizer.decode(outputs[0], skip_special_tokens=True))

4.2 GPTQ(后训练量化)

# gptq_quantize.py

from auto_gptq import AutoGPTQForCausalLM, BaseQuantizeConfig
from transformers import AutoTokenizer

def quantize_with_gptq(
    model_path: str,
    output_path: str,
    bits: int = 4,
    group_size: int = 128,
):
    """使用 GPTQ 量化"""
    
    quantize_config = BaseQuantizeConfig(
        bits=bits,                    # 量化位数
        group_size=group_size,        # 量化组大小
        desc_act=False,               # 是否量化激活值
        damp_percent=0.01,            # 阻尼系数
    )
    
    # 加载模型
    model = AutoGPTQForCausalLM.from_pretrained(
        model_path,
        quantize_config=quantize_config,
        trust_remote_code=True,
    )
    tokenizer = AutoTokenizer.from_pretrained(
        model_path,
        trust_remote_code=True,
    )
    
    # 量化(需要校准数据)
    model.quantize(tokenizer, calib_data="pile", num_calib_data=1000)
    
    # 保存
    model.save_quantized(output_path)
    tokenizer.save_pretrained(output_path)
    
    print(f"✅ GPTQ 量化完成: {output_path}")

# GPTQ vs AWQ 对比
"""
| 特性 | AWQ | GPTQ |
|------|-----|------|
| 原理 | 激活感知权重量化 | 最优脑量化(OBQ) |
| 速度 | 更快(GEMM 优化)| 稍慢 |
| 精度 | 略好 | 基准 |
| 显存 | 更低 | 相近 |
| 推荐 | ✅ 生产环境 | 研究/实验 |
"""

4.3 FP8 量化(H100 专用)

# fp8_quantize.py

"""
FP8 量化:
- 格式:E4M3(4 位指数 + 3 位尾数)或 E5M2
- 优势:硬件加速(H100 Transformer Engine)
- 精度损失:< 1%(几乎无损)
- 速度提升:2-3x(相比 FP16)

适用场景:
- H100 / H200 / L40S / Ada Lovelace(RTX 40 系列)
- 对精度要求高的场景
"""

import torch
from transformers import AutoModelForCausalLM

def convert_to_fp8(model_path: str, output_path: str):
    """转换为 FP8 格式(需 TensorRT-LLM 或 vLLM)"""
    
    # 加载模型
    model = AutoModelForCausalLM.from_pretrained(
        model_path,
        torch_dtype=torch.bfloat16,
        device_map="auto",
    )
    
    # 转换 Linear 层为 FP8
    for name, module in model.named_modules():
        if isinstance(module, torch.nn.Linear):
            # 动态量化(推理时)
            weight_fp8 = module.weight.to(torch.float8_e4m3fn)
            module.weight.data = weight_fp8
            
            # 保存缩放因子
            scale = module.weight.abs().max() / 448  # FP8 范围
            module.register_buffer("weight_scale", scale)
    
    # 保存
    model.save_pretrained(output_path)
    print(f"✅ FP8 转换完成: {output_path}")

五、投机采样(Speculative Decoding)

# speculative_decoding.py

"""
投机采样:
问题:自回归解码每次只能生成一个 token,GPU 利用率低

方案:用小模型(草稿模型)快速生成 K 个 token,
     用大模型(验证模型)并行验证这 K 个 token,
     接受所有正确的 token,拒绝第一个错误的 token 及其后所有 token

效果:
- 加速比:2-3x(取决于草稿模型质量)
- 无损:输出分布与原始大模型完全相同
- 适用:代码生成、翻译等草稿模型质量高的场景
"""

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

class SpeculativeDecoder:
    """投机采样解码器"""
    
    def __init__(
        self,
        draft_model_name: str,
        target_model_name: str,
        num_speculative_tokens: int = 5,
    ):
        self.draft_model = AutoModelForCausalLM.from_pretrained(draft_model_name)
        self.target_model = AutoModelForCausalLM.from_pretrained(target_model_name)
        self.tokenizer = AutoTokenizer.from_pretrained(target_model_name)
        self.num_speculative_tokens = num_speculative_tokens
    
    def generate(
        self,
        prompt: str,
        max_new_tokens: int = 1024,
        temperature: float = 1.0,
    ) -> str:
        """投机采样生成"""
        
        input_ids = self.tokenizer.encode(prompt, return_tensors="pt")
        
        for _ in range(max_new_tokens):
            # ===== 1. 草稿模型生成 K 个 token =====
            draft_tokens = self._generate_draft(input_ids, self.num_speculative_tokens)
            
            # ===== 2. 目标模型并行验证 =====
            # 将 input_ids + draft_tokens 一起输入目标模型
            full_seq = torch.cat([input_ids, draft_tokens], dim=-1)
            with torch.no_grad():
                outputs = self.target_model(full_seq)
                target_logits = outputs.logits[:, -len(draft_tokens[0]):]  # 只取草稿部分
            
            # ===== 3. 接受/拒绝 =====
            accepted_tokens = []
            for i, draft_token in enumerate(draft_tokens[0].tolist()):
                # 计算接受概率
                target_prob = torch.softmax(target_logits[0, i] / temperature, dim=-1)
                draft_prob = ...  # 草稿模型的预测概率(需要存储)
                
                accept_prob = min(1.0, target_prob[draft_token] / (draft_prob + 1e-10))
                
                if torch.rand(1).item() < accept_prob:
                    # 接受
                    accepted_tokens.append(draft_token)
                else:
                    # 拒绝,从修正分布中采样
                    corrected_prob = torch.clamp(
                        target_prob - draft_prob, min=0
                    )
                    corrected_prob = corrected_prob / corrected_prob.sum()
                    new_token = torch.multinomial(corrected_prob, 1).item()
                    accepted_tokens.append(new_token)
                    break  # 拒绝后不再接受后续 token
            
            # ===== 4. 更新序列 =====
            if not accepted_tokens:
                break
            
            new_tokens = torch.tensor([accepted_tokens], device=input_ids.device)
            input_ids = torch.cat([input_ids, new_tokens], dim=-1)
            
            # 检查是否生成了 EOS
            if accepted_tokens[-1] == self.tokenizer.eos_token_id:
                break
        
        return self.tokenizer.decode(input_ids[0], skip_special_tokens=True)
    
    def _generate_draft(self, input_ids: torch.Tensor, k: int) -> torch.Tensor:
        """草稿模型生成 K 个 token"""
        with torch.no_grad():
            for _ in range(k):
                outputs = self.draft_model(input_ids)
                next_token = outputs.logits[:, -1:].argmax(dim=-1)
                input_ids = torch.cat([input_ids, next_token], dim=-1)
        
        return input_ids[:, -k:]  # 返回最后 K 个 token

六、分布式推理

6.1 张量并行(Tensor Parallelism)

# tensor_parallel.py

"""
张量并行(Tensor Parallelism):
将模型的每一层(如 Linear)按列/行切分到多个 GPU 上

切分方式(以 Linear 层为例):
- 按列切分(Column Parallel):A = [A1 | A2],Y = X @ A = [X@A1 | X@A2]
- 按行切分(Row Parallel):A = [A1; A2],Y = X1@A1 + X2@A2(需要 AllReduce)

Megatron-LM 风格:
- Attention 的 Q/K/V 按列切分
- MLP 的第一层按列切分,第二层按行切分
- 每层末尾需要一次 AllReduce

vLLM 中的 TP:
- tensor_parallel_size=N 自动启用
- 使用 NCCL 做 GPU 间通信
- 推荐:2 卡用 TP,4+ 卡用 TP+PP
"""

import torch
import torch.nn as nn
import torch.distributed as dist

class ColumnParallelLinear(nn.Module):
    """按列切分的 Linear 层(Megatron 风格)"""
    
    def __init__(self, in_features: int, out_features: int, device: int):
        super().__init__()
        self.in_features = in_features
        self.out_features = out_features
        
        # 每个 GPU 只存储部分权重
        self.weight = nn.Parameter(
            torch.empty(out_features // dist.get_world_size(), in_features)
        )
        self.bias = nn.Parameter(torch.empty(out_features // dist.get_world_size()))
    
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # 输入需要复制(Broadcast)
        # 输出按列拼接(All-Gather)
        output = F.linear(x, self.weight, self.bias)
        return output


class RowParallelLinear(nn.Module):
    """按行切分的 Linear 层(Megatron 风格)"""
    
    def __init__(self, in_features: int, out_features: int, device: int):
        super().__init__()
        self.in_features = in_features
        self.out_features = out_features
        
        # 每个 GPU 存储部分行
        self.weight = nn.Parameter(
            torch.empty(out_features, in_features // dist.get_world_size())
        )
        self.bias = None  # Row Parallel 无 bias(或各 GPU 单独加 bias 后 AllReduce)
    
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # 输入按行切分(Scatter)
        # 输出需要 AllReduce 求和
        output_partial = F.linear(x, self.weight)
        dist.all_reduce(output_partial)
        return output_partial

6.2 流水线并行(Pipeline Parallelism)

# pipeline_parallel.py

"""
流水线并行(Pipeline Parallelism):
将模型的不同层分配到不同 GPU 上

切分方式:
- GPU 0: Embedding + Layer 0-3
- GPU 1: Layer 4-7
- GPU 2: Layer 8-11
- GPU 3: Layer 12-15 + LM Head

挑战:
- 气泡时间(Pipeline Bubble):GPU 空闲等待
- 解决:GPipe(填充微批次)+ 1F1B(One Forward One Backward)

vLLM 中的 PP:
- pipeline_parallel_size=N 启用
- 使用 NCCL 做 GPU 间激活/梯度传输
- 推荐:8+ 卡考虑 PP
"""

# 简化版流水线并行
class PipelineParallel:
    def __init__(self, model_layers: list[nn.Module], num_stages: int):
        self.num_stages = num_stages
        self.layers_per_stage = len(model_layers) // num_stages
        
        # 分配层到各 GPU
        self.stages = []
        for i in range(num_stages):
            start = i * self.layers_per_stage
            end = start + self.layers_per_stage if i < num_stages - 1 else len(model_layers)
            stage_layers = model_layers[start:end]
            self.stages.append(nn.Sequential(*stage_layers).to(f"cuda:{i}"))
    
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """前向传播(流水线)"""
        for i, stage in enumerate(self.stages):
            x = x.to(f"cuda:{i}")
            x = stage(x)
        
        return x.to("cpu")

6.3 数据并行(Data Parallelism)

# data_parallel.py

"""
数据并行(Data Parallelism):
同一样模型复制到多个 GPU,不同 GPU 处理不同数据批次

方式:
- DP(DataParallel):主卡汇总梯度(低速)
- DDP(DistributedDataParallel):各卡独立计算,AllReduce 同步梯度(推荐)
- ZeRO(DeepSpeed):优化器状态/梯度/参数分片(显存优化)

推理时的数据并行:
- 每个 GPU 运行完整的模型副本
- 请求分配到不同 GPU(Round Robin / 负载均衡)
- vLLM: data_parallel_size=N 启用
"""

import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

def setup_ddp(
    model: nn.Module,
    local_rank: int,
    world_size: int,
):
    """设置 DDP"""
    torch.cuda.set_device(local_rank)
    
    # 初始化进程组
    dist.init_process_group(
        backend="nccl",
        init_method="env://",
        world_size=world_size,
        rank=local_rank,
    )
    
    # 包装模型
    model = DDP(
        model.cuda(),
        device_ids=[local_rank],
        output_device=local_rank,
    )
    
    return model

七、推理服务部署

7.1 vLLM + FastAPI 生产服务

# production_server.py

from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from vllm import LLM, SamplingParams
from typing import List, Optional
import asyncio
from contextlib import asynccontextmanager

# ===== 请求/响应模型 =====
class GenerateRequest(BaseModel):
    prompt: str
    max_tokens: int = 1024
    temperature: float = 0.7
    top_p: float = 0.95
    top_k: int = 50
    stop: Optional[List[str]] = None

class BatchGenerateRequest(BaseModel):
    prompts: List[str]
    max_tokens: int = 1024
    temperature: float = 0.7
    top_p: float = 0.95

class GenerateResponse(BaseModel):
    text: str
    tokens: int
    finish_reason: str

# ===== 全局 LLM 实例 =====
llm = None

@asynccontextmanager
async def lifespan(app: FastAPI):
    """应用生命周期管理"""
    global llm
    
    # 启动时初始化
    llm = LLM(
        model="meta-llama/Llama-3.1-8B-Instruct",
        tensor_parallel_size=2,
        dtype="bfloat16",
        gpu_memory_utilization=0.9,
        max_model_len=8192,
        enable_prefix_caching=True,
    )
    
    yield
    
    # 关闭时清理
    del llm

app = FastAPI(lifespan=lifespan)

# ===== API 端点 =====
@app.post("/generate", response_model=GenerateResponse)
async def generate(request: GenerateRequest):
    """单条生成"""
    sampling_params = SamplingParams(
        temperature=request.temperature,
        top_p=request.top_p,
        top_k=request.top_k,
        max_tokens=request.max_tokens,
        stop=request.stop or [],
    )
    
    outputs = llm.generate([request.prompt], sampling_params)
    output = outputs[0].outputs[0]
    
    return GenerateResponse(
        text=output.text,
        tokens=len(output.token_ids),
        finish_reason=output.finish_reason,
    )

@app.post("/batch_generate", response_model=List[GenerateResponse])
async def batch_generate(request: BatchGenerateRequest):
    """批量生成"""
    sampling_params = SamplingParams(
        temperature=request.temperature,
        top_p=request.top_p,
        max_tokens=request.max_tokens,
    )
    
    outputs = llm.generate(request.prompts, sampling_params)
    
    return [
        GenerateResponse(
            text=out.outputs[0].text,
            tokens=len(out.outputs[0].token_ids),
            finish_reason=out.outputs[0].finish_reason,
        )
        for out in outputs
    ]

@app.get("/health")
async def health():
    """健康检查"""
    return {"status": "ok", "model_loaded": llm is not None}

# ===== 启动 =====
if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=8000)

7.2 性能监控

# metrics.py

from prometheus_client import Counter, Histogram, Gauge, start_http_server
import time

# ===== Prometheus 指标 =====
REQUEST_COUNT = Counter(
    "llm_requests_total", 
    "Total LLM requests", 
    ["model", "endpoint"]
)

REQUEST_LATENCY = Histogram(
    "llm_request_duration_seconds",
    "LLM request latency",
    ["model", "endpoint"],
    buckets=[0.1, 0.5, 1.0, 2.0, 5.0, 10.0, 30.0, 60.0],
)

TOKEN_THROUGHPUT = Gauge(
    "llm_tokens_per_second",
    "Token generation throughput",
    ["model"],
)

GPU_MEMORY_USAGE = Gauge(
    "llm_gpu_memory_mb",
    "GPU memory usage in MB",
    ["model", "gpu_id"],
)

class LLMMetricsMiddleware:
    """LLM 推理监控中间件"""
    
    def __init__(self, model_name: str):
        self.model_name = model_name
    
    def record_request(self, endpoint: str):
        REQUEST_COUNT.labels(model=self.model_name, endpoint=endpoint).inc()
    
    def record_latency(self, endpoint: str, duration: float):
        REQUEST_LATENCY.labels(model=self.model_name, endpoint=endpoint).observe(duration)
    
    def record_throughput(self, tokens_per_second: float):
        TOKEN_THROUGHPUT.labels(model=self.model_name).set(tokens_per_second)
    
    def update_gpu_memory(self):
        """更新 GPU 显存使用"""
        import torch
        for i in range(torch.cuda.device_count()):
            allocated = torch.cuda.memory_allocated(i) / 1024**2
            GPU_MEMORY_USAGE.labels(model=self.model_name, gpu_id=str(i)).set(allocated)
    
    def __call__(self, func):
        """装饰器:自动记录指标"""
        import functools
        
        @functools.wraps(func)
        async def wrapper(*args, **kwargs):
            self.record_request(func.__name__)
            
            start = time.time()
            try:
                result = await func(*args, **kwargs)
                return result
            finally:
                duration = time.time() - start
                self.record_latency(func.__name__, duration)
        
        return wrapper

# 启动 Prometheus metrics 服务
def start_metrics_server(port: int = 9090):
    """启动 Prometheus metrics 暴露端点"""
    start_http_server(port)
    print(f"✅ Prometheus metrics 已启动: http://localhost:{port}/metrics")

八、最佳实践 Checklist

推理优化 Checklist

□ 引擎选型
  □ vLLM:通用场景(PagedAttention + Continuous Batching)
  □ TensorRT-LLM:NVIDIA GPU 极致性能(Kernel Fusion + FP8)
  □ SGLang:RadixAttention 场景(共享前缀)
  □ llama.cpp:CPU / 边缘设备

□ 量化
  □ AWQ(4-bit):生产推荐(速度 + 精度平衡)
  □ GPTQ(4-bit):研究 / 实验
  □ FP8:H100 专用(几乎无损 + 2-3x 加速)
  □ INT8 SmoothQuant:NVIDIA GPU 老架构

□ 分布式
  □ 张量并行(TP):2-4 卡
  □ 流水线并行(PP):8+ 卡
  □ 数据并行(DP):多实例负载均衡

□ 高级优化
  □ Prefix Caching(系统提示复用)
  □ Chunked Prefill(分块预填充,降低延迟)
  □ Speculative Decoding(草稿模型加速)
  □ Continuous Batching(动态批处理)

□ 部署
  □ OpenAI 兼容 API
  □ Prometheus 监控(延迟 / 吞吐量 / GPU 显存)
  □ 健康检查 + 优雅关停
  □ 请求限流 + 队列管理

总结

推理性能对比(Llama-3.1-8B,2×A100 80G)

配置 吞吐量(tokens/s) 延迟(ms) 显存占用
HuggingFace Transformers 800 120 16G
vLLM(无量化) 2400 45 14G
vLLM + AWQ 4-bit 4100 28 6G
TensorRT-LLM + INT8 5200 22 6G
TensorRT-LLM + FP8 5800 20 8G
SGLang + RadixAttention 3600 35 10G

本文覆盖 LLM 推理优化完整链路:vLLM(PagedAttention + Continuous Batching + API 服务)+ TensorRT-LLM(模型编译 + INT8/FP8 量化 + 推理服务)+ SGLang(RadixAttention + Python API)+ 量化技术(AWQ/GPTQ/FP8)+ 投机采样 + 分布式推理(TP/PP/DP)+ 生产服务部署(FastAPI + Prometheus 监控)。


下一步推荐:

  • 向量数据库深度对比(Milvus/Qdrant/Chroma)
  • 实时通信全链路(SSE/WebSocket + 流式响应)
  • 可观测性工程(OpenTelemetry + Grafana)
Logo

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

更多推荐