长上下文优化:注意力稀疏化与 KV Cache 管理

深入长上下文优化的三条主线:注意力稀疏化(滑动窗口、注意力汇、H2O 驱逐)、KV Cache 管理(PagedAttention、前缀共享、KV 量化)与位置编码外推(RoPE scaling / YaRN),附显存计算、配置参数与评估方法。

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 个 tokenMistral 7B(W=4096)
注意力汇 + 窗口(StreamingLLM)保留前 k 个 sink + 最近 W 个StreamingLLM
块稀疏(Block-Sparse)按块划分,只算部分块Longformer、BigBird
动态稀疏(H2O)按累计注意力分数驱逐H2O
检索式(Retrieval Attention)用 ANN 近似找相关 KVRetrievalAttention

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 / GQA4~32x小
位数(dtype)KV INT8 / FP82x小
长度(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
类型显存精度影响硬件要求
FP161x基准通用
FP8 (E4M3)0.5x极小H100/Ada
INT80.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-3Hopper 异步(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 内容维度

策略思路代表
位置维度丢旧、留新、留 sinkStreamingLLM、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)位置除以 s2~4x建议
NTK-aware动态调 base4~8x建议
NTK-by-parts高频不插值、低频插值8~16x建议
YaRNNTK-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 权重)
8k2 GB~32 并发
32k8 GB~8 并发
128k32 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 而非稀疏注意力。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「llm」更多文章

  1. 端侧推理:移动端与浏览器部署
  2. RAG 评估体系:召回、忠实度与自动化指标
  3. 实时语音 Agent:全双工对话与低延迟链路