文章名:Keyformer: KV Cache reduction through key tokens selection for Efficient Generative Inference
论文的关键:
在vllm中,Decode的时候由于矩阵乘法和KV cache的存取,该阶段受限于memory bound
研究发现,在推理过程中,90%的attention权重关注于特定的token
这篇文章的工作就是通过一个新颖的评分函数来找到这些特定的token
在KV cache中就只保存这些特定的token
对于多embedding和多任务上,keyfromer能降低延迟,提高吞吐量,不丢失准确率
背景
生成的新token会带来额外的KV cache,如MPT-7B 模型中,序列长度增加 16 倍(从 512 到 8K)会导致推理延迟增加 50 倍以上。推理总时间的约 40%(绿色突出显示)被 KV 缓存数据移动所消耗,还会延长其他操作所需的时间(蓝色显示)。当序列长度超过 8K 时,KV 缓存大小超过了模型大小。

阻尼系数α
对于上面提到的问题,Keyformer 通过智能地丢弃不必要的标记动态减少 KV cache大小。这种动态是将上一次的key和最近的几个token进行结合,作为下一次的key token。在数学上来看,就是在不断增加的kv cache中,取其中的子集作为key token作为下一次迭代的kv cache。
在这个过程中,如果仅靠key attention是不够的,如下图,key attention(绿色)不如full attention

这是因为本身的softmax函数已经会删掉数值比较小的预测,再加上只保留key attention,所以就更低了。下图就是对比全attention,删掉了之后模型不会关注序列最相关的token,丢失了上下文信息。

在数学上来说,这种数据的减少相当于减少样本域的样本,然后每个样本被抽中的概率不一样了
右下角的项是被去掉的
所以作者就是想找一个折中的办法,去掉一些样本,但是影响不那么大,就用一个阻尼系数α。使得\bar{f}_\theta=\alpha f_\theta,来抑制因为丢弃数据带来产生过多的注意力分数
那么接下来就要设计试验来证明阻尼α能不能抵消去掉数据带来的影响
首先,我要去掉多少数据呢?下图表明,去掉50%数据不会对key attention产生过多的影响

那么证明其实数据里50%左右是key attention,那么就丢弃50%的kv cache。试验结果发现,单一的α无法替代减少50%的kv cache所带来的影响

Gumbel 分布针对上面提到的问题,其实就是加了一个项去调整丢掉一半kv cache之后的数据分布使其接近原分布,只是效果不好,既然出来的数据不能接近原分布,那么就调整数据分布,使用一种类似正则化的方式进行拟合。
首先,对qkv的数据来说,加入一个偏置的分布
其中分布为Gumbel分布,Gumbel 分布在统计理论中具有重要意义。它抓住了甘贝尔极限定理的精髓,即常见的概率分布(如正态分布、指数分布、均匀分布等)都会向甘贝尔分布收敛。这凸显了它在关键标记识别建模中的适当性。
那么如何验证这种分布具有意义呢?面对数据分布的问题,就要用熵来解决,或者用KL散度,这里没用不知道为什么
除了加了一个分布,还需要来抑制分布带来的尖锐或者过缓,所以加入一个除项
这有什么影响呢?当τ无限大,e的0次方是1,分子为1,分母为k,所有预测的结果都是1/k,这意味着所有 token 的概率几乎相等,形成均匀分布(uniform distribution)。随机性最大,每个 token 被选中的机会差不多。在 KV 缓存中移除 token 时有用,因为均匀分布允许更多 token 有机会被保留或采样,而不会过度偏向特定 token。
当 τ 很小时,指数部分无限大,其 exp 值会急剧增大,其他 token 的 exp 值相对变小,这形成尖锐分布(sharp distribution),概率几乎全部集中在最高 logit 的 token 上,趋向于“one-hot”分布(即 argmax 操作)。
至此,文章的操作就很清晰了,在decode的过程中

在prompt阶段,也就是Decode step1,其实是一次full attention,在数据中加入Gumbel分布噪音用于识别关键token,然后选取 w 个最近的token,然后在n - w个token中选取k - w个token。k - w个key attention和w个最近的token组成了新的k个缩小的KV cache。由于是在prompt阶段,τ取无限大来软化最大概率分布。
在decode的时候,也就是Decode step 2,使用缩小的KV cache进行decode,最后得分函数加上上一层的得分函数。
结果
结果就是Keyformer(绿线)在使用60%~70%的KV Cache的情况下,已经能跑出99%水平的full attention结果。

在延迟和吞吐量方面,由于直接减少了KV Cache,那么推理速度肯定是快的

结论
结论的翻译原文:
大型语言模型(LLM)的进步推动了更长的上下文和大量文本的生成,模型是在数百万个词组的序列中训练出来的。然而,这种趋势会占用系统内存带宽,导致执行成本增加。在较长的上下文中,主要造成内存带宽消耗和推理延迟的 KV 缓存大小超过了模型参数的大小。
为了解决这个问题,我们提出了 Keyformer,它可以在不牺牲准确性的情况下,通过丢弃跨头、层和梁的令牌,根据新颖的分数函数识别重要令牌(关键令牌),从而有效地将 KV 缓存的大小减少 50%。Keyformer 可在推理时应用于 LLM,无需微调,同时还能改善延迟和令牌生成吞吐量。
源代码解读
keyformer-kv-cache-reduction-through-key-tokens-selection-for-efficient-generative-inference.pdf