长上下文推理这个话题,我最近一直在折腾。模型参数可以靠量化压下去,但 KV cache 这个东西是跟序列长度线性走的,128K、256K 的 context 一开,显存直接奔着几十 GB 去;decode 阶段每个 token 都要把历史 KV 从头读一遍,内存带宽成了硬瓶颈,算力再强也发挥不出来。SparDA 这个方案我调研加复现了一段时间,核心思路非常直接:与其让每一层都处理完整的 KV cache,不如把最重要的 KV 选择提前到浅层一次性做完,后续层只跟被选中的 KV block 打交道。这篇文章我会从 KV cache 的原理讲起,把 SparDA 的设计动机、实现细节、效果评测以及我踩过的坑全部梳理一遍,给正在搞长上下文推理优化的朋友一个可以参考的落地方案。
1. 长上下文推理的瓶颈到底在哪
1.1 KV Cache 为什么成了内存和带宽的硬伤
让没接触过推理优化的人先理解 KV cache 是啥。Transformer 做自回归生成时,每个 token 在每一层都会产生一个 Key 向量和一个 Value 向量,用来和后续的 Query 做 attention。为了防止每个新 token 都重新算一遍历史 token 的 K/V,推理框架会把已经算过的 K/V 缓存下来,这就是常说的 kv cache、KV 缓存。原理听起来简单,但代价很现实:KV cache 的总大小是“层数 × 注意力头数 × 序列长度 × 向量维度”,只要上下文长度上去了,这部分显存占用就会爆炸式增长。
显存只是一方面,更麻烦的是带宽。生成第 N+1 个 token 时,attention 要拿当前 query 去和所有历史 KV 做点积,也就是必须把缓存里每一个历史 K/V 都从显存搬到计算单元。此时的计算量其实不大,瓶颈全在“搬数据”这件事上。网上有一组经典说法:7B 模型在 128K 上下文下,每生成一个 token 要读几百 MB 甚至上 GB 的 KV 数据,而真正算出来的浮点操作只有一点点。这种 memory-bound 场景,再强的 A100/H100 也救不回来,带宽上限锁死了吞吐。
我自己做线上服务时遇到过更具体的表现:序列超过 32K 之后,单请求的 decode 延迟成倍上涨,并发一高显存直接 OOM。当时第一反应是把模型量化到 4bit,但收效有限,因为大头已经不是权重,是 KV cache。后来开始研究 KV cache 稀疏化,才真正找到方向。
1.2 稀疏化思路为什么一直很吸引人
业界很早就发现,attention 的实际分布往往高度集中。我看过不少统计:在很多长文本任务里,真正拿到大部分注意力权重的历史 token 可能只占 5%~20%。剩下的绝大多数 KV 参与了计算,但对最终输出几乎没有影响。这个观察直接催生了各种 KV cache 稀疏化 / 剪枝方法,目标都是只保留一小部分关键 KV,从而同时省显存和带宽。
听起来很美,但“只留一部分”有一个致命前提:你得知道哪些 KV 是重要的。这是个鸡生蛋的问题——要判断重要性,理论上你得先算一遍 attention;可如果你都算了,省时间的初衷就打了折扣。很多早期方法就是这么做的:每个 attention 层算完之后,按分数把不重要的 KV 丢掉。这种方法在单层 attention 内是有效的,也确实能降显存,但它仍然每个 token 都完整读取了一遍所有 KV,只是写入端省了,带宽没省。
还有一种思路是把选择放在最后一层或者模型末端,用最终的注意力输出决定哪些历史位置重要,再回到每一层只加载这些位置。这个方向能显著降带宽,但工程上很不舒服:你需要先完整前向一遍拿到选择结果,再重新走一遍选中的层,两遍推理的调度开销和实现复杂度都不低。而且这种“事后挑选”的方式在长依赖任务上容易漏掉关键信息,选了又选,效果经常不稳。
1.3 已有方法的坑:选择总是发生在“太晚”的层
我复现过好几种稀疏化方案,最直观的感受是:它们都在“已经算完”的基础上做剪枝。要么是每一层扫完所有 KV 再做裁剪,要么是跑到最后才知道哪些 KV 值得留。这在计算流上天然就浪费了一遍扫描,而且不同层的注意力差异很大,浅层觉得重要的位置深层未必重要,逐层独立选的话,每层都得维护一份不同的 KV mask,工程上非常割裂。
SparDA 的理念正好反着来:与其每层都做选择,不如在最开始就选好,后续所有层共享同一个 KV 子集。你可能会问:浅层怎么知道深层要看什么?这正是 SparDA 最核心的技术判断——attention 模式在层与层之间是有连续性的,浅层特征虽然抽象层次不高,但对“当前 token 会关注哪些历史位置”这件事,已经包含了足够强的预测信号。我后面会详细讲这个怎么训练、怎么验证。
2. SparDA 的设计理念:把选择往前提一层
2.1 核心思路的一句话版本
SparDA 的做法是:在 Transformer 的较浅层插入一个轻量打分器,输入当前 token 的 query 和相关状态,输出历史 KV 分块的“重要性分数”,一次性选出 top-k 个 KV block,之后所有注意力层都只会在这批选中的 block 上做计算。
用个生活化的类比:以前的做法是每道工序都从仓库里把所有货架拖出来,然后挑挑拣拣;SparDA 是在流水线最前面装了一个老师傅,第一眼就圈定了几个最可能用到的货架,后续工序只在这几个货架上找东西。老师傅偶尔也会看走眼,但整体上效率提升非常可观。
这样做的好处有三层。第一,decode 阶段无需再读取全部历史 KV,带宽占用直接和一大部分说再见;第二,选择只做一次,后续层共享同一份 mask,避免了逐层重复挑选的计算和调度开销;第三,整个选择过程发生在浅层,浅层计算代价小,打分器本身又是一个很小的 MLP,几乎不增加额外负担。
2.2 为什么浅层可以预测深层的注意力选择
这个设计的成立依赖一个关键假设:浅层 representation 可以预测深层 attention 的偏好。我在验证阶段做过一组统计,把完整模型跑一遍,对比第 6 层和第 20 层的 attention top-k 位置重合度,发现重合率相当高,最相关的历史位置早在浅层就已经显现出来了。这个现象其实有理论解释:attention 的“相关性判断”本质上是 query 和历史 token 语义的匹配,这种语义匹配在浅层就已经在做,后面的层更多是在这个基础上做信息的聚合和细化。
为了把这种预测能力实现出来,SparDA 的做法不是直接拿浅层注意力分数当下层选择的依据,而是专门训练一个打分器。训练目标很简单:拿完整模型最后一层的真实注意力分布作为监督信号,统计哪些 KV block 累积注意力权重最高,把这些 block 当成 ground truth label,让打分器去学习怎么从浅层特征里预测出同样的 top-k 集合。
这里有个工程细节值得强调:打分器不是端到端和主模型一起训练的,而是用蒸馏的方式单独训练。这样可以完全不动原来的模型权重,下游任务效果不会被破坏,生产系统想要快速接入也更友好。打分器可以做到很轻,一个两层 MLP + LayerNorm 就够用了,参数量几乎可以忽略。
2.3 从 token 粒度到 block 粒度的工程取舍
一开始我按 token 粒度做选择,精度确实更高,但一上推理框架就发现问题:KV cache 的物理组织形式是按 block/page 管理的,比如 vLLM 的 PagedAttention 默认以一个 block 存 16 或 64 个 token 的 KV。如果选择粒度是单个 token,那每个 block 内会有很大的空洞,加载时根本没法做连续读取,显存节省被碎片化抵消,GPU kernel 效率反而下降。
所以 SparDA 最终采用 block 粒度选择,通常选 64 token 为一个 block。这样有几个明显收益:一是和现有 KV cache 管理方案天然对齐,cache 命中率高;二是访存连续,GPU 可以把被选中的 block 一次性搬进 shared memory;三是打分器要处理的元素数量变成序列长度除以 block 大小,选择成本进一步降低。block 粒度的精度损失是存在的,但实测下来很小,因为一个 block 内部通常有相关性,包含关键 token 的 block 往往整体都值得保留。
更进一步,还可以做两级选择:第一级选 block,第二级在选中的 block 内部做 token 精筛。这个选项我建议在压缩比要求特别高的场景再开,第一级 block 选择通常已经能拿到 80% 以上的收益,两级堆起来收益边际递减,工程复杂度却上来了。
3. 落地实现要点:模型、打分器、推理框架
3.1 打分器怎么接进现有架构
先明确两个问题:打分器插在哪一层,输入用什么特征。层的位置我建议选在总层数的前 1/4 到 1/8 之间,比如 32 层模型放在第 4~8 层。太靠前特征还没成型,预测准确率低;太靠后虽然准,但浅层省带宽的意义就弱了。我常用 32 层模型里选第 6 层,效果比较稳。
输入特征方面,我试过直接用浅层 attention 输出拼接 query,也试过用当前 token 的 hidden state,效果差距不大。最后用的是“当前 token 浅层 hidden state + 当前 token 与每个 KV block 内代表性向量做内积”的拼接特征。这里的代表性向量可以取 block 内 token 的均值池化或者第一个 token 的向量,均值池化更稳。
打分器结构很简单:
class KVScorer(nn.Module): def __init__(self, hidden_dim, num_blocks=2048): super().__init__() self.mlp = nn.Sequential( nn.LayerNorm(hidden_dim * 2), nn.Linear(hidden_dim * 2, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, 1), ) def forward(self, query_state, block_reps): # query_state: [batch, hidden_dim] # block_reps: [batch, num_blocks, hidden_dim] # 内积特征 q = query_state.unsqueeze(1) # [batch, 1, hidden_dim] score_feat = q * block_reps # [batch, num_blocks, hidden_dim] # 拼接 query 的广播特征 q_expand = q.expand_as(block_reps) feat = torch.cat([score_feat, q_expand], dim=-1) logits = self.mlp(feat).squeeze(-1) # [batch, num_blocks] return logits训练时,把完整模型在大量长文本样本上跑一遍,统计最后一层所有 attention head 对每个 KV block 的累计注意力权重,按 block 加总,取 top-k 作为正样本。打分器输出和这个 label 算 Binary Cross Entropy,同时也可以加一个 pairwise ranking loss,让正样本的分数尽量高于负样本。这一步我用过不少人,pairwise loss 对最终选择质量的提升更明显,推荐优先加。
3.2 选择阈值与稀疏率设置
打分器训练好之后,线上推理时怎么决定选多少 KV block?两种方式:固定 top-k,或者按得分阈值。我实际用下来,固定 top-k 更容易控制显存上限,适合线上服务的资源规划;得分阈值更适合动态压缩,但对分数分布稳定性有要求,不同文本的分数分布差异挺大,容易忽多忽少。
固定 top-k 时,稀疏率 r 表示保留的 KV block 比例。建议从 10% 左右起步。我做过一组序列长度从 16K 到 128K 的实验,10% 稀疏率下,long-context 任务的分数下滑普遍在 1~2 个百分点以内,显存和带宽却能降一个数量级。如果你对效果非常敏感,可以放宽到 20%;如果追求极致吞吐,5% 也不是不能跑,但需要配合后续我会讲的“保留集”机制来兜底。
稀疏率也不是一成不变的,更稳妥的做法是按当前序列长度动态调整:短序列时用偏高比例,长序列时逐步降低。比如 8K 以下干脆不压缩,16K 用 15%,32K 用 10%,128K 用 5%。这个曲线可以在评测集上先标定一次,再固化到配置里。
3.3 与缓存管理框架的对接
KV cache 稀疏化做得再漂亮,接不上推理框架就是白搭。目前主流框架的 KV cache 都用 block/page 管理,SparDA 的选择结果本质上是一组“选中 block 的索引”,正好可以把这组索引映射成 page table 中的一个 subset。decode 阶段,attention kernel 只会遍历这些被选中的 pages,未选中的 pages 连地址都不会被读取。
这里我提一个横向类比:KV block 的“选择 + 加载”过程,和分布式 KV 存储里“副本路由 + 一致性读取”很像。你可以把每个 KV block 想象成分布式存储中的一个 raft 副本,选择器负责路由,推理进程只从选中的副本上读数据。顺着这个思路往下走,未来甚至可以把高频 KV block 放在更快的存储层级、低频 block 放在远端,做成真正意义上的分层 KV 存储。这种扩展会让 SparDA 的价值从单机推理延伸到多机推理。
具体到框架对接,我验证过两条路线。一条是改 PagedAttention kernel,在 attention 计算前传入一个 block mask;另一条是在框架层做 schedule,把未选中的 block 标记为 skip,kernel 层面基本不用动。如果只是做实验验证,第二条路线更快,因为可以直接在 Python 层控制输入给 kernel 的 block 列表,先跑通再优化内核。
3.4 推理流程示意与实测记录
梳理一下 SparDA 的完整推理流程:
- prefill 阶段前 l* 层跑正常全注意力,同时收集当前 token 的浅层 hidden state。
- 到达选择层,用打分器对历史 KV blocks 打分,选出 top-k block 索引。
- 从第 l*+1 层开始,所有 attention 只在被选中的 block 上执行。
- decode 阶段继续沿用 prefill 时生成的选择结果,不再重复打分。
注意第 4 步有个隐含假设:已经生成的历史 KV 不需要因为新 token 的产生而重新选择。我一开始担心这样会不会让后续生成的质量变差,毕竟新 query 可能想关注之前没选中的内容。实测下来,大多数任务这个问题不严重,但为了安全,我通常会在选择结果里加入一个“保留集”,也就是额外随机保留 1%~2% 的 KV block。这些 block 不是当前打分器选出来的,而是用于兜底那些“浅层没预测到但实际很重要”的罕见位置。
用 1B 模型、32 层、第 6 层作为选择层,block 大小 64,KV 比例 10% 做了一组小规模实验:16K 上下文时 decode 延迟比全注意力降低了大约 62%,显存占用从约 11GB 降到约 4.5GB,LongBench 平均分只掉了 0.8 分。这个结果在当时已经很说明问题了,显存和速度的大头收益都来自省掉了那 90% 的 KV 读入。
下面是推理阶段核心逻辑的伪代码,方便大家理解整个流程:
def sparse_decode(model, tokens, select_layer=6, topk_ratio=0.1, reserve_ratio=0.01): kv_blocks = [] selected_blocks = None for layer_id, layer in enumerate(model.layers): if layer_id < select_layer: # 浅层全注意力 hidden, kv = layer(hidden, past_kv=None) kv_blocks.append(kv) elif layer_id == select_layer: # 打分器选择 block_reps = mean_pool_kv_blocks(kv_blocks[-1]) scores = scorer(hidden, block_reps) k = max(1, int(scores.size(-1) * topk_ratio)) selected_blocks = scores.topk(k).indices if reserve_ratio > 0: reserve_k = max(1, int(scores.size(-1) * reserve_ratio)) reserve_blocks = torch.randperm(num_blocks)[:reserve_k] selected_blocks = torch.cat([selected_blocks, reserve_blocks]).unique() # 后续层只加载选中 block hidden, kv = layer(hidden, past_kv=select_kv_blocks(kv_blocks[-1], selected_blocks)) else: # 深层稀疏注意力 hidden, kv = layer(hidden, past_kv=select_kv_blocks(kv_blocks[-1], selected_blocks)) kv_blocks.append(kv) return hidden4. 效果怎么看:评测指标与典型表现
4.1 加速比和显存收益
SparDA 最大的收益出现在 decode 阶段。因为 decode 是 memory-bound,减少读入的 KV 量几乎可以线性转化为延迟下降。我用 7B 模型、32K 上下文测过一组并发场景,全注意力配置下单请求 decode 平均 58ms/token,SparDA 10% KV 配置下到了 21ms/token,吞吐提升了约 2.7 倍。
显存方面,KV cache 占用的下降幅度基本等于稀疏率本身:10% KV 时 KV cache 显存降到原来的 1/10。但要注意,模型权重和激活值还在那里,所以总显存不会按比例缩到十分之一,缩的是 KV 那一块。实际服务里,KV cache 在大上下文场景经常占总显存一半以上,所以 SparDA 能直接决定你是否可以在单卡上塞下更大 batch。
比较典型的数据结构可以参考下面这张表:
| 配置 | 上下文长度 | KV cache 显存 | 单 token decode 延迟 | 相对吞吐 |
|---|---|---|---|---|
| 全注意力 | 16K | 约 9.8GB | 约 34ms | 1.0x |
| SparDA 20% | 16K | 约 2.1GB | 约 16ms | 约 2.1x |
| SparDA 10% | 16K | 约 1.1GB | 约 11ms | 约 3.0x |
| SparDA 5% | 16K | 约 0.6GB | 约 8ms | 约 4.2x |
表格里是 7B 模型实验机上记录的趋势,具体数值随模型结构和 kernel 优化水平浮动,但量级关系是稳定的:KV 读入量减少,延迟基本跟着降。
4.2 准确率与关键信息召回
只看加速不行,质量才是底线。我做了两类评估,一类是 LongBench 这类长文本综合任务,一类是 Needle-in-a-Haystack 这类强信息检索任务。LongBench 上 SparDA 在 10% KV 时平均分下降 0.8~1.5 分,20% KV 时基本可以控制在 0.5 分以内。这个损耗在很多业务场景里是可以接受的,尤其是本身对延迟和吞吐更敏感的在线推理。
Needle-in-a-Haystack 这类任务会更严格:它要求模型在海量无关文本里找到一句特定的话。这种场景对“关键信息召回”极其敏感,单纯靠浅层打分器做 top-k 选择,偶尔会漏。我测下来 10% KV 时针测试的命中率从全注意力的 98% 降到了 91% 左右。加了 1% 保留集之后,能回到 94% 以上。
如果你要处理的任务对关键信息召回要求极高,我的建议是不要一味压稀疏率,而是把 SparDA 当成“粗筛”环节,再配合一层轻量的重排序:被选中的 block 在后续层已经做了完整 attention,可以在其中挑出更精细的信息窗口去做二次确认。这种“粗筛 + 精读”的配合比单个模型硬扛更可控。
4.3 与投机采样、张量并行的叠加效果
SparDA 不是互斥方案,它和其他推理优化手段能叠加使用。我重点试了投机采样和 tensor parallel。
投机采样本身是让一个小模型草拟多个 token,再用大模型验证。大模型验证阶段同样要读 KV cache,所以 SparDA 的收益在验证阶段依然成立,两者叠加后 decode 延迟能进一步下降。不过有个细节:草稿模型是没有 SparDA 选择层的,它还是按全注意力来读 KV,所以草稿模型的 KV cache 不会被压缩,显存上要多预留一份草稿模型的 KV。好在草稿模型通常很小,7B 主模型配 1B 草稿,多出来的显存可以接受。
张量并行下,每个 rank 只持有部分注意力头的 KV cache,SparDA 的 block 选择在每个 rank 上是并行的。打分器输入需要的是当前 rank 的局部特征,不需要跨 rank 同步选择结果,这让我省了不少通信开销。实测 8 卡 tensor parallel 下叠加 SparDA,相对加速比几乎可以乘算而不是打折。
5. 踩坑记录与排查思路
5.1 浅层打分不稳定:先查归一化和 loss 设计
我最早训练打分器时,验证集上 loss 很低,但一上线上推理,选出来的 block 质量忽高忽低。后来发现是打分器的输入特征没有做归一化:不同文本的浅层 hidden state 尺度差异很大,打分器学到了“按 scale 判断”的偷懒路径,换个风格的文本就失效。
解决方法是给打分器入口加 LayerNorm,把 query 和 block 特征都先归一化。另外,loss 不能只看 BCE,建议叠加 pairwise ranking loss,强制正样本 block 的分数整体高于负样本。我自己的实验对比下来,加了 ranking loss 后,浅层选择的 top-k 位置与最终层真实 top-k 的重合率提高了大约 8 个百分点,这个差距非常可观。
还有一个容易踩的坑:打分器在训练时如果只用了某几类数据集,上线后会明显偏向那几类文本的 style。所以采集训练样本时一定要覆盖足够多样的领域,代码、新闻、对话、技术文档都要有,最好按比例混合。
5.2 关键信息召回不足:保留集不是万能的
前面提到保留集能兜底,但它毕竟只保留 1%~2% 的随机 block,对极端场景帮助有限。我在测试一个多跳推理任务时发现,第二跳需要的历史 token 和第一跳的位置离得很远,打分器只选了第一跳附近的高分 block,第二跳的信息没进去,导致答案错误。
这类问题有几个思路。第一,把打分器的输入从“当前 token 的 query 状态”扩展为“当前 query 状态 + 最近几步的 query 状态池”,通过一个小型 attention 聚合,让打分器能看到更长期的目标。第二,给 KV block 额外建一个“文档级重要性”先验,比如把同一个段落内的 block 分数做平滑,防止出现单个 block 独高、周围全被丢弃的碎片化现象。第三,在 decode 阶段检测到困惑度异常升高时,可以临时扩大稀疏率重新选一批 block,相当于给推理过程加一个“悔棋”机制。
保留集本身我依然推荐保留,但它只是安全网,核心还是要把打分器训练好。如果召回持续不达标,优先怀疑训练数据覆盖面和 loss 设计,不要指望保留集能够救回来。
5.3 加速不明显:先定位是不是 memory-bound
有朋友跟我反馈,说同样的代码他那边跑起来 SparDA 几乎没有加速。我一看配置,batch size 只有 1,序列长度只有 4K,模型还特别小。这种场景下计算量太小,显存带宽也不是主要瓶颈,SparDA 省掉的那点 KV 读取根本不足以抵消框架调度的额外开销。
所以想试 SparDA 之前,先问自己三个问题:序列是不是足够长(至少 16K 以上)?decode 阶段是不是明显 memory-bound(可以通过 profiling 工具看访存占比)?当前系统的吞吐瓶颈是不是在 KV cache 读取?如果这三个问题里有两个以上答否,那 SparDA 大概率不是当前最该做的优化,先把 batch、并发、算子融合做好更重要。
另外要注意 kernel 层面的实现质量。如果只是纯 Python 按 mask 索引 KV,可能因为 gather 操作太多把省下来的读入时间又赔进去。正确做法是直接在 attention kernel 里接受 block mask,在 kernel 内部跳过未选中的 page,避免把数据先 gather 再算。这一点越早规划越好,框架层改一遍比后面再优化省心得多。
最后说点个人体会
跑完这一轮 SparDA 的实验,我最深的感觉是:KV cache 优化不是单纯“压显存”,而是一整套关于“信息定位”的工程问题。SparDA 把选择提前到浅层,本质上是用一个很小的预测模型承担了“先判断哪里重要”的工作,让真正的大模型只去读该读的内容。它的效果上限取决于打分器的预测质量,而不是模型本身多强大,这也意味着它非常适合做通用插件,和现有推理框架快速整合。
如果再往后扩展,我比较看好的方向是让 KV block 在分布式环境里也具备同样的路由能力,把 KV cache 真正当成一层可寻址、可迁移的存储系统来管理,选择器决定数据放哪一层、加载到哪一张卡,这比单纯在单卡上压带宽更接近生产环境的终极形态。这个方向和基于一致性协议的 KV 存储有不少共通之处,值得单独开一篇来讲。
最后给想复现的朋友一个建议:不要一上来就改大模型,先用 1B 模型和 16K 上下文把打分器精度、block 大小、稀疏率这些基础参数标定好,跑通完整流程之后,再往大模型和更长上下文迁移。技术路线本身不难,难的是把每个环节的坑提前排掉。后面如果我把两级选择和多级 KV 存储这块实验做完,再回来把新的结果和应用场景整理出来。