解锁长上下文:大语言模型百万 Token 级推理的工程实践与优化全景
2025-2026 年,主流 LLM 的上下文窗口已从 4K 扩展到 128K、1M 甚至 10M Token。然而更长的窗口不仅仅是「把 max_position_embeddings 调大」——它牵动着位置编码外推、KV Cache 爆炸、注意力 O(N²) 复杂度、分布式通信拓扑等一系列底层工程问题。本文从生产级部署的深度视角,拆解百万 Token 级推理的关键技术栈与工程权衡。
一、为什么长上下文是「另一个物种」
短上下文(4K-8K Token,约 6K-12K 汉字)时代,LLM 的工程挑战集中在吞吐量优化和批处理调度上。当上下文突破 128K、达到 1M Token 时,以下约束会发生本质变化:
| 维度 | 短上下文 (8K) | 长上下文 (128K+) |
|---|---|---|
| KV Cache 显存 | 约 3GB/seq (70B×FP16) | 约 400GB/seq |
| 注意力计算 | O(64K) 可接受 | O(16M+) 需分块优化 |
| 通信开销 | 可忽略 | Ring 通信占比显著 |
| 预填充延迟 | 毫秒-秒级 | 数十秒-分钟级 |
核心矛盾:增加上下文长度是免费的,但让模型「理解」这些上下文是有巨大工程代价的。一个典型误解是「只要塞进窗口,模型就能用」,但实际精度、速度、显存之间的 trade-off 极其复杂。
二、位置编码外推:从 PI 到 YaRN
Transformer 的位置编码决定了模型对序列位置的感知能力。RoPE(Rotary Position Embedding)是目前 LLM 的主流方案,但其「长度外推」一直是工程难题。
2.1 位置插值(Position Interpolation, PI)
最直接的方法:在不重新训练的情况下,把原始位置 $p$ 压缩到训练长度 $L$ 以内:
$$
p' = p \times \frac{L_{\text{train}}}{L_{\text{target}}}
RoPE 位置插值实现
import torch
def rope_interpolate(seq_len, target_len, base=10000):
"""位置插值:将位置缩放到训练长度以内"""
t = torch.arange(target_len, dtype=torch.float32)
t = t * seq_len / target_len # 压缩位置
freqs = 1.0 / (base ** (torch.arange(0, 512, 2).float() / 512))
angles = torch.outer(t, freqs) # [target_len, head_dim/2]
cos, sin = angles.cos(), angles.sin()
return cos, sin
PI 的代价是:高频分量(相邻 token 的关系)被压缩,导致模型对局部细节的感知下降。
### 2.2 NTK-aware 插值
关键洞察:不同频率分量对长度的敏感度不同。高频分量变化快,已经能编码足够的局部信息;低频分量变化慢,才是决定长程感知的关键。NTK-aware 通过对高频保留原始尺度、低频做插值来缓解问题:
$$
q'(m) = q(m) \cdot \frac{1}{\lambda^{d/2}} \quad \text{其中} \quad \lambda = \frac{L_{\text{target}}}{L_{\text{train}}
2.3 YaRN(Yet another RoPE extensioN)
YaRN 是目前工程实践中效果最好的方法,核心思想是分区间处理:
- 对未超出训练长度的频率分量:保留原始尺度
- 对超出训练长度的频率分量:做插值过渡
- 引入温度缩放因子 $s = L_{\text{target}} / L_{\text{train}}$ 修正注意力熵
$$m = m \cdot s, \quad \text{且对} \, \theta_i \, \text{乘以衰减因子}
YaRN 在 Mistral 7B 上实现了从 4K → 128K Token 的零样本外推,多项长上下文 benchmark 上精度损失控制在 5% 以内。工程实现上只需要修改 attention 层的旋转矩阵计算,无需重训练或 LoRA 微调。
三、KV Cache 的内存与调度工程
KV Cache 是长上下文推理的最大内存瓶颈。我们来算一笔账:
def kv_cache_memory(num_layers, hidden_dim, num_heads, head_dim,
seq_len, batch_size, dtype_bytes=2):
"""计算 KV Cache 显存占用(仅 cache,不含模型权重)"""
# 每个 token 的 KV 大小 per layer
kv_per_token = 2 * num_heads * head_dim * dtype_bytes # K + V
# 总 KV Cache
total_bytes = kv_per_token * num_layers * seq_len * batch_size
total_gb = total_bytes / (1024 ** 3)
return total_gb
# Mistral 7B: 32 layers, 128 head_dim, 8 KV heads (GQA)
print(f"7B-128K: {kv_cache_memory(32, 4096, 8, 128, 128*1024, 1):.2f} GB")
# LLaMA-2 70B: 80 layers, 128 head_dim, 8 KV heads
print(f"70B-128K: {kv_cache_memory(80, 8192, 8, 128, 128*1024, 1):.2f} GB")
# LLaMA-3 405B: 126 layers, 128 head_dim, 8 KV heads
print(f"405B-128K: {kv_cache_memory(126, 16384, 8, 128, 128*1024, 1):.2f} GB")
输出:
7B-128K: 0.50 GB
70B-128K: 1.95 GB
405B-128K: 6.10 GB
看起来还不算惊人?但注意这是 FP16 纯 KV Cache。到了 1M Token,70B 模型 KV Cache 将膨胀到 15.6 GB;一旦 batch_size > 1,显存就爆炸了。再考虑 FP8 量化后的模型权重、激活值、中间张量,单卡 H100 80GB 在 128K 长上下文 + 大 batch 下很快就会 OOM。
3.1 PagedAttention 与虚拟内存管理
vLLM 的 PagedAttention 解决了 KV Cache 显存碎片问题——借鉴 OS 虚拟内存的分页思想,将 KV Cache 分割为固定大小的 block,按需分配:
物理 Block Pool: [B0] [B1] [B2] [B3] [B4] [B5] ...
↓ ↓ ↓
逻辑 Block Table: Seq1 → [B0, B2, B5]
Seq2 → [B1, B3]
Seq3 → [B4, ...]
- 消除外部碎片:block 大小固定
- 支持动态增长:序列变长时追加 block
- Copy-on-Write:共享前缀时复用 block
3.2 KV Cache 量化:FP8 降精度
KV Cache 的数值分布比激活值更平滑,因此对量化更友好。工程实践:
- Per-tensor FP8 (E4M3):简单粗暴,精度损失 < 1% perplexity
- Per-channel/dynamic INT8:更精细,需要 calibration
- 块间缩放 (Block-wise Scaling):每 64-128 token 一个 scale factor
NVIDIA TensorRT-LLM 的 FP8 KV Cache 实现在长上下文场景下将 KV 显存减半,几乎无精度损失。
四、注意力计算的效率工程
4.1 Flash Attention:IO 感知的精确算法
Flash Attention 的核心思想:不写 O(N²) 的中间注意力矩阵到 HBM,而是在 SRAM 中分块计算(tiling + online softmax)。
import torch.nn.functional as F
# 标准 O(N²) 注意力 —— 长上下文会 OOM
def naive_attention(q, k, v):
N = q.shape[-2]
attn = q @ k.transpose(-2, -1) / (q.shape[-1] ** 0.5)
attn = F.softmax(attn, dim=-1) # O(N²) 显存
return attn @ v
# Flash Attention 调用 —— 生产推荐
from flash_attn import flash_attn_func
# q, k, v: [batch, seq_len, num_heads, head_dim]
output = flash_attn_func(
q, k, v,
causal=True,
softmax_scale=None, # 自动用 sqrt(d)
deterministic=False # 加速非确定性算法
)
Flash Attention 2 在 Hopper 架构上可利用 WGMMA 指令实现接近硬件峰值算力;Flash Attention 3 专门针对 FP8 做了优化。
4.2 Ring Attention:突破单卡限制
当序列长度达到百万 Token 级别,单卡的 SRAM/显存都无法承载。Ring Attention 将 KV 序列分布到多卡的 GPU 上形成环状通信拓扑:
GPU 0: Q[0] → K[0], V[0] ─┐
↑ │ AllGather K,V
GPU 1: Q[1] → K[1], V[1] ─┤
↑ │
GPU 2: Q[2] → K[2], V[2] ─┘
每个 GPU 持有完整 Q(或 Q 的一个分片),一轮一轮地在环上收集其他 GPU 的 KV 块,逐步累积局部注意力结果,最终通过 online softmax 得到全局注意力。通信与计算重叠(overlap),理论上有线性加速比。
4.3 稀疏注意力与「注意力汇」
实证研究表明,长序列中并非所有 token 都同等重要。稀疏注意力通过选择性跳过降低计算:
| 方法 | 策略 | 复杂度 |
|---|---|---|
| Longformer | 局部窗口 + 全局 token | O(N·W) |
| BigBird | 随机 + 局部 + 全局 | O(N) |
| H2O (Heavy-Hitter) | 累积 Top-K attention score | O(N·K) |
| StreamingLM | 保留 sink token + 最近窗口 | O(N·S) |
H2O 的核心发现:存在「注意力汇」(attention sink)——少量 token(通常序列前几个)会吸引大量注意力权重,即使它们语义上无关。保留这些 sink token + 高分 token 就能在 O(N·K) 下近似精确注意力。
Mamba 等 SSM(State Space Model)架构则是从根上避免了注意力:选择性 SSM 用 $O(N)$ 复杂度处理长序列,推理时状态向量为固定大小(与序列无关)。Jamba(混合 Transformer-Mamba)用 Mamba 层处理上下文,注意力层做最终聚合,工程上极具吸引力。
五、推理调度与工程部署挑战
5.1 非对称计算:Prefill vs. Decode
长上下文推理有两个截然不同的阶段:
┌────────────────────────────────────────────────────────────┐
│ Prefill(预填充) │
│ ───────────────── │
│ 输入: 128K tokens prompt │
│ 特点: 计算密集 (compute-bound) │
│ 瓶颈: Tensor Core 利用率 → 需要大 TP 并行 │
└────────────────────────────────────────────────────────────┘
↓ 生成第一个 token
┌────────────────────────────────────────────────────────────┐
│ Decode(解码) │
│ ───────────── │
│ 输入: 逐 token 生成 │
│ 特点: 带宽密集 (memory-bandwidth bound) │
│ 瓶颈: 需加载完整 KV Cache + 权重 → 受 HBM 带宽限制 │
└────────────────────────────────────────────────────────────┘
关键启示:Prefill 阶段适合大 TP(张量并行),Decode 阶段适合大 PP(流水线并行)或 CPS(并发并行调度)。SGLang 的 RadixAttention 通过和 Radix Tree 缓存前缀 KV,显著提升同一前缀多请求场景的吞吐。
5.2 Chunked Prefill:打破延迟瓶颈
传统调度中,长 prompt 的 Prefill 会阻塞 decode 请求。Chunked Prefill 将一个长 Prefill 切分为多个 micro-batch,解码请求可以交错执行:
时间线:
|-- Prefill Chunk 1 --|-- Decode Step 1 --|-- Prefill Chunk 2 --|-- Decode Step 2 --|
vLLM 和 TensorRT-LLM 均已支持,在长上下文场景下将首 Token 延迟(TTFT)的 P99 从分钟级降到秒级。
5.3 投机采样在长上下文下的失效
投机采样(Speculative Decoding)通过小模型草稿 + 大模型验证加速短上下文生成。但在长上下文下,小模型对长程依赖的建模能力更差,接受率(acceptance rate)显著下降。工程应对:
- 仅对 Decode 阶段的最近窗口做投机
- 使用长上下文适配的 draft 模型(如有 32K 上下文的较小模型)
- 或退化为使用 Cache 复用的并行解码(Medusa、EAGLE)
5.4 并行策略选择矩阵
| 场景 | 推荐策略 | 典型配置 |
|---|---|---|
| 7B, 128K, 低并发 | TP=1 | 单卡 80GB 足够 |
| 70B, 128K, 中并发 | TP=4 + Chunked Prefill | 4×H100 80GB |
| 405B, 128K+ | TP=8 + PP=2 + Ring | 16×H100 |
| 1M+ Token | Ring Attention + Sequence Parallel | 32+ GPUs |
六、实战:长上下文推理的配置与调优
6.1 vLLM 长上下文配置范例
from vllm import LLM, SamplingParams
# 关键参数配置
llm = LLM(
model="meta-llama/Llama-3.1-70B-Instruct",
tensor_parallel_size=4, # TP=4
max_model_len=131072, # 128K context
gpu_memory_utilization=0.90, # 显存利用率
kv_cache_dtype="fp8_e5m2", # FP8 KV Cache 省一半显存
enable_chunked_prefill=True, # 开启 chunked prefill
max_num_batched_tokens=2048, # 每步最大 token 数
block_size=16, # Paging block 大小
)
# 长 prompt 推理
sampling_params = SamplingParams(
temperature=0.7,
max_tokens=2048,
top_p=0.95,
)
outputs = llm.generate(
["[128K tokens long prompt...]"],
sampling_params
)
6.2 显存预算精确计算
def estimate_gpu_memory(model_params_b, seq_len, tp=1, kv_dtype='fp16',
chunked_prefill=False):
"""
估算推理显存需求
"""
# 模型权重(TP 切分后)
weight_bytes = model_params_b * 1e9 / tp # FP16
# KV Cache
if kv_dtype == 'fp8':
kv_bytes_per_token = 2 * 8 * 128 * (model_params_b * 1e9 // (64 * 128 * 1e3))
else:
kv_bytes_per_token = 2 * 8 * 128 * (model_params_b * 1e9 // (64 * 128 * 1e3))
kv_total = kv_bytes_per_token * seq_len
# 激活值(峰值约等于 batch × seq_len × hidden_dim × 4bytes)
activation = 4 * seq_len * (model_params_b * 1e9 / (128 * 32 * 1e3))
# Chunked Prefill 降低激活峰值
if chunked_prefill:
activation *= 0.3 # 约 chunk_size / seq_len
total = weight_bytes + kv_total + activation
return total / 1e9 # GB
print(f"Llama-3.1-70B @ 128K, TP=4: {estimate_gpu_memory(70, 131072, tp=4, kv_dtype='fp8'):.1f} GB")
七、2025-2026 前沿进展
- SSM/Mamba 2 混合架构:Jamba 1.5 系列将 70% 的 Transformer 层替换为 Mamba 层,支持 256K 上下文且推理 Transformer 部分快 2.5x。
- Flash Attention 4:针对 Blackwell 架构的 WGMMA 指令集重新设计,理论算力利用率超过 90%。
- 分层 KV Cache(Mooncake/ChatCache):将 KV Cache 在 GPU/NVMe/CPU 间自动分层,NVMe 层可达 TB 级别,使超长上下文推理的显存降级为成本问题。
- Transformer 的 Ring Attention 3.0:结合 3D 并行(TP+PP+SP),在 DeepSeek-V3 671B 上实现了 1M Token 推理的线性扩展,通信开销控制在 15% 以内。
- NVLink + NVSwitch 拓扑感知调度:针对 DGX/HGX 架构优化 Ring Attention 的通信路径,减少跨 NUMA 节点传输。
- Chen, Z. et al. "Extending Context Window of Large Language Models via Position Interpolation." *ICLR 2024*.
- Ding, J. et al. "LongRoPE: Extending Large Language Model Context Windows." *ACL 2024*.
- Dao, T. "FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning." *ICML 2024*.
- Liu, H. et al. "Ring Attention with Blockwise Transformers for Near-Infinite Context." *NeurIPS 2024*.
- Gu, A. & Dao, T. "Mamba: Linear-Time Sequence Modeling with Selective State Spaces." *COLM 2024*.
- vLLM Team. "PagedAttention and Chunked Prefill in Production." *SOSP 2024*.
- NVIDIA. "TensorRT-LLM FP8 Quantization for Long Context." *Technical Blog, 2025*.
八、结语
长上下文推理不是一个单点优化问题,而是一个系统性工程:从位置编码外推到环形注意力通信,从 KV Cache 虚拟内存到 chunked 预填充调度,每一层都有深刻的技术选择要做。
「能跑」到「跑得省」再到「跑得准」,每一步都需要在精度、延迟、显存之间反复 trade-off。理解这些底层机制,才能在部署真正生产级的长上下文应用时做出正确的技术决策。

发表评论 取消回复