全部术语

GLOSSARY ENTRY

KV Cache 量化

  • KV Cache Quantization
  • FP8 KV Cache
  • Int8 KV Cache

用低位宽 payload 与配套 scale/metadata 存储历史 K/V,降低容量和 decode 读取带宽;真实收益取决于量化粒度、误差、packing 与消费 kernel,而不是只把公式中的字节数替换。

它改变缓存表示,不改变逻辑 shape

标准 KV Cache 逻辑 shape 仍是 [B,Hkv,T,D][B,H_{kv},T,D]。量化把浮点向量 xx 映射为整数/FP8 payload,并保存 scale:

s=maxixi127,qi=clip(round(xi/s),127,127)s=\frac{\max_i|x_i|}{127},\qquad q_i=\operatorname{clip}(\operatorname{round}(x_i/s),-127,127)

消费时 x^i=sqi\hat x_i=sq_i,或在 attention kernel 中边加载边反量化。它不减少 token 数、KV heads 或 attention 历史长度;降低的是每元素 payload,代价是 scale metadata、量化/反量化计算和误差。

Granularity 决定 metadata 与误差

粒度scale shape 示例特点
per-tensor[1]metadata 最少,outlier 容易压低整体分辨率
per-head[B,Hkv,1,1] 或静态 head scale适配 heads 动态范围,仍较粗
per-token/per-head[B,Hkv,T,1]每个 token head 单独 scale,误差小些,metadata 随 T 增长
per-group[B,Hkv,T,D/G]更细但 scale/packing 与 kernel 更复杂

对 asymmetric quantization 还需 zero point;FP8 可能使用 per-tensor/per-head scale 与不同格式。配置中的 “int8/fp8” 不足以推断真实字节数。

可运行 fake-quant 实验

import math,platform,torch
torch.manual_seed(137)
DEVICE=torch.device("cuda" if torch.cuda.is_available() else "cpu")
DTYPE=torch.float32

def quantize_per_token_head(x):
    scale=x.abs().amax(-1,keepdim=True).clamp_min(1e-8)/127
    q=torch.round(x/scale).clamp(-127,127).to(torch.int8)
    return q,scale

def attend(q,k,v):
    scores=q@k.transpose(-2,-1)/math.sqrt(q.shape[-1])
    return torch.softmax(scores,-1)@v

def main():
    b,h,t,d=2,4,16,32
    q=torch.randn(b,h,1,d,device=DEVICE,dtype=DTYPE)
    k=torch.randn(b,h,t,d,device=DEVICE,dtype=DTYPE); v=torch.randn_like(k)
    qk,sk=quantize_per_token_head(k); qv,sv=quantize_per_token_head(v)
    k_hat,v_hat=qk.float()*sk,qv.float()*sv
    reference=attend(q,k,v); quantized=attend(q,k_hat,v_hat)
    torch.testing.assert_close(k_hat,k,rtol=0,atol=.02)
    torch.testing.assert_close(v_hat,v,rtol=0,atol=.02)
    fp16_payload=2*k.numel()*2
    int8_payload=qk.numel()+qv.numel()
    scale_bytes=(sk.numel()+sv.numel())*sk.element_size()
    total=int8_payload+scale_bytes
    print(f"python={platform.python_version()} torch={torch.__version__}")
    print(f"device={DEVICE} dtype={DTYPE} Q={tuple(q.shape)} K=V={tuple(k.shape)}")
    print(f"K_max_abs_error={(k_hat-k).abs().max():.6f} V_max_abs_error={(v_hat-v).abs().max():.6f}")
    print(f"attention_max_abs_error={(quantized-reference).abs().max():.6f}")
    print(f"fp16_payload_bytes={fp16_payload}")
    print(f"int8_payload_bytes={int8_payload} scale_metadata_bytes={scale_bytes} total={total}")
    print(f"effective_compression_ratio={fp16_payload/total:.3f}")

if __name__=="__main__":
    with torch.inference_mode(): main()

完整文件位于 examples/kv_quantization_reference.py。实际运行:K/V 最大误差约 0.0140/0.0145,attention 输出最大误差 0.007099。FP16 K+V payload 16384 bytes;INT8 payload 8192 bytes,加 FP32 per-token/head scales 1024 bytes,总压缩比只有 1.778×,不是理想的 2×。

这是 fake quant:在 PyTorch 中立刻反量化回 FP32 再做 dense attention,没有证明低位宽 cache kernel 能省带宽或加速。生产收益要求 payload 保持 packed 低位宽直到消费,并将反量化融合进 attention load path。

容量公式与生产开销

原公式 M=2LBTHkvDsM=2LBTH_{kv}Ds 只覆盖 payload。量化后:

Mtotal=Mpayload(bq/8)+Mscale+Mzero+Mpadding/metadataM_{total}=M_{payload}(b_{q}/8)+M_{scale}+M_{zero}+M_{padding/metadata}

bqb_q 是量化位数。4-bit 还要说明 two values/byte 的 packing、group alignment 和尾部 padding。Paged blocks 可能为每 block/head 保存 scale,实际 metadata 粒度与数学 quant group 不完全相同。

误差会在长上下文、不同层/head、outlier 与位置分布下变化;一个随机小张量的 max error 不能替代 perplexity、任务准确率和长上下文 eval。K 与 V 的敏感性也可能不同,runtime 可使用不同策略。

何时不成立

  • backend 不支持低位宽 KV,先展开为 FP16 后读取;
  • scale metadata/反量化占比在小 D/短序列过高;
  • 量化降低 batch 可用容量却因 kernel 更慢导致 TPOT 恶化;
  • calibration/static scale 不覆盖服务数据 outlier;
  • prefix cache blocks 使用不同量化配置或 scale identity,不能安全共享;
  • 与 MLA 混用时仍按标准 2HkvD2H_{kv}D 估算,缓存字段算错。

上下游关系

  • 推理量化基础给出通用 scale、granularity、weight-only/W8A8 与 fake quant 边界;本页把同一数值表示问题放到跨 step 持久的 K/V 上。
  • KV Cache给出生命周期与基础容量;本页替换 payload 表示并补 metadata。
  • PagedAttention按 blocks 分配量化 payload,block layout 必须与消费 kernel 匹配。
  • Prefix Cache共享 blocks 时要把量化配置/scale 语义纳入 identity。
  • MLA先改变缓存元素字段,KV quantization 再改变这些字段的表示,二者效果不能机械相乘。

参考资料