资讯动态

KV Cache 4-bit量化实践:LLM推理显存占用压缩近4倍

发布时间:2026/9/5 12:03:21 来源:尧图企业网站定制
最近在调一套 LLM 的推理服务模型权重明明还能塞进一张卡但只要 context 一拉长、batch 一起来就 OOM。排查到最后显存大头不是权重而是 KV Cache——那东西随生成的 token 数线性涨解码到 20K token、batch 拉到 8 的时候光 KV Cache 就能吃掉十几个 GB。既然 Flash Attention 解决不了 KV Cache 的存储占用PagedAttention 也只是治碎片不治总量我决定试试把 KV Cache 直接从 FP16 压到 4-bit。实验做下来效果超过预期代码也整理成了一个开源小项目下面把这套方案的设计过程、实现细节和踩坑记录都摊开讲。1. 显存账本里的隐性大户KV Cache 是怎么一步步吃掉一张卡的1.1 解码过程的必然代价Key/Value 缓存为什么省不掉LLM 推理基本分成两个阶段prefill 和 decode。prefill 阶段一次性读入整段 prompt并行算出所有位置的注意力结果decode 阶段每生成一个 token都要拿这个新 token 的 query 去和之前所有 token 的 key、value 做注意力计算。关键点在于decode 阶段每算一个新 token旧 token 的 key 和 value 如果重新算一遍代价等于把整段历史重新过一遍模型这在长上下文场景里是完全不可接受的。所以实现上会为每一层都维护一份 cache存下已经算好的 K 和 V 矩阵每次新 token 算完 K/V 后追加到 cache 尾部后续注意力计算直接读取。这个缓存就是 KV Cache。理论上这是用显存换算力典型的空间换时间。可问题在于显存是硬约束而 KV Cache 是跟着序列长度和 batch size 线性增长的。在多轮对话、长文档阅读、代码仓库级补全这些场景里序列长度动不动就几千甚至上万KV Cache 往往比模型权重本身还要占地方。1.2 KV Cache 的显存公式从 7B 到 13B 的一笔实算KV Cache 的占用有个很简单的估算公式M 2K 和 V× L层数× H_kv × d_head头维度× S序列长度× Bbatch size× dtype_size如果是 MHA 架构且 KV 头数等于 Q 头数公式可以进一步简化成M ≈ 4 × L × H × S × BbytesFP16 下拿常见的 7B 模型举例L32 层H4096FP16 下每个 token 每层要存 K 和 V 两份一共 2×4096 个元素每个元素 2 字节那就是 16KB32 层叠加每生成一个 token 就要多占 512KB 显存。这个数字听起来不大但乘上序列长度和 batch size 就很吓人了场景单条序列长度batchFP16 KV Cache 占用单条长文档20481约 1GB常规并发40964约 8GB长上下文推理81928约 32GB一个 40GB 显存的卡7B 模型权重 FP16 大概 14GB如果 KV Cache 占到 32GB光这两项加起来就已经超过显存容量这还没算激活值、临时缓冲和推理框架里各种中间状态。即使到了 8B 级别的 GQA 模型KV 头从 32 降到 8KV Cache 能做到原来的四分之一每 token 约占 128KB但序列拉到 32K、batch 拉到 4 之后照样要 16GB。KV Cache 是实打实地在跟权重抢显存而且抢得越来越凶。1.3 Flash Attention 和 PagedAttention 为什么都没解决存储问题很多人一开始会想我有 Flash Attention注意力矩阵不落显存KV Cache 应该小了吧这是个常见误解。Flash Attention 解决的问题是计算过程中的中间矩阵——也就是 S QK^T 这样的完整打分矩阵——不显式存到显存里从而省掉了这部分临时开销。但它不改变 K 和 V 本身需要缓存的事实。你依然要把每一层的 K、V 写进 cacheFlash Attention 优化的只是“怎么去读 cache、怎么算注意力”不是“cache 要占多大”。PagedAttention 的思路也类似。它把 KV Cache 切成固定大小的块像操作系统分页一样按需分配尽量减少显存碎片和预分配浪费。在服务多个请求时它可以让 batch 内的显存利用率明显提升但每个 token 需要的字节数没有变总存储量依然是由序列长度和 batch size 决定的。真正能让 KV Cache 存储体积下降的路径除了换 GQA/MQA 这种模型结构层面的改动剩下的其实就在数据表示上把默认的 FP16 降成 8-bit、4-bit甚至更低。这也是我做这个 4-bit 压缩实验的出发点。2. 4-bit 压缩的关键不是位宽是量化粒度2.1 对称均匀量化的基本盘每组一个 scale把 FP16 张量压成 4-bit最直接的做法是均匀量化。思路很简单拿到一组数值记录这组数里的最大绝对值然后用这个最大值把整个范围映射到 [-7, 7] 的整数区间。每个数除以 scale 再四舍五入存成 int4读的时候乘回 scale 就能得到一个近似原值的 FP16 数。在 PyTorch 里按组量化的原型代码不长def quantize_4bit_group(x, group_size64): x: 任意形状最后一维按 group_size 切分 orig_shape x.shape x_flat x.reshape(-1, group_size) amax x_flat.abs().max(dim-1, keepdimTrue).values scale amax / 7.0 # 避免全 0 分组除零 scale torch.where(scale 0, torch.ones_like(scale), scale) q torch.round(x_flat / scale).clamp(-7, 7) return q.to(torch.int8), scale, orig_shape def dequantize_4bit_group(q, scale, orig_shape): return (q * scale).reshape(orig_shape)这个实现里q 才是真正按 4-bit 存储的主体scale 是反推原值时必需的辅助信息也得显式保存。scale 自身也有显存开销所以量化不是无成本地把 16 位变成 4 位实际存储量要加上 scale 的部分。这里有个 trade-offgroup size 越小每组覆盖的数值范围越窄量化误差越小但 scale 的数量越多额外开销越大group size 越大scale 省了但一组里摊到不同量级的数值误差会明显变大。2.2 为什么要按组量化而不是整个张量给一个 scale量化参数里最忌讳的就是整个 KV Cache 张量只给一个 scale。原因很直观不同 head、不同位置的 K/V 数值分布差异非常大。有的 head 的值集中在 [-0.1, 0.1] 之间有的 head 某些通道会出现绝对值几十的 outlier。如果全局只用一个 scale那这个 scale 会被最大的 outlier 顶得很高细颗粒度的小值全部变成 0信息损失惨重。按组量化时数值相似的相邻位置共享同一个 scale能比较有效地隔离 outlier 对局部精度的影响。从存储开销看scale 本身的成本也可以量化计算group size每个数值的实际存储位宽相对 FP16 的压缩比324 16/32 4.5 bit约 3.56 倍644 16/64 4.25 bit约 3.76 倍1284 16/128 4.125 bit约 3.88 倍group_size64 时压缩比接近 3.8 倍group_size128 时省得更多但误差也更大。实际测试下来4-bit KV Cache 的精度瓶颈往往不在位宽而在于 group size 选得是否合适。我最终的默认参数是 K 用 group_size64V 用 group_size128理由后面细说。2.3 K 和 V 要区别对待

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价