GQA原理:如何减少KV缓存

黎 浩然/ 11 10 月, 2026/ 大语言模型/LARGELANGUAGEMODEL/LLM, 机器学习/MACHINELEARNING, 研究生/POSTGRADUATE, 计算机/COMPUTER/ 0 comments

估算LLM的KV缓存时,不能直接把query头数当成KV头数。GQA让一组query头共享同一组key和value头,因此query头可以保持较多,而需要缓存的KV数据更少。关键在于共享的是K与V,不是把所有query输出变成同一个结果。

文章目录
  1. 8个query头,不一定缓存8组KV
  2. 缓存比例由KV头数决定
  3. 共享KV,输出仍然可以不同
  4. 论文结果不能变成统一保证
  5. 部署时怎样核对
  6. 资料来源

8个query头,不一定缓存8组KV

用8个query头说明三种结构:MHA为每个query头配置对应的KV头;MQA让全部query头共享一组KV;GQA处在两者之间,例如4个query头共享一组KV,共2组。这里“一个KV头”表示一对同编号的key头与value头,不是把K和V合成一个张量。

Q0 · Q1 · Q2 · Q3
共享 K0 / V0
Q4 · Q5 · Q6 · Q7
共享 K1 / V1
原创分组示意:8个query头共享2组KV头;编号仅为教学说明,不对应某个模型权重。
结构 query头 KV头 共享关系
MHA 8 8 每个query对应一组
GQA 8 2 每4个query共享一组
MQA 8 1 全部query共享一组

GQA论文于2023年5月先提交arXiv,同年12月收入EMNLP会议论文集,是经过同行评审的历史研究,不是今天的新消息。论文将GQA描述为MHA与MQA之间的分组方案。

缓存比例由KV头数决定

假设单个请求,L层、T个已缓存token、每层Hkv个KV头、每头d个分量,每个数占b字节,K和V形状一致且以紧凑格式存储。原始缓存字节数为:

KV bytes = 2 × L × T × Hkv × d × b

取L=32、T=4096、d=128、b=2。8个KV头需要512 MiB,2个需要128 MiB,1个需要64 MiB。1 MiB=2²⁰字节。在这些条件保持一致时,2头GQA的原始KV缓存是8头MHA的四分之一。

这个比例不包含权重、激活、内存对齐、分配器预留及多设备复制,也不意味着总显存或推理时间缩小到四分之一。批量请求长度不同时,应按各请求实际缓存长度相加,而不是拿一个短请求代替全部请求。

共享KV,输出仍然可以不同

对同组两个query,q₁和q₂可以不同。即使用相同K与V,softmax(q₁Kᵀ/√d)V和softmax(q₂Kᵀ/√d)V仍可能不同。共享KV减少表示的自由度,但没有强制所有query生成同一个注意力分布。

下面的原创教学代码已运行:复算三种缓存大小,检查8到2的分组映射,并用两个不同query读取完全相同的key和value。最后两个输出不同。标量value只用于把例子写短,完整注意力的value通常是向量。

import math

def kv_bytes(layers, tokens, kv_heads, head_dim, bytes_per_value):
    return 2 * layers * tokens * kv_heads * head_dim * bytes_per_value

for name, heads in [("MHA", 8), ("GQA", 2), ("MQA", 1)]:
    size = kv_bytes(32, 4096, heads, 128, 2)
    assert size == {"MHA": 512, "GQA": 128, "MQA": 64}[name] * 1024**2
    print(name, size // (1024 ** 2), "MiB")

query_heads, kv_heads = 8, 2
assert query_heads % kv_heads == 0
mapping = [h // (query_heads // kv_heads) for h in range(query_heads)]
assert mapping == [0,0,0,0,1,1,1,1]

def attend(query):
    keys = [(1.0,0.0), (0.0,1.0)]
    values = [10.0,20.0]
    scores = [sum(a*b for a,b in zip(query,k))/math.sqrt(2)
              for k in keys]
    weights = [math.exp(s-max(scores)) for s in scores]
    return sum(w*v for w,v in zip(weights,values))/sum(weights)

# Two queries share exactly the same keys and values.
a, b = attend((1.0,0.0)), attend((0.0,1.0))
assert not math.isclose(a,b)
print(round(a,6), round(b,6))

这段代码没有加载模型,没有训练GQA,也没有使用GPU。缓存算术与局部注意力差异通过断言核对,不构成任何模型的质量或性能测试。

论文结果不能变成统一保证

原论文的主要实验基于T5.1.1,讨论了检查点转换与追加训练。作者在摘要、翻译及问答任务中报告了GQA的质量与推理权衡。计时设置采用TPUv4;这些是作者的实验条件,不是任意模型、显卡与负载都能复现的保证。

不能只修改一个配置中的KV头数,就认为既有MHA权重会成为效果相同的GQA模型。权重形状、转换方法和后续训练都需要与架构一致。比较现成模型时,也要避免把不同模型规模和训练数据造成的质量差异,全部归因于KV头数。

部署时怎样核对

先从模型的实际配置确认query头数、KV头数、头维度和缓存数据类型,再用上述公式计算原始预算。运行时测量峰值内存,并同时记录输入长度、输出长度、并发和缓存策略,才容易定位估算与实测的差额。

GQA与PagedAttention作用在不同层面:前者改变共享结构,后者文章讨论KV缓存的组织。今天上午的LLM量化则解释数值表示与误差。结构、存储管理与位宽应分别核对,不要用其中一个名词代替全部内存优化。

English version

资料来源

GQA, EMNLP 2023; 会议论文原文; arXiv submission history.

支持

如果这篇文章对你有帮助,欢迎支持本站。

微信支持二维码,点击查看大图
微信
支付宝支持二维码,点击查看大图
支付宝
Buy Me a Coffee,点击支持本站
Buy Me a Coffee

二维码可点击放大。更多支持方式见支持页面。

Share this Post

Leave a Comment

您的邮箱地址不会被公开。 必填项已用 * 标注

*
*