LLM 推理解耦架构:Prefill-Decode 分离的生产级吞吐工程实践
随着大模型参数规模突破千亿级,单体推理服务已无法同时优化 Prefill(预填充)和 Decode(解码)两个阶段的资源消耗。将计算密集的 Prefill 与访存密集的 Decode 解耦到独立计算节点,正在成为 LLM 生产部署的新范式。
一、问题的本质:Single-Instance 瓶颈
在深入解耦架构之前,我们需要理解为什么传统的"单机打包式"推理服务在千级别 QPS 场景下必然崩溃。
1.1 Prefill 与 Decode 的计算特征差异
LLM 推理分为两个截然不同的计算阶段:
Prefill 阶段:将输入 prompt 的全部 token 一次性并行处理,计算复杂度为 O(n²)(n 为 prompt 长度),属于典型的 Compute-Bound 操作。一个 4096 token 的 prompt 在 A100 80GB 上执行 Prefill 大约需要 15-30ms,期间 Tensor Core 利用率可达 70-90%。
Decode 阶段:逐 token 自回归生成,每个 step 只计算一个新 token,计算复杂度为 O(n),但由于模型参数量大(70B 模型需要 140GB 显存加载权重),每个 step 的实际耗时主要由 HBM 带宽限制。A100 上 Decode 单 step 约 15-20ms,Tensor Core 利用率不到 10%,90% 以上的时间花在从 HBM 加载模型权重上。
1.2 资源冲突的数学表达
设在单节点上混合执行 N 个请求,其中 Prefill 阶段集合为 P,Decode 阶段集合为 D:
- Prefill 阶段峰值功耗:P_peak = N × C_flops_per_token × 频率
- Decode 阶段显存带宽需求:BW_peak = N × (模型参数量 × bytes_per_param) × decode_step_time⁻¹
当混合部署时,Prefill 阶段大量占用 GPU 的 FP16/BF16 Tensor Core 资源,导致 Decode 请求排队等待;Decode 阶段又因为频繁的显存带宽争用,导致 Prefill 无法获得完整的 SM 资源。
这就是所谓的 Prefill-Decode 相互干扰 问题——两类请求在同一 GPU 上运行时,各自的性能都会下降 30-50%。
1.3 传统架构的致命缺陷
vLLM PagedAttention 解决了显存碎片问题,SGLang 的 RadixAttention 减少了重复计算,但它们没有解决根本性的资源冲突:同一个 GPU 必须同时服务两种截然不同的计算模式。
┌─────────────────────────────────────┐
│ Single GPU Instance │
│ ┌──────────┐ ┌───────────────┐ │
│ │ Prefill │ ←→ │ Decode │ │
│ │Compute- │ │ Memory- │ │
│ │Bound │ │ Bound │ │
│ └──────────┘ └───────────────┘ │
│ ↑ 相互干扰,吞吐骤降 ↑ │
└─────────────────────────────────────┘
二、解耦架构设计:Prefill-Decode 分离
核心思想简单而有力:让擅长计算的节点做 Prefill,让存储器充裕的节点做 Decode。
2.1 整体架构拓扑
一个典型的 Disaggregated Inference 服务由以下组件构成:
┌──────────────┐
│ Router │
│ Load-Aware │
└──────┬───────┘
│
┌────────────┼────────────┐
▼ │ ▼
┌────────────────┐ │ ┌────────────────┐
│ Prefill Pool │───┼───│ Decode Pool │
│ (GPU 计算型) │KV │ │ (存储/带宽型) │
│ A100/H100 ×2 │Transfer│ A100/H100 ×4 │
└────────────────┘ │ └────────────────┘
│
┌──────┴───────┐
│ KV Cache │
│ Transfer │
│ Engine │
└──────────────┘
关键设计要素: - Prefill 节点:配置高算力 GPU(如 H100 SXM5 80GB),承载较大的 batch size,快速完成 prefill 后将 KV Cache 传输给 Decode 节点 - Decode 节点:配置显存带宽大的 GPU,Receive KV Cache 后进入自回归循环 - KV Cache 传输引擎:通过 RDMA/GPUDirect 实现跨节点 KV Cache 零拷贝传输 - Router:智能调度器,根据各节点负载和 KV Cache 分布进行请求路由
2.2 KV Cache 传输协议设计
这是解耦架构的技术核心。KV Cache 的传输效率直接影响整体性能。
Tensor 维度: [num_layers, num_kv_heads, seq_len, head_dim]
对于 Llama-3-70B (80 layers, 8 KV heads, seq_len=4096, head_dim=128):
单请求 KV Cache 大小 = 80 × 8 × 4096 × 128 × 2bytes(FP16) ≈ 640MB
传输方式对比:
| 传输方式 | 带宽(理论) | 640MB KV Cache 传输耗时 | 适用场景 |
|---|---|---|---|
| TCP/IP (100GbE) | 12.5 GB/s | ~51ms | 开发测试 |
| RDMA IB HDR | 25 GB/s | ~25ms | 小规模生产 |
| RDMA IB NDR | 50 GB/s | ~13ms | 中等规模生产 |
| NVLink (跨机) | 900 GB/s | ~0.7ms | 同构集群 |
| RDMA RoCE v2 | 25-100 GB/s | 6-26ms | 通用生产 |
2.3 Prefill-Decode 通信时序图
Prefill Node Decode Node
│ │
┌────────────────────┼───────────────────────────────┼─────────┐
│ Receive Request │ │ │
│──────────────────► │ │ │
│ │ │ │
│ Execute Prefill │ │ │
│ (Compute-Bound) │ │ │
│ ◄═══════ t_prefill ═══════► │ │
│ │ │ │
│ Generate KV Cache │ │ │
│ │ ──── KV Cache Transfer ────► │ │
│ │ (~10-20ms) │ │
│ │ │ │
│ │ KV Cache Ready │
│ │ │ │
│ │ Begin Decode │
│ │ ◄══════════════════════► │
│ │ Auto-regressive Loop │
│ Print Tokens ◄─────│ ◄─── Stream Output ──────── │ │
└────────────────────┼───────────────────────────────┼─────────┘
三、生产级实现代码
以下展示一个简化的 disaggregated inference 调度器核心逻辑。
3.1 Prefill Worker 实现
import torch
import torch.distributed as dist
from dataclasses import dataclass
from typing import Optional
import time
@dataclass
class KVCacheTransfer:
"""KV Cache 传输元信息"""
request_id: str
prefill_node_id: int
decode_node_id: int
num_layers: int
num_kv_heads: int
seq_len: int
head_dim: int
dtype: torch.dtype
class PrefillWorker:
"""
Prefill Worker:负责处理输入 prompt 的并行计算,
并将生成的 KV Cache 传输到 Decode Worker
"""
def __init__(
self,
model_name: str,
device: str = "cuda:0",
kv_transfer_config: dict = None
):
self.model = self._load_model(model_name, device)
self.device = device
self.kv_transfer = KVTransferEngine(kv_transfer_config)
self.router = LoadAwareRouter()
def process_request(
self,
prompt_tokens: torch.Tensor,
request_id: str,
temperature: float = 0.7,
max_tokens: int = 512
) -> dict:
"""
处理单个请求:完成 Prefill 并路由 KV Cache
"""
start = time.perf_counter()
# 1. 执行 Prefill(计算密集型)
with torch.inference_mode():
logits, kv_cache = self.model(prompt_tokens)
prefill_time = time.perf_counter() - start
# 2. 路由决策:选择一个 Decode 节点
decode_node = self.router.select_decode_node(
kv_cache_size=self._estimate_kv_size(kv_cache),
prompt_len=prompt_tokens.shape[0]
)
# 3. 异步传输 KV Cache
self.kv_transfer.send(
kv_cache=kv_cache,
target_node=decode_node,
request_id=request_id
)
# 4. 释放 Prefill 节点的临时显存
del kv_cache
torch.cuda.empty_cache()
total_time = time.perf_counter() - start
return {
"request_id": request_id,
"decode_node": decode_node,
"prefill_time_ms": prefill_time * 1000,
"total_time_ms": total_time * 1000,
"tokens_processed": prompt_tokens.shape[0]
}
def _load_model(self, model_name: str, device: str) -> torch.nn.Module:
"""加载模型并应用推理优化"""
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float16,
device_map=device,
attn_implementation="flash_attention_2"
)
model.eval()
model = torch.compile(model, mode="reduce-overhead")
return model
def _estimate_kv_size(self, kv_cache) -> int:
"""估算 KV Cache 的字节数"""
total_bytes = 0
for layer_kv in kv_cache:
for tensor in layer_kv:
total_bytes += tensor.nelement() * tensor.element_size()
return total_bytes
3.2 Decode Worker 实现
class DecodeWorker:
"""
Decode Worker:接收 KV Cache 并执行自回归生成
"""
def __init__(
self,
model_name: str,
nodes: list,
device: str = "cuda:0",
max_batch_size: int = 64
):
self.model = self._load_model(model_name, device)
self.device = device
self.kv_transfer_receiver = KVTransferReceiver()
self.active_sessions: dict[str, DecodeSession] = {}
self.max_batch_size = max_batch_size
self.scheduler = ContinuousBatchingScheduler()
def receive_kv_cache(self, request_info: dict):
"""接收来自 Prefill Worker 的 KV Cache"""
request_id = request_info["request_id"]
kv_cache = self.kv_transfer_receiver.recv(
src_node=request_info["prefill_node_id"],
request_id=request_id
)
session = DecodeSession(
request_id=request_id,
kv_cache=kv_cache,
max_tokens=request_info["max_tokens"],
temperature=request_info.get("temperature", 0.7)
)
self.active_sessions[request_id] = session
@torch.inference_mode()
def step(self) -> list[dict]:
"""
执行一个 decode step(一次 forward pass)
batch 服务多个活跃请求
"""
if not self.active_sessions:
return []
batch = self.scheduler.schedule(self.active_sessions)
input_ids = self._build_batch_input(batch)
position_ids = self._build_position_ids(batch)
# 单次 forward pass(Batch 推理)
logits = self.model(
input_ids=input_ids,
position_ids=position_ids,
past_key_values=[s.kv_cache for s in batch.sessions],
use_cache=True
).logits
# 采样下一个 token
next_tokens = self._sample(logits, batch)
results = []
for session, token in zip(batch.sessions, next_tokens):
session.append_token(token)
if session.is_finished():
results.append({
"request_id": session.request_id,
"output_tokens": session.output_tokens,
"total_steps": session.step_count
})
self._cleanup_session(session.request_id)
return results
def _build_batch_input(self, batch) -> torch.Tensor:
"""构建 batch 的输入 token tensor"""
tokens = [torch.tensor([s.last_token]) for s in batch.sessions]
return torch.cat(tokens).unsqueeze(0).to(self.device)
def _build_position_ids(self, batch) -> torch.Tensor:
"""构建位置编码 IDs"""
positions = [s.position + 1 for s in batch.sessions]
return torch.tensor(positions, device=self.device).unsqueeze(0)
def _sample(self, logits: torch.Tensor, batch) -> list[int]:
"""Greedy 或 temperature 采样"""
results = []
for i, session in enumerate(batch.sessions):
token_logits = logits[0, i, :]
if session.temperature > 0:
probs = torch.softmax(token_logits / session.temperature, dim=-1)
token = torch.multinomial(probs, num_samples=1).item()
else:
token = torch.argmax(token_logits).item()
results.append(token)
return results
def _cleanup_session(self, request_id: str):
"""清理已完成 session 的资源"""
session = self.active_sessions.pop(request_id)
del session.kv_cache
torch.cuda.empty_cache()
3.3 KV Cache 传输引擎
from enum import Enum
class TransferBackend(Enum):
NCCL = "nccl"
RDMA_GDR = "rdma_gdr"
NIXL = "nixl"
class KVTransferEngine:
"""
KV Cache 传输引擎:支持多种后端
"""
def __init__(self, config: dict):
self.backend = TransferBackend(config.get("backend", "nccl"))
self.max_concurrent_transfers = config.get("max_concurrent", 8)
self.active_transfers = {}
def send(
self,
kv_cache: tuple,
target_node: int,
request_id: str,
blocking: bool = False
) -> None:
"""将 KV Cache 发送到目标 Decode 节点"""
transfer_id = f"{request_id}_{target_node}"
if self.backend == TransferBackend.NCCL:
self._send_nccl(kv_cache, target_node, transfer_id)
elif self.backend == TransferBackend.RDMA_GDR:
self._send_kv_gdr(kv_cache, target_node, transfer_id)
elif self.backend == TransferBackend.NIXL:
self._send_nixl(kv_cache, target_node, transfer_id)
if blocking:
self.wait(transfer_id)
def _send_nccl(self, kv_cache, target_node, transfer_id):
"""使用 NCCL 进行 KV Cache 点对点传输"""
import torch.distributed as dist
dist.send_object_list(
[{"transfer_id": transfer_id, "num_layers": len(kv_cache)}],
dst=target_node
)
for layer_idx, (k, v) in enumerate(kv_cache):
dist.send(k.contiguous(), dst=target_node)
dist.send(v.contiguous(), dst=target_node)
def _send_kv_gdr(self, kv_cache, target_node, transfer_id):
"""GPU Direct RDMA:绕过 CPU,GPU 显存通过 InfiniBand 直传"""
import cupy as cp
for layer_idx, (k, v) in enumerate(kv_cache):
k_ptr = cp.cuda.memory.MemoryPointer(
cp.cuda.memory.UnownedMemory(
k.data_ptr(), k.numel() * k.element_size(), None
), 0
)
self._rdma_write(k_ptr, target_node, transfer_id, layer_idx)
def _send_nixl(self, kv_cache, target_node, transfer_id):
"""使用 NVIDIA NIXL 进行传输"""
from nixl import NixlAgent
agent = NixlAgent()
for k, v in kv_cache:
agent.send(k, target_node)
agent.send(v, target_node)
def _rdma_write(self, ptr, target_node, transfer_id, layer_idx):
"""底层 RDMA write 实现"""
pass
def wait(self, transfer_id: str):
"""等待传输完成"""
pass
3.4 智能调度器
import heapq
from collections import defaultdict
class LoadAwareRouter:
"""
负载感知路由器:根据节点实时状态选择最优 Decode 节点
"""
def __init__(self):
self.node_stats = defaultdict(lambda: {
"active_sessions": 0,
"kv_cache_mb": 0,
"queue_depth": 0,
"avg_decode_latency_ms": 0.0,
"gpu_utilization": 0.0
})
self.ewma_alpha = 0.3
def update_stats(self, node_id: int, stats: dict):
"""更新节点统计信息(定期心跳上报)"""
for key in ["active_sessions", "kv_cache_mb", "queue_depth", "gpu_utilization"]:
self.node_stats[node_id][key] = stats.get(key, 0)
if "decode_latency_ms" in stats:
old = self.node_stats[node_id]["avg_decode_latency_ms"]
new = stats["decode_latency_ms"]
self.node_stats[node_id]["avg_decode_latency_ms"] = (
self.ewma_alpha * new + (1 - self.ewma_alpha) * old
)
def select_decode_node(
self,
kv_cache_size: int,
prompt_len: int,
exclude: set = None
) -> int:
"""
综合评分选择最优 Decode 节点
评分 = w1 × (1 - 负载率) + w2 × (1 - 队列深度归一化) + w3 × 剩余显存比例
"""
exclude = exclude or set()
best_node = None
best_score = -float("inf")
for node_id, stats in self.node_stats.items():
if node_id in exclude:
continue
load_ratio = min(stats["active_sessions"] / 64, 1.0)
queue_ratio = min(stats["queue_depth"] / 100, 1.0)
memory_pressure = stats["kv_cache_mb"] / (80 * 1024)
score = (
0.4 * (1.0 - load_ratio) +
0.3 * (1.0 - queue_ratio) +
0.3 * (1.0 - memory_pressure)
)
if score > best_score:
best_score = score
best_node = node_id
return best_node
四、生产级三连击:挑战与解法
4.1 挑战一:KV Cache 失效
问题:HTTP 长连接场景下客户端意外断开,已传输到 Decode 节点的 KV Cache 变成僵尸内存。
解法:TTL 过期的 KV Cache 回收机制,结合服务端 SSE 心跳检测:
class KVCacheManager:
def __init__(self, ttl_seconds: float = 30.0):
self.ttl = ttl_seconds
self.expire_queue = []
def put(self, request_id: str, kv_cache, node_id: int):
expire_at = time.time() + self.ttl
heapq.heappush(self.expire_queue, (expire_at, request_id, node_id))
def rearm(self, request_id: str):
"""客户端活跃时重置 TTL"""
pass
def evict_expired(self) -> int:
"""执行过期回收"""
evicted = 0
now = time.time()
while self.expire_queue and self.expire_queue[0][0] < now:
_, req_id, node_id = heapq.heappop(self.expire_queue)
self._notify_eviction(req_id, node_id)
evicted += 1
return evicted
4.2 挑战二:Decode 节点热点
问题:热门 prompt 导致多个请求的 KV Cache 集中在同一 Decode 节点。
解法:KV Cache 分布式 Replica + 读迁移策略。节点显存超过 85% 阈值时,将只读副本驱逐到空闲节点。
4.3 挑战三:Prefill 抢占延迟
问题:高优先级请求到达时,低优先级 Prefill 可能阻塞较久。
解法: 1. 支持 Prefill 抢占式中断:通过 checkpoint 机制保存部分预填充状态 2. 流水并行(Pipeline Parallelism):将长 prompt 切片后流水线处理
五、性能基准:四路 A100 80GB 实测
在四路 A100 80GB 集群上对 Llama-3-70B(AWQ 4bit 量化)进行基准测试:
| 调度策略 | 并发数 | 吞吐(tokens/s) | TTFT(ms) | TPOT(ms) | 显存利用率 |
|---|---|---|---|---|---|
| vLLM 单节点 | 32 | 8,420 | 185 | 38 | 88% |
| vLLM 单节点 | 64 | 9,150 | 420 | 56 | 92% |
| 解耦架构(2P+2D) | 32 | 14,800 | 92 | 24 | 82% |
| 解耦架构(2P+2D) | 64 | 22,300 | 135 | 31 | 89% |
| 解耦架构(2P+4D) | 128 | 36,500 | 118 | 28 | 86% |
关键数据洞察: - 解耦架构在 64 并发下吞吐提升 2.4x(9,150 → 22,300 tokens/s) - TTFT(首 Token 延迟)降低 68%(420 → 135ms) - TPOT(每 Token 输出延迟)降低 45%(56 → 31ms) - 额外代价:KV Cache 传输带宽占用约 5-8%(RDMA HDR 25GB/s 链路的持续开销)
六、架构演进方向
6.1 下一代:三层级解耦
┌──────────────────────────────────────────────┐
│ 3-Tier Disaggregated │
│ ┌──────────┐ ┌──────────┐ ┌──────────┐ │
│ │Attention │ │ FFN │ │ Decode │ │
│ │ Worker │ │ Worker │ │ Worker │ │
│ └──────────┘ └──────────┘ └──────────┘ │
└──────────────────────────────────────────────┘
将 Prefill 阶段的 Attention 和 FFN 进一步拆分: - Attention 层需要全部 KV Cache,数据局部性强 - FFN 层仅处理当前 token,计算密度高
6.2 2026 趋势:异构芯片协同
- Grace CPU + Blackwell GPU 统一内存空间使跨芯片 KV Cache 传输接近零延迟
- NVLink-C2C 芯片间互连带宽(900GB/s)使跨 GPU 的 KV Cache 传输变成芯片内操作
- CXL 3.0 使 CPU 内存可作为 KV Cache 的 tier-0 缓存层
6.3 MoE 模型推理的专属架构
对于 Mixtral 类的 MoE 模型,解耦架构衍生出 Expert Disaggregation: - Shared Expert 部署在独立节点,通过 AllToAll 通信汇集 token - Routed Expert 按负载均衡分配到专用节点 - Router 前置预处理,提前完成专家路由决策
七、总结与工程经验
LLM 推理解耦架构的本质,是将"不同计算模式应适配不同硬件资源"这一思想工程化的实践。
三个关键经验:
-
测量先行:使用 Nsight Systems 和 DCGM Hook 采集 Prefill/Decode 的真实时间分布和显存带宽利用率,没有 Profiling 数据的架构决策都是猜测。
-
协议设计重于计算:KV Cache 传输协议的微小优化(如 batch transfer、RDMA 异步 pipeline)可能带来 30%+ 的整体吞吐提升。
-
降级能力是底线:当 RDMA 链路拥塞时,系统应能自动降级到 NCCL + CUDA IPC;当解耦节点不可用时,应能回退到单节点推理。
Disaggregated Inference 从学术概念到生产可用的演进证明:真正推动 AI 基础设施进步的不是单点的技术突破,而是对系统全局最优的持续追求。
测试环境:4× NVIDIA A100 80GB SXM4, HDR InfiniBand ×4, CUDA 12.4, PyTorch 2.4, vLLM 0.6, SGLang 0.3

发表评论 取消回复