给Transformer加第四维,NVIDIA把长上下文解码速度拉到1.7倍
NVIDIA和MIT联合推出的新Transformer变体SparDA,只在每一层多做一个小投影,就解决了长上下文推理中KV缓存 offload 的两大核心痛点,实测解码速度提升1.7倍,长推理准确率涨6.5个点。
NVIDIA和MIT的研究人员给Transformer做了一个很小的改动,把长上下文LLM推理的解码速度提了1.7倍,长推理准确率还涨了6.5个点。
先补一点基础背景,方便没技术背景的读者看懂:很多人用ChatGPT的时候都发现,第一个字出得慢,之后出字就变快。这个差异来自KV缓存技术,已经是现在LLM推理的标配,能把推理速度提5倍。
Akshay Pachaar之前写过一篇KV缓存的入门讲解,把这个技术说的非常清楚:简单来说,大模型生成文字是一个字一个字往外蹦的,每生成一个新字,都需要用到之前所有字的Key和Value向量。如果每次都重新算这些向量,会产生巨量的重复计算,KV缓存就是把算好的向量存起来,只算新字的向量,省下大量重复工作。
但KV缓存不是没有问题,上下文越长,缓存需要占的内存就越大。10万token以上的长上下文,缓存根本装不下GPU显存,只能把一部分缓存放到CPU内存里,要用的时候再拷回GPU。这就是现在长上下文推理最大的瓶颈。
之前的稀疏注意力方案,已经做了一步优化:不会给所有缓存块都做计算,只选top-k最重要的块保留计算。但这个方案还是绕不开两个问题:
第一,选块要用到当前层的查询向量Q,Q算出来才能去CPU内存捞需要的块,GPU只能空等着拷数据,每一步解码都要卡一次。
第二,选块本身也有计算成本。分组查询注意力GQA里,每个分组多个查询头共享KV头,原来的方案要给每个查询头算一次分,再做softmax,成本跟着上下文长度一起涨。
这次NVIDIA推出的SparDA,解决问题的思路很简单,就是给每一层多加一个投影,叫Forecast。原来只有Q、K、V三个投影,现在变成四个。
这个Forecast的作用,就是由当前层L预测下一层L+1需要哪些KV块。这么一改,两个痛点同时解决了:
1. 下一层需要的块,在当前层计算的时候就已经知道了,可以用独立的CUDA流提前从CPU内存预取,拷数据的过程和当前层的计算并行,GPU不用再空等。
2. Forecast和查询向量Q完全解耦,不需要给每个查询头都做一次打分,一个GQA分组只需要一个Forecast头就够了,直接省掉了每个头的打分循环和softmax步骤,选块的计算成本降了一大块。
这个改动的成本非常小。在8B参数的模型上,Forecast只增加了3350万参数,占总参数的0.41%,而且只需要训练新增的Forecast投影,用KL损失对齐原来选块器的分布就行,不需要全参数重训。
实测结果也很漂亮,在MiniCPM4.1-8B和NOSA-8B两个模型上,准确率和原来的稀疏方案持平甚至更高,其中NOSA-8B的长推理准确率涨了6.5个点。prefill阶段最快快1.25倍,解码阶段最快快1.7倍。
还有一个额外的好处:因为预取把offload的成本藏住了,大部分KV缓存都可以放在CPU内存里,省出来的GPU显存能放下更大的批量,解码吞吐量比不做offload的稀疏基线最高提了5.3倍。
当然这个优化不是全场景生效,它的收益主要来自解码阶段用CPU offload的场景。prefill阶段所有KV本来都已经在GPU上,收益主要来自降低选块成本。
有意思的是,DeepSeek之前在DSA里也用过类似的思路:让一个小索引器选重要token,而不是让查询自己选。SparDA把这个思路用到了块级别,还加上了DSA没有覆盖的预取优化,属于站在之前的思路上补完了痛点。
论文已经公开在arXiv,代码也放到了NVIDIA的GitHub仓库,感兴趣可以看原文:
https://arxiv.org/abs/2606.04511
很多从业者评论,大模型推理优化走到现在,很多这种小改动反而能拿到大收益,大家对这个优化怎么看?
发布时间: 2026-08-13 17:16