27届大模型面试准备(七十四):大模型长上下文推理与服务化工程------分块 Prefill、KV 卸载与 Ring Attention
引言
A13 讲了长上下文的"原理与位置编码",A29 讲了高效注意力,A69 讲了 PD 分离。本篇把这三块在工程服务化层面串起来:当用户输入 32K、128K 甚至百万 token 时,推理服务怎么不爆显存、不卡首包、还能扛住吞吐。这是长文档问答、代码仓库级理解、4MRAG 长多跳推理的必过工程关。
和 A53 多模态服务化、A63 推理缓存是一伙的------缓存解决重复前缀,本篇解决"超长单次上下文"本身的显存与计算瓶颈。华为多模态 LLM 岗常考"长视频/长文档怎么喂"。
前文链接:A13 长上下文与长度外推、A29 长文本推理与高效注意力、A69 Prefill-Decode 分离、A53 多模态大模型推理服务化。
长上下文 serving 数据流(单请求 128K)
┌──────────┐ ┌──────────────┐ ┌──────────────────────┐
│ 请求 128K │──▶│ Chunked Prefill│──▶│ KV Cache 管理 │
│ token │ │ 分块进(8K/块) │ │ (GPU/CPU 分层 + 量化) │
└──────────┘ └──────────────┘ └──────────────────────┘
│
┌─────────────┼──────────────┐
▼ ▼ ▼
GPU HBM KV CPU 内存 KV Sequence 并行分片
(热层/近期) (冷层/远端) (Ring Attention)
│ │ │
└─────▶ Decode 逐 token ◀───┘
表:长上下文主要优化手段对比
| 手段 | 解决什么 | 代价 | 与谁配合 |
|---|---|---|---|
| Chunked Prefill | 长 prefill 阻塞 decode、首包慢 | 实现复杂、需调度 | 连续批处理 |
| KV 量化(FP8/INT4) | KV 显存随 seq 线性爆 | 轻微精度损 | 任何长上下文 |
| KV 卸载(CPU/NVMe) | 超长序列显存不够 | 取回有带宽延迟 | 分层_cache |
| Sequence/Ring Attention | 单卡放不下整序列注意力 | 通信开销 | 张量并行 |
| 稀疏/滑动窗口 | 并非所有 token 都需全注意 | 可能丢长程依赖 | 局部+全局混合 |
一、长上下文的成本来源
KV Cache 随序列长度 L 线性增长:每 token 每层的 KV 占用 = 2 × L × n_head × d_head × 2(byte, BF16)。以 7B 模型、L=128K,KV 约 2×128000×32×128×2 ≈ 2GB/请求。多并发下显存迅速吃满。Prefill 阶段对 L 做全注意力,计算量 O(L²),长序列首包延迟陡增。
二、分块 Prefill(Chunked Prefill)
把长 prefill 切成小块(如 8K)与 decode 请求穿插进同一 batch,避免一次长 prefill 饿死所有 decode、拖垮 TTFT。
# 伪代码:调度器把 prefill 切成 chunk,与 decode 同 batch 执行
def schedule_step(waiting, running, chunk_size=8192):
batch = []
for req in running: # decode 请求优先保留
batch.append(req.next_decode_step())
for req in waiting: # prefill 请求按 chunk 喂
if req.prefill_done:
continue
n = min(chunk_size, req.uncomputed)
batch.append(req.prefill_chunk(n)) # 只算前 n 个 token 的注意力
req.uncomputed -= n
if req.uncomputed == 0:
req.prefill_done = True
return run_batch(batch) # 一次前向覆盖 decode+prefill 片段
关键:chunk 内注意力需做因果掩码,且跨 chunk 的 KV 要正确拼接进 KV Cache,不能重复计算。vLLM/SGLang 均内置该机制。
三、KV Cache 量化与卸载
显存放不下时,先量化(FP8/INT4 压缩 KV),仍不够就分层卸载到 CPU 内存甚至 NVMe,按需取回。
import torch
def quantize_kv(k, bits=8):
# per-token 缩放的 INT8 量化,省 ~50% 显存
qmin, qmax = -128, 127
s = k.abs().amax(-1, keepdim=True) / 127.0 + 1e-6
kq = (k / s).round().clamp(qmin, qmax).to(torch.int8)
return kq, s # 推理时反量化为 BF16 再算注意力
class LayeredKVCache:
def __init__(self):
self.gpu = {} # 热点层/近期 token 留 GPU
self.cpu = {} # 冷层/远端 token 卸到 CPU
def get(self, layer, idx):
if idx in self.gpu.get(layer, {}):
return self.gpu[layer][idx]
t = self.cpu[layer].pop(idx).to("cuda") # 取回,有 PCIe 带宽代价
self.gpu[layer][idx] = t
return t
取舍:量化几乎无损且零延迟;卸载省显存但以取回延迟换空间,适合"超长但大部分 token 是背景"的场景(如长文档 RAG 仅少数段落被反复注意)。
四、序列并行与 Ring Attention
单卡 HBM 放不下整层 KV 时,把序列维度切到多卡,注意力用 Ring Attention:每张卡持一段 KV,通过 ring 通信轮转传递 Q/K/V,分块算 softmax 的全局归约。
# Ring Attention 核心思想(伪代码,2 卡示例)
def ring_attention(q_shard, k_shards, v_shards, world):
out = 0; m = -inf; l = 0
for step in range(world): # 每轮收到一卡的分片 KV
k, v = k_shards[step], v_shards[step]
s = q_shard @ k.T / sqrt(d) # 局部 logits
m_new = max(m, rowmax(s)); # 在线 softmax 归约
p = exp(s - m_new)
l = exp(m - m_new) * l + p.sum(-1)
out = exp(m - m_new) * out + (p @ v)
m = m_new
k_shards = rotate(k_shards); v_shards = rotate(v_shards) # 环形传递
return out / l.unsqueeze(-1)
Ring Attention 把显存峰值从 O(L) 降到 O(L/world),代价是 world−1 次 all-to-all 通信。配合 Ulysses 序列并行可进一步降通信。
五、工程权衡清单
- 显存 vs 延迟:量化优先(零延迟),卸载兜底(有延迟),序列并行解决单卡放不下。
- 精度 vs 长度:长上下文下注意力数值范围大,softmax 用 FlashAttention 稳定;KV 量化选 per-token 缩放减误差。
- 吞吐:Chunked Prefill + 连续批处理让长/短请求混跑,GPU 利用率最高;纯长请求反而该限制并发保 TTFT。
- 与 4MRAG 结合:长多跳推理不必一次性把全文档塞上下文------用检索把相关片段动态注入(A55/A63 缓存命中),只在真正需要长程时展开全序列,省显存。
面试速答
- 长上下文推理显存主要耗在哪?KV Cache 随序列长度线性增长,每 token 每层都要存 K/V。
- Chunked Prefill 解决什么?长 prefill O(L²) 计算阻塞 decode 导致首包慢,切块后与 decode 混批,平滑 TTFT、提升利用率。
- KV 量化 vs 卸载怎么选?量化几乎零延迟优先用;单卡仍放不下再分层卸载到 CPU/NVMe,以取回延迟换空间。
- Ring Attention 干嘛用?把序列切到多卡,环形传递 KV 分块做全局 softmax,显存峰值随卡数下降。
高频追问清单
- FlashAttention 为什么对长上下文关键,它怎么避免物化完整注意力矩阵?
- KV 量化用 per-token 还是 per-channel 缩放,哪种对长序列更稳?
- 序列并行(Ulysses)和 Ring Attention 怎么配合,通信量差多少?
- 4MRAG 长多跳任务里,何时该"全序列展开"、何时该"检索注入",工程判据是什么?
- Prefix Cache(A63) 和长上下文 KV 管理怎么叠加省显存?
- 超长上下文下位置编码(RoPE, A24)的外推与训练长度怎么设才不退化?