它改变缓存表示,不改变逻辑 shape
标准 KV Cache 逻辑 shape 仍是 。量化把浮点向量 映射为整数/FP8 payload,并保存 scale:
消费时 ,或在 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。
容量公式与生产开销
原公式 只覆盖 payload。量化后:
是量化位数。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 混用时仍按标准 估算,缓存字段算错。
上下游关系
- 推理量化基础给出通用 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 再改变这些字段的表示,二者效果不能机械相乘。