1. 长上下文的三个代价
把上下文从 8k 拉到 128k,不是"把窗口改大"那么简单,代价出现在三个维度。
| 维度 | 增长规律 | 128k 时的量级(8B 模型) |
|---|---|---|
| 计算量(prefill) | O(n²) | 8k 的 256 倍 |
| KV Cache 显存 | O(n) | 约 16 GB(FP16) |
| 精度退化 | 中段信息被忽略 | “Lost in the Middle” |
先看显存账,它是硬约束:
KV bytes per token = 2 (K 和 V) × n_layers × n_kv_heads × head_dim × dtype_bytes
Llama-3-8B: 2 × 32 层 × 8 kv_heads × 128 head_dim × 2 (FP16)
= 131072 bytes ≈ 128 KB / token
128k token → 128000 × 128KB ≈ 16 GB ← 仅 KV,不含权重
一个 8B 模型权重才 16GB,KV Cache 却能吃掉同样多。所以长上下文优化的第一战场是 KV Cache。
2. 注意力稀疏化
2.1 为什么要稀疏
注意力矩阵里绝大多数权重接近零。可视化研究显示:除了少数"注意力汇(attention sink)“token,大部分位置只关注局部邻域。稀疏化就是利用这个先验。
| 方法 | 稀疏模式 | 代表模型 |
|---|---|---|
| 滑动窗口(Sliding Window) | 只关注前 W 个 token | Mistral 7B(W=4096) |
| 注意力汇 + 窗口(StreamingLLM) | 保留前 k 个 sink + 最近 W 个 | StreamingLLM |
| 块稀疏(Block-Sparse) | 按块划分,只算部分块 | Longformer、BigBird |
| 动态稀疏(H2O) | 按累计注意力分数驱逐 | H2O |
| 检索式(Retrieval Attention) | 用 ANN 近似找相关 KV | RetrievalAttention |
2.2 滑动窗口注意力
def sliding_window_mask(seq_len: int, window: int):
"""生成滑动窗口注意力 mask。"""
import torch
idx = torch.arange(seq_len)
# |i - j| <= window 的位置可见
mask = (idx[None, :] - idx[:, None]).abs() <= window
return torch.where(mask, 0.0, float("-inf"))
代价:窗口外的信息永久丢失。Mistral 通过堆叠多层让信息"逐层传递"间接覆盖更长距离,但这不是精确回忆。
2.3 StreamingLLM 与注意力汇
关键发现:如果简单地把旧 KV 丢掉,模型输出会崩溃;但只要保留最前面 4 个 token(它们吸收了大部分注意力质量),再配上最近的窗口,就能稳定流式推理。
class StreamingKVCache:
"""注意力汇 + 滑动窗口的 KV 管理。"""
def __init__(self, num_sinks: int = 4, window: int = 4096):
self.num_sinks = num_sinks
self.window = window
self.keys, self.values = [], []
def append(self, k, v):
self.keys.append(k)
self.values.append(v)
total = len(self.keys)
if total > self.num_sinks + self.window:
# 保留前 num_sinks 个 + 最近 window 个,丢弃中间
keep = self.keys[: self.num_sinks] + self.keys[-self.window :]
self.keys = keep
self.values = self.values[: self.num_sinks] + self.values[-self.window :]
这条路线让"无限长流式对话"在固定显存下可行,代价是无法精确回忆很久之前的内容。
2.4 H2O:按重要性驱逐
H2O(Heavy-Hitter Oracle)认为:累计注意力分数高的 KV 才重要,其余可驱逐。
class H2OEviction:
def __init__(self, budget: int = 2048, recent: int = 512):
self.budget = budget # 总保留 KV 数
self.recent = recent # 最近窗口强制保留
self.scores = None # 每个位置的累计注意力
def update_scores(self, attn_weights):
# attn_weights: [heads, q_len, kv_len],对 query 维度求和
cur = attn_weights.sum(dim=-2).mean(dim=0)
if self.scores is None:
self.scores = cur
else:
self.scores[: cur.shape[0]] += cur
def select(self, kv_len: int):
keep = self.budget - self.recent
# 最近 recent 个无条件保留
recent_idx = list(range(kv_len - self.recent, kv_len))
# 其余按分数取 top
head_idx = self.scores[: kv_len - self.recent].topk(keep).indices.tolist()
return sorted(set(head_idx + recent_idx))
H2O 报告在 20% KV 预算下保持接近全量的精度,但驱逐是不可逆的——被丢掉的 KV 再也找不回来,这对需要精确引用原文的场景有风险。
3. KV Cache 的结构与显存账
3.1 为什么 KV Cache 必须存在
自回归解码时,第 t 步的注意力需要前 t-1 个位置的 K、V。若不缓存,每步都要重算全部历史,复杂度从 O(n) 变成 O(n²)。
无缓存: 生成 n 个 token 需要 O(n²) 次注意力计算
有缓存: 每步 O(n) 读取缓存,总计 O(n²) 读取但 O(n) 计算
decode 阶段是显存带宽瓶颈:每生成一个 token,都要把全部 KV 从显存读一遍。所以 KV 的体积直接决定吞吐。
3.2 三个压缩维度
| 维度 | 手段 | 压缩比 | 精度影响 |
|---|---|---|---|
| 层数(n_layers) | 跨层共享 KV(YOCO、CLA) | 2~4x | 中 |
| 头数(n_kv_heads) | MQA / GQA | 4~32x | 小 |
| 位数(dtype) | KV INT8 / FP8 | 2x | 小 |
| 长度(seq) | 驱逐 / 窗口 | 2~10x | 中~大 |
GQA(Grouped-Query Attention)是性价比最高的一项:Llama-3-8B 有 32 个 Q head 但只有 8 个 KV head,KV 直接降到 1/4,精度几乎无损。它已被几乎所有现代模型采用。
MHA: n_kv_heads = n_heads (32/32,KV 最大)
GQA: 1 < n_kv_heads < n_heads (8/32,折中,主流)
MQA: n_kv_heads = 1 (1/32,KV 最小,精度略降)
4. KV Cache 管理
4.1 PagedAttention
朴素实现的 KV Cache 需要为每个请求预分配最大长度的连续显存,导致严重碎片(实测利用率常低于 40%)。
PagedAttention 借鉴操作系统的虚拟内存分页:把 KV 切成固定大小的块(block,通常 16 个 token),用块表(block table)把逻辑位置映射到物理块。
逻辑 KV: [t0 t1 ... t15][t16 ... t31][t32 ...]
物理块: 块 #7 块 #2 块 #19
块表: [7, 2, 19, ...] ← 非连续,按需分配
收益:
| 指标 | 朴素实现 | PagedAttention |
|---|---|---|
| 显存利用率 | < 40% | > 90% |
| 前缀共享 | 不支持 | 支持(块级 COW) |
| 内存碎片 | 严重 | 无外部碎片 |
vllm serve Qwen/Qwen2.5-7B-Instruct \
--max-model-len 32768 \
--gpu-memory-utilization 0.90 \
--block-size 16
块大小是权衡点:块越大,内部碎片越多;块越小,块表开销越大。16 是社区默认值。
4.2 前缀共享(Prefix Caching)
多轮对话与 RAG 场景中,大量请求共享相同前缀(系统提示词、few-shot 示例)。前缀缓存让这些请求复用同一份 KV 块,只对增量部分计算。
# vLLM 自动启用,OpenAI 兼容接口无需改代码
resp = client.chat.completions.create(
model="Qwen/Qwen2.5-7B-Instruct",
messages=[
{"role": "system", "content": LONG_SYSTEM_PROMPT}, # 长且固定 → 命中缓存
{"role": "user", "content": user_input},
],
)
效果:系统提示词 2000 token 时,TTFT 可降 40~70%。这类"以缓存换成本"的通用手段见 /llm-cost-optimization/。
4.3 KV 量化
KV Cache 的量化与权重量化不同:KV 是运行时产生的动态张量,需要 per-token 或 per-channel 的动态 scale。
# vLLM 开启 KV INT8/FP8
# --kv-cache-dtype fp8 或 auto
| 类型 | 显存 | 精度影响 | 硬件要求 |
|---|---|---|---|
| FP16 | 1x | 基准 | 通用 |
| FP8 (E4M3) | 0.5x | 极小 | H100/Ada |
| INT8 | 0.5x | 小 | 通用 |
注意:KV 量化省的是显存,不一定省时间——若内核不支持融合反量化,反而变慢。KV Cache 的更多优化技巧见 KV Cache 优化 。
5. 注意力内核优化
5.1 FlashAttention
标准注意力要把 [n, n] 的注意力矩阵写回显存(HBM),这是真正的瓶颈。FlashAttention 用分块(tiling)+ 在线 softmax(online softmax),把中间结果留在 SRAM 里。
标准: QK^T → 写 HBM → softmax → 写 HBM → 乘 V (多次 HBM 往返)
Flash: 分块载入 SRAM → 在线 softmax 累加 → 直接出结果(HBM 只读写 O(n))
| 版本 | 关键改进 | 相对加速 |
|---|---|---|
| FlashAttention-1 | 分块 + 重计算 | 2~4x |
| FlashAttention-2 | 更好的并行划分与 warp 调度 | 再 2x |
| FlashAttention-3 | Hopper 异步(TMA + WGMMA) | 再 1.5~2x |
# 安装并验证
# pip install flash-attn --no-build-isolation
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen2.5-7B-Instruct",
attn_implementation="flash_attention_2",
torch_dtype="bfloat16",
device_map="auto",
)
注意 FlashAttention 是精确注意力,不改变数值结果,只是更快更省显存。内核实现的更多细节见 FlashAttention 内核 。
5.2 长上下文的 prefill 优化
128k 的 prefill 计算量是 8k 的 256 倍,且是算力瓶颈。优化手段:
- 分块 prefill(Chunked Prefill):把长 prompt 切成块,与 decode 请求混合批处理,避免长 prompt 阻塞在线请求。
- 序列并行(Sequence Parallelism):把序列维度切到多卡,Ring Attention 就是其代表。
- 投机 prefill:用 draft 模型预测,减少大模型的计算量。
# vLLM 开启分块 prefill,降低长 prompt 对在线请求的干扰
vllm serve Qwen/Qwen2.5-7B-Instruct --enable-chunked-prefill --max-num-batched-tokens 8192
6. 上下文压缩与选择
6.1 位置维度 vs 内容维度
| 策略 | 思路 | 代表 |
|---|---|---|
| 位置维度 | 丢旧、留新、留 sink | StreamingLLM、H2O |
| 内容维度 | 只保留与当前 query 相关的 | RetrievalAttention、Quest |
| 摘要维度 | 把旧上下文压成摘要 | 递归摘要、MemGPT |
| 表示维度 | 压缩成 latent(如 gist token) | ICAE、Gist |
6.2 递归摘要
最工程化、最通用的做法:超过阈值就把最早的对话轮次摘要成一段短文本。
async def maybe_summarize(history: list[dict], max_tokens: int = 8000):
if count_tokens(history) <= max_tokens:
return history
# 保留最近 N 轮原文
keep_recent = 4
old, recent = history[:-keep_recent], history[-keep_recent:]
summary = await llm.complete(
"把以下对话压缩成要点摘要,保留事实、数字、约定与未决问题:\n"
+ format_turns(old)
)
return [{"role": "system", "content": f"[历史摘要] {summary}"}] + recent
要点:摘要必须保留数字与专有名词,否则后续问答会丢失关键事实。上下文工程的完整方法论见 /llm-context-engineering/。
7. 位置编码外推
模型训练时的最大长度是硬约束。想在推理时突破,需要修改 RoPE 的旋转频率。
7.1 RoPE 与频率
RoPE 把位置 m 编码为旋转角度 θ_i · m,其中 θ_i = base^(-2i/d),base 默认 10000
直接外推(train 4k 推 32k)会因高频维度"绕圈"导致崩溃。
7.2 主流外推方法
| 方法 | 做法 | 扩展倍数 | 是否需要微调 |
|---|---|---|---|
| 线性插值(PI) | 位置除以 s | 2~4x | 建议 |
| NTK-aware | 动态调 base | 4~8x | 建议 |
| NTK-by-parts | 高频不插值、低频插值 | 8~16x | 建议 |
| YaRN | NTK-by-parts + 注意力温度缩放 | 16~32x | 推荐微调 |
| LongRoPE | 分维度搜索最优缩放 | 100x+ | 需要 |
# transformers 中启用 YaRN
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen2.5-7B-Instruct",
rope_scaling={
"type": "yarn",
"factor": 4.0,
"original_max_position_embeddings": 32768,
},
)
7.3 外推不等于有效
关键认知:声称支持 128k ≠ 在 128k 上有效。必须在你的任务上实测"有效上下文长度”。经典现象是 “Lost in the Middle”:把关键信息放在上下文中间,模型准确率显著低于放在开头或结尾。
8. 工程实践与评估
8.1 长上下文评估
| 基准 | 测什么 | 特点 |
|---|---|---|
| Needle in a Haystack | 在长文中找一句"针" | 位置敏感度 |
| RULER | 多任务(检索/多跳/聚合) | 比 NIAH 严格 |
| LongBench | 中文长文本任务集 | 中文场景 |
| ∞Bench | 超长(100k+)任务 | 极限测试 |
自建评估的最小方案:把你的真实文档随机插入一句可验证的事实,在多个位置(10%、50%、90%)测试召回率。
def needle_test(model, doc: str, needle: str, question: str, positions=(0.1, 0.5, 0.9)):
results = {}
for p in positions:
idx = int(len(doc) * p)
injected = doc[:idx] + f"\n{needle}\n" + doc[idx:]
answer = model.generate(f"{injected}\n\n{question}")
results[f"pos_{int(p*100)}%"] = needle in answer
return results
8.2 参数配置建议
1. 开启 FlashAttention-2:零成本加速
2. 开启 Chunked Prefill:保护在线请求的 TTFT
3. 开启 Prefix Caching:多轮/RAG 场景必开
4. 设置合理的 max-model-len:不要盲目拉满,KV 显存按需
5. 显存吃紧时:KV FP8 + GQA 模型 + 上下文压缩
6. 需要精确回忆:禁用驱逐类稀疏,改用 RAG 外挂检索
8.3 常见误区
- 盲目拉长 max-model-len:KV 显存按线性增长,会导致并发数暴跌。
- 用稀疏注意力替代 RAG:稀疏是"近似回忆",RAG 是"精确检索",长文档问答仍应优先 RAG。
- 忽略 prefill 排队:长 prompt 会把在线请求的 TTFT 顶到几秒,必须开分块 prefill。
- 外推后不做评估:声称 128k 的模型在你的数据上可能只有 16k 有效。
9. 长上下文服务的成本模型
9.1 显存换算器
部署前先算清楚"这个配置能跑多大上下文、多少并发"。KV 显存公式:
def kv_cache_gb(
n_layers: int,
n_kv_heads: int,
head_dim: int,
max_len: int,
batch: int = 1,
dtype_bytes: int = 2, # FP16=2, FP8/INT8=1
) -> float:
per_token = 2 * n_layers * n_kv_heads * head_dim * dtype_bytes
return per_token * max_len * batch / (1024 ** 3)
# Llama-3-8B,FP16,8k 上下文
print(round(kv_cache_gb(32, 8, 128, 8192), 2)) # 2.0 GB
# 同配置拉到 128k
print(round(kv_cache_gb(32, 8, 128, 131072), 2)) # 32.0 GB
# 换成 FP8 KV
print(round(kv_cache_gb(32, 8, 128, 131072, dtype_bytes=1), 2)) # 16.0 GB
注意 GQA 的 n_kv_heads=8 已经帮你省了 4 倍;如果换成 MHA(32 heads),128k 需要 128GB,单卡根本放不下。
9.2 并发数与 max-model-len 的取舍
单卡显存固定,max-model-len 与并发数此消彼长:
| max-model-len | 单请求 KV(8B, FP16) | 80GB 卡可承载(留 16GB 权重) |
|---|---|---|
| 8k | 2 GB | ~32 并发 |
| 32k | 8 GB | ~8 并发 |
| 128k | 32 GB | ~2 并发 |
结论:不要把 max-model-len 设成模型的理论上限,而应按业务实际长度分档部署——短上下文请求走高并发实例,长文档请求走专门的长上下文实例。
# 短上下文高并发实例
vllm serve Qwen/Qwen2.5-7B-Instruct --max-model-len 8192 --gpu-memory-utilization 0.92
# 长文档实例(低并发,开 KV FP8 省显存)
vllm serve Qwen/Qwen2.5-7B-Instruct \
--max-model-len 131072 --kv-cache-dtype fp8 --enable-chunked-prefill
这类"按请求特征路由到不同实例"的做法,本质是把请求按上下文长度分层,让每层都跑在最适合的并发档位上。
9.3 长上下文的 Token 成本
即使显存扛得住,Token 计费也是真实成本。把 100k token 全塞进上下文,单次调用成本可能是 RAG 方案的 20 倍以上。
方案 A(全量上下文):100k 输入 token × $3/M = $0.30 / 次
方案 B(RAG top-5): 4k 输入 token × $3/M = $0.012 / 次
+ 检索成本(可忽略)
差 25 倍
所以选型顺序应该是:先问"能不能用 RAG 缩小上下文",再问"怎么优化大上下文"。只有"必须看到全文才能回答"的任务(如跨章节推理、代码库全局重构)才值得付长上下文的成本。
9.4 混合架构
实践中最优的是混合方案:
用户请求
├─ 短问题(< 8k) → 直接进上下文
└─ 长文档 → 先检索定位相关段落(RAG)
└─ 若需全局推理 → 送入长上下文实例
└─ 配合前缀缓存复用文档 KV
关键技巧:文档部分放前缀(命中 Prefix Caching),问题放后缀(每次变化)。这样同一份文档服务多次提问时,文档的 KV 只算一次。
小结
长上下文优化可以归纳成一句话:用可接受的近似,换取显存与延迟。
三条主线各司其职:注意力稀疏化解决 prefill 计算量与 KV 体积(代价是不可逆的信息丢失),KV Cache 管理(PagedAttention + 前缀共享 + KV 量化)解决显存利用率与复用,位置编码外推解决训练长度与推理长度的鸿沟。工程上的默认组合是:FlashAttention-2 + GQA 模型 + PagedAttention + Prefix Caching,显存不足时再叠加 KV FP8 与上下文压缩,而需要精确回忆原文时始终优先 RAG 而非稀疏注意力。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。