Skip to content

推理服务与性能优化 — 从能跑到跑得稳 ​

#Inference · #vLLM · #ContinuousBatching · #TensorParallel · #FlashAttention · #SpeculativeDecoding · #吞吐量 · #延迟 · #Token优化 · #上下文管理 · #PromptCaching

模型训好了,怎么让它在生产环境跑得又快又稳?本专题覆盖推理引擎选型、显存优化、并行策略、以及吞吐/延迟/成本三角平衡。


推理引擎全景对比 ​

mermaid
graph TD
    subgraph "推理引擎选择"
        Q1{"你的场景?"} -->|"本地单机<br/>消费级显卡"| A1["llama.cpp<br/>CPU/GPU 通用<br/>量化格式最丰富"]
        Q1 -->|"生产服务<br/>高并发"| A2["vLLM<br/>PagedAttention<br/>Continuous Batching"]
        Q1 -->|"NVIDIA 专用<br/>极致性能"| A3["TensorRT-LLM<br/>编译优化<br/>FP8/INT4 量化"]
        Q1 -->|"HuggingFace<br/>生态集成"| A4["TGI (Text-Generation-Inference)<br/>水塘采样<br/>Flash Attention"]
    end

    style A2 fill:#e74c3c,color:#fff
    style A1 fill:#2ecc71,color:#fff
引擎核心优化吞吐(7B on A100)首个 Token 延迟主要局限
vLLMPagedAttention + Cont. Batching~3000 tok/s~200ms显存管理有时不够极致
TensorRT-LLM编译优化 + FP8~5000 tok/s~100ms部署复杂,模型支持有限
TGIFlash Attention + 水塘采样~2500 tok/s~250ms生态绑定 HuggingFace
llama.cpp量化 + Metal/CUDA~80 tok/s (Mac)~500ms非专门服务端设计
SGLangRadixAttention + 结构化输出~3500 tok/s~150ms较新,社区还在成长

推理性能的核心瓶颈 ​

mermaid
sequenceDiagram
    participant User as 用户请求
    participant Engine as 推理引擎

    Note over User,Engine: 🔥 Prefill 阶段(计算密集)
    User->>Engine: 发送 Prompt(如 4000 token)
    Engine->>Engine: 并行计算所有 Token 的注意力
    Note over Engine: 瓶颈:GPU FLOPS<br/>RTX 4090: 82.6 TFLOPS (FP16)

    Note over User,Engine: 📝 Decode 阶段(显存带宽密集)
    loop 逐 Token 生成
        Engine->>Engine: 加载 KV Cache + 计算下一个 Token
        Engine->>User: 输出 1 Token
        Note over Engine: 瓶颈:显存带宽<br/>RTX 4090: 1.0 TB/s
    end

两个阶段的瓶颈差异 ​

阶段瓶颈类型受什么限制优化方向
PrefillCompute-boundGPU FLOPS更大 batch、FP8、Flash Attention
DecodeMemory-bound显存带宽KV Cache 压缩、GQA、Speculative Decoding

延迟拆解 ​

总延迟 = Prefill 延迟 + Decode 延迟 × 输出 Token 数 + 网络/排队延迟

Prefill 延迟 ≈ Prompt Token 数 / Prefill 吞吐 (tokens/s)
Decode 延迟 ≈ 1 / Decode 吞吐 (tokens/s)  ≈ 单 Token 生成时间

示例(vLLM + LLaMA-7B on A100):
  Prompt: 2000 token
  Prefill: 2000 / 5000 = 0.4s
  生成 500 token: 500 / 50 = 10s
  总延迟 ≈ 10.4s(Decode 占 96%)

vLLM 核心优化 ​

PagedAttention — 零碎显存的救星 ​

传统 KV Cache 管理方式要求为每个请求预分配连续的大块显存,导致严重的显存碎片和浪费。PagedAttention 将 KV Cache 分成固定大小的"页"(page),按需分配:

传统方式:
  请求1: [████████████████░░░░░░░░░░░░]  预留 4096,只用 2000
  请求2: [██████████░░░░░░░░░░░░░░░░░░░░░░]  预留 4096,只用 1000
  浪费: ~60% 显存

PagedAttention:
  请求1: [████][████][████][████]──  按需分配页
  请求2: [████][████]──
  利用率: ~95%
python
"""
PagedAttention 的核心思想示意(简化版)

原理:
1. 将 KV Cache 划分为固定大小的 Block(如 16 token/block)
2. 每个请求维护一个 Block Table,记录其 KV Cache 所在的物理 Block
3. 推理时根据 Block Table 找到对应的物理 Block 即可
"""

from dataclasses import dataclass
from typing import List, Optional

@dataclass
class KVCacheBlock:
    """KV Cache 的一个 Block"""
    block_id: int
    key: "torch.Tensor"    # (num_layers, block_size, num_heads, head_dim)
    value: "torch.Tensor"  # 同上
    ref_count: int = 0     # 引用计数(用于 prefix caching)

class BlockManager:
    """KV Cache Block 管理器(简化版)"""

    def __init__(self, num_blocks: int, block_size: int = 16):
        self.block_size = block_size
        self.free_blocks = list(range(num_blocks))
        self.block_table = {}  # request_id → [block_ids]

    def allocate(self, request_id: str, num_tokens: int) -> List[int]:
        """为一个请求分配 Block"""
        num_needed = (num_tokens + self.block_size - 1) // self.block_size
        if len(self.free_blocks) < num_needed:
            # 触发抢占(preemption):驱逐低优先级请求的 Block
            self._preempt(num_needed)

        blocks = self.free_blocks[:num_needed]
        self.free_blocks = self.free_blocks[num_needed:]
        self.block_table[request_id] = blocks
        return blocks

    def free(self, request_id: str):
        """释放一个请求的所有 Block"""
        if request_id in self.block_table:
            self.free_blocks.extend(self.block_table[request_id])
            del self.block_table[request_id]

    def _preempt(self, num_needed: int):
        """抢占 Block(实际实现更复杂,涉及 swap 到 CPU)"""
        pass

Continuous Batching — 告别串行等待 ​

传统静态批处理:一批请求必须全部完成才能开始下一批。一个长回答会阻塞整个 batch。

Continuous Batching:每生成一个 token,检查是否有请求完成或有新请求到达,动态调整 batch 组成:

mermaid
gantt
    title Continuous Batching 示意
    dateFormat X
    axisFormat %s

    section 请求A (短)
    生成 Token 1-10  : 0, 10
    完成  : 10, 11

    section 请求B (长)
    生成 Token 1-5   : 0, 5
    继续生成 Token 6-20 : 5, 20
    完成  : 20, 21

    section 请求C (新到达)
    等待  : 0, 10
    生成 Token 1-8   : 10, 18
    完成  : 18, 19

    section 请求D (新到达)
    等待  : 0, 18
    生成 Token 1-6   : 18, 24
    完成  : 24, 25
python
"""
Continuous Batching 调度核心逻辑(简化版)
"""

class ContinuousBatchScheduler:
    """Continuous Batching 调度器"""

    def __init__(self, max_batch_size: int = 32, max_tokens: int = 8192):
        self.max_batch_size = max_batch_size
        self.max_tokens = max_tokens
        self.running: dict = {}    # request_id → RequestState
        self.waiting: list = []    # 等待队列

    def step(self):
        """每一步:生成一个 token,然后重新调度"""
        # 1. 对正在运行的请求各生成一个 token
        finished = []
        for rid, state in self.running.items():
            token = self._generate_one(state)
            state.generated_tokens.append(token)

            if state.is_finished():
                finished.append(rid)

        # 2. 移除已完成的请求
        for rid in finished:
            self._free_resources(self.running.pop(rid))

        # 3. 从等待队列中拉取新请求
        while (len(self.running) < self.max_batch_size
               and self.waiting):
            req = self.waiting.pop(0)
            self.running[req.id] = req
            # Prefill 阶段:一次处理完所有 Prompt Token
            self._prefill(req)

    def _can_add_request(self) -> bool:
        """判断是否还能加入新请求"""
        total_tokens = sum(
            len(s.prompt_tokens) + len(s.generated_tokens)
            for s in self.running.values()
        )
        return (len(self.running) < self.max_batch_size
                and total_tokens < self.max_tokens)

# 优势:吞吐量比静态批处理高 2-10x
# 尤其在长短请求混合时效果显著

Prefix Cache — 共享 Prompt 只算一次 ​

当多个请求共享相同的 System Prompt 或前缀时,Prefix Cache 避免重复计算:

python
"""
Prefix Cache 原理

场景:
  System Prompt (500 token) + User Query 1 (100 token)
  System Prompt (500 token) + User Query 2 (100 token)

无 Prefix Cache:
  请求1: 计算 600 token 的 KV Cache
  请求2: 重新计算 600 token 的 KV Cache → 浪费!

有 Prefix Cache:
  请求1: 计算 600 token,缓存 System Prompt 的 KV (500 token)
  请求2: 复用缓存的 500 token KV,只计算 User Query 的 100 token
  节省: 500/600 ≈ 83% 的 Prefill 计算量

RadixAttention — SGLang 的 Prefix Caching 实现 ​

SGLang 使用 Radix Tree(基数树) 管理 Prefix Cache。每个节点存储一个 token 的 KV Cache,通过自动检测前缀匹配来复用:

mermaid
graph TD
    ROOT["Root"] --> SYS["System Prompt<br/>500 tokens<br/>KV cached"]
    SYS --> Q1["User Q1<br/>100 tokens<br/>KV cached"]
    SYS --> Q2["User Q2<br/>100 tokens<br/>复用 SYS KV"]
    SYS --> Q3["User Q3<br/>150 tokens<br/>复用 + 增量命中"]

    style SYS fill:#2ecc71,color:#fff
    style Q2 fill:#3498db,color:#fff
python
"""
RadixAttention 前缀匹配逻辑(简化)

核心优势:
1. 自动前缀匹配 — 不需要显式声明"这是前缀"
2. 增量缓存 — 即使只匹配部分前缀,也能复用
3. 树结构 — 支持多级分支(不同的 System Prompt 变体)
"""

class RadixNode:
    """Radix Tree 节点(存储 KV Cache 块)"""
    def __init__(self):
        self.children: Dict[int, "RadixNode"] = {}  # token → 子节点
        self.kv_cache: Optional[KVCacheBlock] = None
        self.ref_count: int = 0

class RadixCache:
    """Radix Tree 管理的 Prefix Cache"""

    def __init__(self):
        self.root = RadixNode()

    def match_prefix(self, tokens: List[int]) -> Tuple[int, List[KVCacheBlock]]:
        """
        匹配最长前缀

        Returns:
            (matched_length, cached_kv_blocks)
        """
        node = self.root
        matched = 0
        cached_blocks = []

        for token in tokens:
            if token in node.children:
                node = node.children[token]
                if node.kv_cache:
                    cached_blocks.append(node.kv_cache)
                matched += 1
            else:
                break

        return matched, cached_blocks

    def insert(self, tokens: List[int], kv_blocks: List[KVCacheBlock]):
        """插入新的前缀到 Radix Tree"""
        node = self.root
        for i, token in enumerate(tokens):
            if token not in node.children:
                node.children[token] = RadixNode()
            node = node.children[token]
            if i < len(kv_blocks):
                node.kv_cache = kv_blocks[i]
                node.ref_count += 1


# 性能对比:
# 传统 Attention:每个请求完整计算 Prefill(O(N²))
# RadixAttention:前缀匹配 → 只计算增量部分
# 典型场景(System Prompt 500 + User 100 token):
#   传统: 600 token Prefill
#   Radix: 100 token Prefill(节省 83%)

Chunked Prefill — 让长短请求和平共处 ​

Prefill 阶段一次性处理大量 Token 会阻塞其他请求。Chunked Prefill 将长的 Prefill 拆成多个小块,与 Decode 交替执行:

python
"""
Chunked Prefill 调度策略

问题场景:
  请求A: Prompt 10000 token(需要 2s Prefill)
  请求B: Prompt 100 token(只需要 20ms Prefill)

无 Chunked Prefill:
  请求B 必须等请求A 完成 2s Prefill → 总延迟 2s+

有 Chunked Prefill:
  请求A Prefill 500 token → 请求B Prefill 100 token → 请求A Prefill 500 token → ...
  请求B 在 20ms 内得到响应
"""

class ChunkedPrefillScheduler:
    """Chunked Prefill 调度器"""

    def __init__(self, chunk_size: int = 512):
        """
        Args:
            chunk_size: 每个 Chunk 的 Token 数
        """
        self.chunk_size = chunk_size

    def schedule_prefill(self, prefill_queue: List[Request]) -> List[Tuple[Request, int]]:
        """
        调度 Prefill,按 Chunk 交替执行

        Returns:
            [(request, tokens_to_process), ...]
        """
        schedule = []
        remaining = {}

        for req in prefill_queue:
            remaining[req.id] = req.prompt_length

        # 轮询每个请求,每次分配一个 Chunk
        while remaining:
            for req_id in list(remaining.keys()):
                if remaining[req_id] > 0:
                    chunk = min(remaining[req_id], self.chunk_size)
                    schedule.append((req_id, chunk))
                    remaining[req_id] -= chunk
                else:
                    del remaining[req_id]

        return schedule


# vLLM 配置:
# --max-num-batched-tokens 8192     # 每次迭代最多处理 Token 数
# --max-num-seqs 256                # 最大并发请求数
# --enable-chunked-prefill          # 启用 Chunked Prefill

Prefix Cache 的典型共享模式 ​

模式场景缓存命中率收益
System Prompt所有请求共享同一 System Prompt90%+Prefill 降低 80%+
Few-shotN 个示例在所有请求中使用80%+示例越多收益越大
RAG Context同一文档被多次查询60-80%只缓存检索到的文档片段
Multi-turn同一会话的历史消息50-70%对话越长收益越大
Batch 批处理同一 Prompt 处理多个输入95%+几乎零 Prefill 开销
bash
# vLLM 中自动启用 Prefix Caching:
python -m vllm.entrypoints.openai.api_server \
  --model Qwen/Qwen4-7B-Instruct \
  --enable-prefix-caching \  --enable-chunked-prefill \       # Chunked Prefill
  --max-num-batched-tokens 8192    # 每次迭代最大 Token 数

并行策略 ​

Tensor Parallelism (TP) — 单层切分 ​

将一层 Transformer 的权重矩阵切分到多张 GPU:

mermaid
graph LR
    subgraph "单卡"
        IN1["输入 (B, H)"] --> W1["W (H, 4H) 完整"]
        W1 --> OUT1["输出 (B, 4H)"]
    end

    subgraph "Tensor Parallel 2卡"
        IN2["输入 (B, H)"] --> W2A["W_A (H, 2H) GPU0"]
        IN2 --> W2B["W_B (H, 2H) GPU1"]
        W2A --> ALLREDUCE["All-Reduce"]
        W2B --> ALLREDUCE
        ALLREDUCE --> OUT2["输出 (B, 4H)"]
    end

    style ALLREDUCE fill:#e74c3c,color:#fff
bash
# vLLM 启用 Tensor Parallelism
python -m vllm.entrypoints.openai.api_server \
  --model Qwen/Qwen4-72B-Instruct \
  --tensor-parallel-size 4 \    # 4 张 GPU 切分
  --gpu-memory-utilization 0.9

Pipeline Parallelism (PP) — 按层切分 ​

将模型按层切分到多张 GPU:GPU0 放前 20 层,GPU1 放后 20 层。

三种并行的适用场景 ​

策略切分维度通信量适用场景
Data Parallel (DP)Batch参数同步(低)多卡训练(最常用)
Tensor Parallel (TP)算子每层都要通信(高)单机多卡推理
Pipeline Parallel (PP)层层间传递(中)跨机推理(超大模型)

推理优先选 TP:虽然通信多,但延迟最低。TP 不够再叠加 PP。

Data Parallel 推理 — 简单但有效 ​

bash
# 方案:每张 GPU 独立加载完整模型,请求轮询分发
# 优点:零通信开销,线性扩展吞吐
# 缺点:每张 GPU 需要加载完整模型(显存要求高)

# vLLM 方式:启动多个实例 + 前置负载均衡
# 或使用 vLLM 内置的 --data-parallel-size

更极致的优化 ​

Flash Attention — 2-4x 加速注意力 ​

传统 Attention 需要显式构造 N×N 的注意力矩阵,Flash Attention 通过分块(tiling)和重计算(recomputation)避免完整矩阵的显存写入:

python
# PyTorch 2.0+ 自动使用 Flash Attention(如果可用)
import torch.nn.functional as F

# 传统 Self-Attention
# Q, K, V: (batch, heads, seq_len, head_dim)
# attention = softmax(Q @ K.T / sqrt(d_k)) @ V  # O(N²) 显存

# Flash Attention(PyTorch 2.0 内置)
output = F.scaled_dot_product_attention(
    Q, K, V,
    is_causal=True,       # 因果掩码
    dropout_p=0.0,
)
# 自动选择最优 backend: FlashAttention-2 > Memory Efficient > 原始实现

效果:训练时显存降低 5-20x,推理时延迟降低 2-4x。几乎所有现代推理引擎都默认启用。

Speculative Decoding — 用小模型猜大模型 ​

核心思想:用一个小模型快速生成候选 Token,然后让大模型一次性验证:

mermaid
sequenceDiagram
    participant Draft as 草稿模型 (小)
    participant Target as 目标模型 (大)

    Draft->>Draft: 快速生成 5 个候选 Token
    Draft->>Target: [Token₁, Token₂, Token₃, Token₄, Token₅]
    Target->>Target: 并行验证 5 个 Token
    Target-->>Draft: 接受前 3 个,拒绝后 2 个
    Draft->>Draft: 从第 4 个重新生成...
python
"""
Speculative Decoding 的性能计算

假设:
  大模型 Decode: 40 tok/s
  小模型 Decode: 200 tok/s
  大小模型验证: 500 tok/s(批处理 5 个一起验)
  接受率: 80%

无优化: 40 tok/s
有优化:
  小模型生成 5 个: 5/200 = 0.025s
  大模型验证 5 个: 5/500 = 0.01s
  接受 4 个 (80%)
  吞吐 = 4 / (0.025 + 0.01) ≈ 114 tok/s
  加速: 114/40 ≈ 2.85x
"""

# vLLM 中使用:
# --speculative-model <draft_model_name>
# 草稿模型通常选同系列的 0.5B/1.5B 版本

量化推理 — 精度换速度 ​

量化格式显存节省 (vs FP16)推理加速质量损失推荐引擎
FP16baseline1x无所有引擎
INT8~50%1.3-1.5x几乎无TensorRT-LLM
FP8~50%1.5-2x几乎无(H100)TensorRT-LLM, vLLM
INT4 (GPTQ/AWQ)~75%1.5-2x轻微vLLM, TGI
INT4 (GGUF)~75%1.2-1.5x (GPU)轻微llama.cpp

性能监控与压测 ​

关键指标 ​

指标定义目标值计算方式
TTFT首个 Token 延迟<500ms请求发出 → 收到第一个 Token
TPOT每 Token 生成间隔<30ms相邻 Token 之间的时间差
Throughput吞吐量最大化tokens/s(所有并发请求)
QPS每秒请求数根据业务1 / (总延迟)
Queue Time排队时间<100ms请求到达 → 开始处理
GPU UtilizationGPU 利用率>80%nvidia-smi

压测脚本 ​

python
"""
LLM 推理服务压测框架
"""

import time
import asyncio
import aiohttp
import numpy as np
from dataclasses import dataclass, field
from typing import List

@dataclass
class LoadTestResult:
    """压测结果"""
    total_requests: int
    successful: int
    failed: int

    # 延迟指标 (ms)
    ttft_p50: float
    ttft_p95: float
    ttft_p99: float
    total_latency_p50: float
    total_latency_p95: float

    # 吞吐指标
    total_tokens: int
    tokens_per_second: float
    requests_per_second: float

    # 错误
    errors: List[str] = field(default_factory=list)

async def load_test(
    url: str,
    prompts: List[str],
    concurrency: int = 10,
    max_tokens: int = 256,
) -> LoadTestResult:
    """
    对 LLM API 进行并发压测

    用法:
    result = await load_test(
        "http://localhost:8000/v1/chat/completions",
        ["写一个排序算法"] * 50,
        concurrency=10,
    )
    """
    ttfts = []
    total_latencies = []
    total_tokens = 0
    successful = 0
    failed = 0
    errors = []

    semaphore = asyncio.Semaphore(concurrency)

    async def single_request(session, prompt, idx):
        nonlocal successful, failed, total_tokens
        async with semaphore:
            start = time.perf_counter()
            try:
                async with session.post(
                    url,
                    json={
                        "model": "default",
                        "messages": [{"role": "user", "content": prompt}],
                        "max_tokens": max_tokens,
                        "temperature": 0,
                        "stream": True,  # 流式才能测 TTFT
                    },
                    timeout=120,
                ) as resp:
                    first_token_time = None
                    async for line in resp.content:
                        if first_token_time is None:
                            first_token_time = time.perf_counter()
                            ttfts.append((first_token_time - start) * 1000)
                        total_tokens += 1

                    total_latencies.append((time.perf_counter() - start) * 1000)
                    successful += 1
            except Exception as e:
                failed += 1
                errors.append(f"请求 {idx}: {str(e)}")

    async with aiohttp.ClientSession() as session:
        tasks = [
            single_request(session, p, i)
            for i, p in enumerate(prompts)
        ]
        start_time = time.perf_counter()
        await asyncio.gather(*tasks)
        elapsed = time.perf_counter() - start_time

    return LoadTestResult(
        total_requests=len(prompts),
        successful=successful,
        failed=failed,
        ttft_p50=np.percentile(ttfts, 50) if ttfts else 0,
        ttft_p95=np.percentile(ttfts, 95) if ttfts else 0,
        ttft_p99=np.percentile(ttfts, 99) if ttfts else 0,
        total_latency_p50=np.percentile(total_latencies, 50) if total_latencies else 0,
        total_latency_p95=np.percentile(total_latencies, 95) if total_latencies else 0,
        total_tokens=total_tokens,
        tokens_per_second=total_tokens / elapsed if elapsed > 0 else 0,
        requests_per_second=successful / elapsed if elapsed > 0 else 0,
        errors=errors,
    )

# ===== 运行压测 =====
# asyncio.run(load_test(...))

性能退化监控 ​

python
"""
推理服务运行时性能监控
"""

class InferenceMonitor:
    """持续监控推理服务的关键指标"""

    def __init__(self):
        self.metrics = {
            "ttft_samples": [],
            "tpot_samples": [],
            "queue_wait_samples": [],
            "error_count": 0,
            "total_requests": 0,
        }

    def record_request(self, ttft_ms: float, tpot_ms: float,
                       queue_ms: float, success: bool):
        """记录一次请求的指标"""
        self.metrics["total_requests"] += 1
        if success:
            self.metrics["ttft_samples"].append(ttft_ms)
            self.metrics["tpot_samples"].append(tpot_ms)
            self.metrics["queue_wait_samples"].append(queue_ms)
        else:
            self.metrics["error_count"] += 1

    def check_health(self) -> dict:
        """检查服务是否健康"""
        issues = []

        # 最近 100 个请求的 P95 TTFT
        recent_ttft = self.metrics["ttft_samples"][-100:]
        if recent_ttft and np.percentile(recent_ttft, 95) > 1000:
            issues.append(f"P95 TTFT > 1s: {np.percentile(recent_ttft, 95):.0f}ms")

        # 错误率
        error_rate = self.metrics["error_count"] / max(self.metrics["total_requests"], 1)
        if error_rate > 0.01:
            issues.append(f"错误率 > 1%: {error_rate:.2%}")

        # 排队时间
        recent_queue = self.metrics["queue_wait_samples"][-100:]
        if recent_queue and np.percentile(recent_queue, 95) > 500:
            issues.append("P95 排队时间 > 500ms,考虑扩容")

        return {
            "healthy": len(issues) == 0,
            "issues": issues,
            "error_rate": error_rate,
        }

推理服务最佳实践 ​

mermaid
graph TD
    subgraph "部署清单"
        C1["✅ 选型:vLLM (通用) / TensorRT-LLM (极致性能)"] --> C2["✅ 启 Flash Attention"]
        C2 --> C3["✅ 配 Continuous Batching"]
        C3 --> C4{"显存不够?"}
        C4 -->|是| C5["开量化 (AWQ/GPTQ INT4)"]
        C4 -->|否| C6{"延迟太高?"}
        C6 -->|是| C7["开 Speculative Decoding"]
        C6 -->|否| C8{"并发太大?"}
        C8 -->|是| C9["加 Data Parallel / 多实例"]
        C8 -->|否| C10["✅ 部署上线 + 持续监控"]
    end

    style C10 fill:#2ecc71,color:#fff
场景推荐配置预期效果
7B 模型 + 单 RTX 4090vLLM + AWQ INT4~100 tok/s, 并发 4-8
7B 模型 + A100vLLM + FP16 + Continuous Batching~3000 tok/s, 并发 32+
70B 模型 + 4×A100vLLM + TP=4 + INT4~500 tok/s per request
低延迟实时对话小模型 + Speculative DecodingTTFT < 100ms
离线批量推理vLLM 离线模式 + 大 batch最大化 GPU 利用率

线上部署 Checklist ​

  • [ ] 模型预加载到 GPU(避免首次推理冷启动)
  • [ ] 配置健康检查端点(/health)
  • [ ] 配置请求超时 + 优雅降级
  • [ ] 配置速率限制 + 请求队列上限
  • [ ] 接入 Prometheus + Grafana 监控
  • [ ] 准备降级模型(小模型兜底)
  • [ ] 压测验证:确定最大 QPS、P95 延迟、长尾延迟
  • [ ] 灰度发布:先切 10% 流量观察 24h

Token 优化与上下文管理 — 成本与质量平衡 ​

推理性能优化的另一面是 Token 经济性。每次 API 调用都在消耗 Token,优化 Token 使用直接影响成本和延迟。

Token 消耗全景 ​

mermaid
graph LR
    subgraph "Token 消耗来源"
        SYS["System Prompt<br/>固定开销 200-2000 tok"] --> TOTAL["总 Token 消耗"]
        CTX["历史对话<br/>增长最快 1k-100k+ tok"] --> TOTAL
        DOC["检索文档<br/>RAG 上下文 2k-10k tok"] --> TOTAL
        QUERY["用户问题<br/>100-500 tok"] --> TOTAL
        OUT["模型输出<br/>200-2000 tok"] --> TOTAL
    end

    style CTX fill:#e74c3c,color:#fff
    style SYS fill:#f39c12,color:#fff

策略一:Prompt 压缩 — 减负不降质 ​

python
"""
Prompt 压缩策略
"""

from typing import List
import hashlib


class PromptCompressor:
    """Prompt 压缩器"""

    def __init__(self, llm_client):
        self.llm = llm_client
        self._compression_cache = {}  # 压缩结果缓存

    # ===== 1. 长对话压缩:将早期对话提炼为摘要 =====
    async def compress_conversation(
        self, messages: List[dict], keep_recent: int = 6
    ) -> List[dict]:
        """压缩对话历史:保留最近 N 条,更早的压缩为摘要"""
        if len(messages) <= keep_recent:
            return messages

        recent = messages[-keep_recent:]
        old = messages[:-keep_recent]

        # 用 LLM 将早期对话压缩为摘要
        old_text = "\n".join(
            [f"[{m['role']}]: {m['content'][:200]}" for m in old]
        )
        cache_key = hashlib.md5(old_text.encode()).hexdigest()

        if cache_key in self._compression_cache:
            summary = self._compression_cache[cache_key]
        else:
            summary = await self.llm.chat(
                f"将以下对话历史压缩为一条简洁摘要(保留关键事实、决策和结论):\n\n{old_text}"
            )
            self._compression_cache[cache_key] = summary

        # 构造压缩后的消息列表
        compressed = [
            {"role": "system", "content": f"[历史对话摘要]: {summary}"}
        ] + recent

        saved_tokens = sum(len(m["content"]) // 4 for m in old) - len(summary) // 4
        print(f"对话压缩: 节省约 {saved_tokens} tokens")

        return compressed

    # ===== 2. 检索文档压缩:去噪、去重、截断 =====
    def compress_documents(
        self, documents: List[dict], max_tokens: int = 3000
    ) -> List[dict]:
        """压缩检索到的文档"""
        seen = set()
        compressed = []
        total_tokens = 0

        for doc in sorted(documents, key=lambda d: d.get("score", 0), reverse=True):
            # 去重
            doc_hash = hashlib.md5(doc["content"][:100].encode()).hexdigest()
            if doc_hash in seen:
                continue
            seen.add(doc_hash)

            # Token 预算控制
            content = doc["content"]
            estimated_tokens = len(content) // 3  # 粗略估算:3 字符 ≈ 1 token

            if total_tokens + estimated_tokens > max_tokens:
                # 截断当前文档
                remaining = max_tokens - total_tokens
                if remaining > 100:
                    content = content[: remaining * 3] + "..."
                else:
                    break

            compressed.append({**doc, "content": content})
            total_tokens += min(estimated_tokens, max_tokens - total_tokens)

        return compressed

    # ===== 3. System Prompt 精简 =====
    def optimize_system_prompt(
        self, system_prompt: str, target_tokens: int = 500
    ) -> str:
        """精简 System Prompt"""
        rules = system_prompt.split("\n")
        essential = []
        trivial = []
        current_tokens = 0

        for rule in rules:
            rule = rule.strip()
            if not rule:
                continue
            est = len(rule) // 3
            if "必须" in rule or "禁止" in rule or "重要" in rule:
                essential.append((0, rule, est))  # priority 0 = 最高
            elif "建议" in rule or "可以" in rule:
                trivial.append((2, rule, est))
            else:
                essential.append((1, rule, est))

        # 按优先级填充
        result = []
        for _, rule, est in sorted(essential + trivial):
            if current_tokens + est > target_tokens:
                break
            result.append(rule)
            current_tokens += est

        return "\n".join(result)

策略二:Token Budget 管理 ​

python
"""
Token Budget 管理 — 动态分配 Token 配额
"""

from dataclasses import dataclass, field
from enum import Enum


class BudgetStrategy(Enum):
    BALANCED = "balanced"   # 均衡分配
    CONTEXT_HEAVY = "context_heavy"   # 上下文优先
    OUTPUT_HEAVY = "output_heavy"     # 输出优先


@dataclass
class TokenBudget:
    """Token 预算分配"""
    model_max_tokens: int       # 模型上下文窗口上限
    reserved_output: int        # 预留输出 Token
    available: int = 0          # 可用输入 Token
    system_prompt: int = 0      # 系统提示词占用
    conversation: int = 0       # 对话历史占用
    documents: int = 0          # 检索文档占用
    user_query: int = 0         # 用户问题占用
    remaining: int = 0          # 剩余可用

    @classmethod
    def create(
        cls,
        model_max_tokens: int,
        strategy: BudgetStrategy = BudgetStrategy.BALANCED,
        max_output_tokens: int = 2048,
    ) -> "TokenBudget":
        """创建预算"""
        # 输出分配策略
        output_map = {
            BudgetStrategy.BALANCED: max_output_tokens,
            BudgetStrategy.CONTEXT_HEAVY: min(max_output_tokens, 1024),
            BudgetStrategy.OUTPUT_HEAVY: min(max_output_tokens, 4096),
        }
        reserved = output_map[strategy]

        return cls(
            model_max_tokens=model_max_tokens,
            reserved_output=reserved,
            available=model_max_tokens - reserved,
            remaining=model_max_tokens - reserved,
        )

    def allocate(self, category: str, tokens: int) -> bool:
        """分配 Token,返回是否成功"""
        if tokens > self.remaining:
            return False
        setattr(self, category, getattr(self, category) + tokens)
        self.remaining -= tokens
        return True

    def report(self) -> str:
        """预算报告"""
        return (
            f"Token Budget [{self.remaining}/{self.available} available]:\n"
            f"  System Prompt: {self.system_prompt}\n"
            f"  Conversation:  {self.conversation}\n"
            f"  Documents:     {self.documents}\n"
            f"  User Query:    {self.user_query}\n"
            f"  Reserved Out:  {self.reserved_output}\n"
            f"  Total:         {self.available - self.remaining}/{self.model_max_tokens}"
        )


class TokenBudgetManager:
    """Token 预算管理器 — 确保不超出模型上下文窗口"""

    def __init__(self, llm_client):
        self.llm = llm_client
        self.compressor = PromptCompressor(llm_client)

    async def build_context(
        self,
        system_prompt: str,
        messages: List[dict],
        documents: List[dict],
        user_query: str,
        model_max_tokens: int = 128000,
        strategy: BudgetStrategy = BudgetStrategy.BALANCED,
    ) -> dict:
        """构建上下文并管理 Token 预算"""
        budget = TokenBudget.create(model_max_tokens, strategy)

        # 1. 分配 System Prompt
        sys_tokens = len(system_prompt) // 3
        if not budget.allocate("system_prompt", sys_tokens):
            # System Prompt 太长,压缩
            system_prompt = self.compressor.optimize_system_prompt(
                system_prompt, budget.remaining
            )
            sys_tokens = len(system_prompt) // 3
            budget.allocate("system_prompt", sys_tokens)

        # 2. 分配对话历史(先压缩)
        compressed_msgs = await self.compressor.compress_conversation(
            messages, keep_recent=6
        )
        conv_tokens = sum(len(m["content"]) // 3 for m in compressed_msgs if isinstance(m.get("content"), str))
        budget.allocate("conversation", conv_tokens)

        # 3. 分配检索文档(剩余空间)
        remaining_for_docs = budget.remaining - (len(user_query) // 3) - 500  # 500 buffer
        if remaining_for_docs > 0:
            documents = self.compressor.compress_documents(
                documents, max_tokens=remaining_for_docs
            )

        doc_tokens = sum(len(d["content"]) // 3 for d in documents)
        budget.allocate("documents", doc_tokens)

        # 4. 分配用户问题
        query_tokens = len(user_query) // 3
        budget.allocate("user_query", query_tokens)

        return {
            "system_prompt": system_prompt,
            "messages": compressed_msgs,
            "documents": documents,
            "user_query": user_query,
            "budget": budget,
        }

策略三:API 级 Prompt Caching ​

主流 LLM API 都提供了 Prompt Caching 功能,对重复的 Prompt 前缀自动缓存,节省成本和延迟:

平台功能名缓存规则成本节省目前状态
AnthropicPrompt Caching1024+ token 的重复前缀写入 +25%,读取 -90%正式可用
OpenAIAutomatic Caching1024+ token 的重复前缀,5-10min TTL自动 -50%正式可用
Google GeminiContext Caching32K+ token,可手动管理 TTL按存储收费,读取 -75%正式可用
DeepSeekContext Caching命中缓存时自动应用自动折扣正式可用
python
"""
API Prompt Caching 最佳实践
"""

class PromptCacheOptimizer:
    """利用 API Prompt Caching 优化 Token 成本"""

    @staticmethod
    def structure_for_cache(
        system_prompt: str,
        static_context: str,    # 静态上下文(如项目文档、规则)
        dynamic_context: str,   # 动态上下文(如对话历史)
        user_query: str,
    ) -> List[dict]:
        """按缓存友好顺序组织消息

        关键原则:将静态内容放在前面,动态内容放在后面。
        API 对连续重复的 token 前缀进行缓存。
        """
        return [
            # 第一层:完全静态(每次都一样,HIT 率最高)
            {"role": "system", "content": system_prompt},

            # 第二层:准静态(按需加载的文档,复用率高)
            {"role": "system", "content": f"[参考文档]\n{static_context}"},

            # 第三层:动态(变化频繁,基本不 HIT)
            {"role": "system", "content": f"[对话历史]\n{dynamic_context}"},

            # 第四层:用户问题
            {"role": "user", "content": user_query},
        ]

    @staticmethod
    def estimate_savings(
        static_tokens: int,
        requests_per_minute: int,
        input_price_per_1k: float = 0.003,
    ) -> dict:
        """估算缓存带来的成本节省"""
        # 假设缓存命中率 80%(静态部分)
        hit_rate = 0.8
        cached_tokens_per_request = static_tokens * hit_rate
        saved_per_request = cached_tokens_per_request * input_price_per_1k / 1000
        saved_per_hour = saved_per_request * requests_per_minute * 60
        saved_per_month = saved_per_hour * 24 * 30

        return {
            "tokens_cached_per_request": int(cached_tokens_per_request),
            "saved_per_request": f"${saved_per_request:.6f}",
            "saved_per_hour": f"${saved_per_hour:.4f}",
            "saved_per_month": f"${saved_per_month:.2f}",
        }


# 示例:每天 100 万次请求,System Prompt 2000 tokens
# optimizer = PromptCacheOptimizer()
# savings = optimizer.estimate_savings(static_tokens=2000, requests_per_minute=700)
# print(savings)
# => {'tokens_cached_per_request': 1600, 'saved_per_request': '$0.000005',
#     'saved_per_hour': '$0.2016', 'saved_per_month': '$145.15'}

Token 优化清单 ​

优化策略预期节省实施复杂度适用场景
长对话压缩为摘要30-70%低多轮对话 Agent
System Prompt 精简20-50%低所有场景
检索文档去重截断20-40%低RAG 场景
静态内容前置(API Caching)50-90%(静态部分)低有固定 System Prompt 的场景
Token Budget 动态管理防溢出中超长上下文场景
使用更便宜的模型做预处理30-50%中分类、实体提取等辅助任务
缓存常见问题的回答100%(命中时)中FAQ、客服场景

核心原则:Token 优化不是"无脑削减",而是用最少的 Token 传递最大的信息量。压缩摘要时保留关键事实,截断文档时保留高相关性内容,精简 Prompt 时保留硬性约束。


长上下文推理技术 — 突破 Context Window 限制 ​

为什么长上下文是难题? ​

标准 Attention 的计算复杂度为 O(n2),KV Cache 显存占用与序列长度线性增长。当 Context Window 从 4K 扩展到 200K+ 时,推理成本急剧上升:

上下文长度KV Cache (70B, FP16)Attention 计算量首 Token 延迟
4K~5 GB基准~200ms
32K~40 GB64x~1.5s
128K~160 GB1024x~6s
1M~1.2 TB62500x~45s

Attention Sink — 无限长度的流式推理 ​

核心发现(StreamingLLM, 2023):LLM 推理时,无论上下文多长,前几个 Token(Sink Tokens)始终获得极高的注意力分数。即使这些 Token 的语义无关紧要,删除它们会导致模型崩溃。

mermaid
graph LR
    subgraph "传统滑动窗口(会崩溃)"
        W1["Token 5-1024<br/>丢弃了前4个Token"] --> FAIL["❌ 注意力分布异常<br/>PPL 爆炸"]
    end

    subgraph "Attention Sink(稳定)"
        SINK["Sink Tokens<br/>(前4个Token)"] --> OK["✅ 注意力锚点保留"]
        WINDOW["滑动窗口<br/>(最近1020个Token)"] --> OK
        OK --> STABLE["稳定推理<br/>无限长度"]
    end

    style SINK fill:#e74c3c,color:#fff
    style STABLE fill:#2ecc71,color:#fff
python
class StreamingLLMCache:
    """
    Attention Sink + 滑动窗口 = 无限长度推理
    保留前 N 个 Sink Token + 最近 W 个 Token 的 KV Cache
    """

    def __init__(self, num_sink_tokens: int = 4, window_size: int = 1020):
        self.num_sink = num_sink_tokens
        self.window_size = window_size
        self.kv_cache = None  # (K, V) tensors

    def update(self, new_k, new_v):
        """每生成一个新 Token,更新 KV Cache"""
        if self.kv_cache is None:
            self.kv_cache = (new_k, new_v)
            return self.kv_cache

        cached_k, cached_v = self.kv_cache
        # 拼接新 Token
        cached_k = torch.cat([cached_k, new_k], dim=-2)
        cached_v = torch.cat([cached_v, new_v], dim=-2)

        seq_len = cached_k.shape[-2]
        max_len = self.num_sink + self.window_size

        if seq_len > max_len:
            # 保留 Sink Tokens + 最近的 Window Tokens
            sink_k = cached_k[:, :, :self.num_sink, :]
            sink_v = cached_v[:, :, :self.num_sink, :]
            recent_k = cached_k[:, :, -(self.window_size):, :]
            recent_v = cached_v[:, :, -(self.window_size):, :]

            cached_k = torch.cat([sink_k, recent_k], dim=-2)
            cached_v = torch.cat([sink_v, recent_v], dim=-2)

        self.kv_cache = (cached_k, cached_v)
        return self.kv_cache

    @property
    def effective_length(self):
        if self.kv_cache is None:
            return 0
        return self.kv_cache[0].shape[-2]


# 使用示例
cache = StreamingLLMCache(num_sink_tokens=4, window_size=1020)
print(f"最大 KV Cache 长度: {cache.num_sink + cache.window_size}")
print(f"可处理的输入长度: 无限(流式)")

YaRN — 位置编码外推(训练短、推理长) ​

YaRN(Yet another RoPE extensioN)解决的问题:模型在 4K 上下文训练,但需要在 128K 上推理。核心思想是对 RoPE 的不同频率分量采用不同的缩放策略:

python
import math

def yarn_rope_scaling(
    dim: int,
    max_position: int = 131072,    # 目标长度 128K
    original_max: int = 4096,       # 训练长度 4K
    beta_fast: float = 32.0,
    beta_slow: float = 1.0,
):
    """
    YaRN RoPE 缩放:
    - 高频分量(局部位置信息):不缩放
    - 低频分量(全局位置信息):线性缩放
    - 中间频率:平滑插值
    """
    scale = max_position / original_max  # 缩放因子 = 32

    freqs = []
    for i in range(0, dim, 2):
        # 原始 RoPE 频率
        freq = 1.0 / (10000.0 ** (i / dim))

        # 计算该频率的波长
        wavelength = 2 * math.pi / freq

        # 根据波长决定缩放策略
        if wavelength < beta_fast * original_max:
            # 高频:不缩放(保留局部位置精度)
            scaled_freq = freq
        elif wavelength > beta_slow * original_max:
            # 低频:线性缩放(扩展全局范围)
            scaled_freq = freq / scale
        else:
            # 中间频率:平滑插值
            t = (wavelength - beta_fast * original_max) / (
                (beta_slow - beta_fast) * original_max
            )
            scaled_freq = freq / (1 + (scale - 1) * t)

        freqs.append(scaled_freq)

    return freqs


# 对比不同外推方案
print("=== 位置编码外推方案对比 ===")
schemes = {
    "NTK-aware": "对所有频率统一缩放 base",
    "YaRN": "分频率差异化缩放(最优)",
    "Linear Interpolation": "所有位置等比压缩",
    "Dynamic NTK": "根据当前长度动态调整",
}
for name, desc in schemes.items():
    print(f"  {name}: {desc}")
外推方案训练长度 → 推理长度PPL 退化额外训练
Linear Interpolation4K → 32K中等需要少量微调
NTK-aware4K → 64K较小无需微调
YaRN4K → 128K+最小少量微调效果最佳
Dynamic NTK4K → 任意小无需微调

Ring Attention — 多卡分布式长上下文 ​

当单卡显存无法容纳完整 KV Cache 时,Ring Attention 将序列切分到多张 GPU,通过环形通信实现分布式注意力计算:

mermaid
graph TD
    subgraph "Ring Attention (4 GPUs)"
        GPU0["GPU 0<br/>Q[0:S/4], KV[0:S/4]"]
        GPU1["GPU 1<br/>Q[S/4:S/2], KV[S/4:S/2]"]
        GPU2["GPU 2<br/>Q[S/2:3S/4], KV[S/2:3S/4]"]
        GPU3["GPU 3<br/>Q[3S/4:S], KV[3S/4:S]"]

        GPU0 -->|"传递 KV"| GPU1
        GPU1 -->|"传递 KV"| GPU2
        GPU2 -->|"传递 KV"| GPU3
        GPU3 -->|"传递 KV"| GPU0
    end

    NOTE["每轮:各 GPU 计算本地 Q × 收到的 KV<br/>4轮后所有 GPU 都看过完整 KV"]

    style GPU0 fill:#3498db,color:#fff
    style GPU1 fill:#2ecc71,color:#fff
    style GPU2 fill:#e74c3c,color:#fff
    style GPU3 fill:#9b59b6,color:#fff
python
"""
Ring Attention 伪代码 — 理解核心思想
实际生产使用 DeepSpeed Ulysses 或 Megatron-LM 的序列并行
"""

def ring_attention_step(
    local_q,      # 本 GPU 的 Query 块
    local_kv,     # 本 GPU 的 KV 块
    world_size,   # GPU 总数
    rank,         # 当前 GPU 编号
):
    """
    Ring Attention 一轮计算:
    1. 每个 GPU 持有 Q 的一个块(固定不动)
    2. KV 块在 GPU 之间环形传递
    3. 每轮计算 local_Q × received_KV 的部分注意力
    4. world_size 轮后,每个 GPU 都计算了完整注意力
    """
    # 初始化:用本地 KV 计算第一块注意力
    partial_attn = compute_attention_block(local_q, local_kv)
    running_max = partial_attn.max_score
    running_sum = partial_attn.sum_exp
    running_output = partial_attn.output

    # 环形传递 world_size - 1 轮
    for step in range(1, world_size):
        # 异步发送本地 KV 给下一个 GPU,接收上一个 GPU 的 KV
        send_to = (rank + 1) % world_size
        recv_from = (rank - 1) % world_size
        received_kv = ring_send_recv(local_kv, send_to, recv_from)

        # 计算新的部分注意力
        new_attn = compute_attention_block(local_q, received_kv)

        # Online Softmax 合并(Flash Attention 风格)
        running_output, running_max, running_sum = online_softmax_merge(
            running_output, running_max, running_sum,
            new_attn.output, new_attn.max_score, new_attn.sum_exp,
        )

        local_kv = received_kv  # 准备下一轮传递

    return running_output


# 生产部署建议
print("""
Ring Attention 部署建议:
- 适用场景:单卡放不下完整 KV Cache(>100K tokens)
- 通信要求:GPU 间需要高带宽互联(NVLink/InfiniBand)
- 框架支持:DeepSpeed Ulysses, Megatron-LM Context Parallel
- 与 Flash Attention 兼容:Ring + Flash = 最优长上下文方案
""")

长上下文技术选型决策 ​

mermaid
graph TD
    START["需要处理长上下文"] --> Q1{"上下文长度?"}
    Q1 -->|"< 32K"| A1["标准推理<br/>Flash Attention 即可"]
    Q1 -->|"32K - 200K"| Q2{"单卡显存够?"}
    Q1 -->|"> 200K / 无限流式"| Q3{"需要全局注意力?"}

    Q2 -->|"够"| A2["YaRN 外推<br/>+ Prefix Cache"]
    Q2 -->|"不够"| A3["Ring Attention<br/>序列并行"]

    Q3 -->|"是"| A3
    Q3 -->|"否(流式/对话)"| A4["Attention Sink<br/>StreamingLLM"]

    style A1 fill:#2ecc71,color:#fff
    style A2 fill:#3498db,color:#fff
    style A3 fill:#e74c3c,color:#fff
    style A4 fill:#9b59b6,color:#fff
技术适用场景显存节省精度损失实现复杂度
Flash Attention所有场景内存 O(n²)→O(n)无低(框架内置)
Attention Sink无限流式对话固定上限丢失中间信息低
YaRN训练短推理长无极小中
Ring Attention超长文档分析线性分摊无高
Chunked Prefill长 Prompt 首次处理峰值显存降低无低(vLLM 内置)
批注模式

💬 文章评论

暂无评论,来说点什么吧 👇

编程学习笔记