用 AI 算 KV 缓存:查询头有十六个,缓存也可能只按四组保存

前天 3阅读

本地模型的权重已经装得下,为什么对话变长后显存还会涨?原因之一是生成过程保存了过去位置的键和值。可以让AI按张量形状做一份容量预算,但先要确认该数查询头,还是实际存储的KV头。

用 AI 算 KV 缓存:查询头有十六个,缓存也可能只按四组保存

AI模型生成的概念插图:成对缓存随序列延长;不是硬件照片、显存监控截图或精确内存布局。

只计算一类明确的缓存

假设一个教学解码器有12层,每层保存完整历史。批量B=2,两条序列都已缓存T=1024个位置;查询头为16,KV头为4,每个键头和值头维度都为64,每个元素占2字节。没有量化、前缀共享、窗口裁剪或额外副本。

Hugging Face文档说明缓存按层保存K与V,基本形状包含批量、头数、序列长度和头维度。GQA论文区分了查询头与共享的键值头。本文假设实现以4个KV头的紧凑形状持久存储,不能直接把16个查询头填进容量公式。

先算一层,再乘层数

一层的K共有2×4×1024×64=524288个元素,每个2字节,正好1048576字节,即1MiB。V同样是1MiB,所以一层2MiB,12层合计24MiB,也就是25165824字节。这里1MiB=1048576字节,不是十进制的一百万字节。

把计算合起来,普通等形状KV载荷为2×B×层数×KV头数×T×头维度×每元素字节数。最前面的2代表K和V两份;B本身又是2,两个乘二的来源要分清。若误用16个查询头,就会算成96MiB,比本设定多四倍。

每多一个位置,要给两条序列都留空间

在批量中的每条序列各增加一个缓存位置,总增量为2×2×12×4×64×2=24576字节,即24KiB。若两条都从1024延长到1536,增加512个位置,载荷增加12MiB,总共36MiB。

本文用整数运算独立核对了逐层算法与总公式。T是实际存入缓存的token位置数,不是中文字数,也不是只计算新回复的长度。提示词处理过的位置也可能已经在缓存里;最后刚选出的token是否已进入缓存,要看它是否又完成了下一次前向计算。

让AI同时交假设和算式

“估算一个12层解码器的完整历史KV载荷:B=2,每条T=1024;查询头16,实际持久存储KV头4;K/V头维度均64,元素2字节,无共享、无量化、无裁剪。先分别算单层K与V,再算总字节与MiB。计算每条各多1位置,以及T变1536时的增量。指出误用查询头会得到什么,并列出未包含的显存项目。”

这份预算不含模型权重、注意力临时张量、其它激活、分配器空闲块与框架开销。静态缓存可能按预定最大长度先分配,滑动窗口可能停止增长;不同架构也可能采用不同表示。实际容量规划还要观察所用实现,不能把24MiB宣称为整机运行需求或实测显存。

资料核对日期:2026年10月2日。本文使用原创合成设定,独立核算不代表真实模型训练或效果测试。

参考资料

Hugging Face:How caching works

Ainslie等:GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints

文章版权声明:除非注明,否则均为云鹊BLOG原创文章,转载或复制请以超链接形式并注明出处。