GQA原理:如何减少KV缓存
估算LLM的KV缓存时,不能直接把query头数当成KV头数。GQA让一组query头共享同一组key和value头,因此query头可以保持较多,而需要缓存的KV数据更少。关键在于共享的是K与V,不是把所有query输出变成同一个结果。
8个query头,不一定缓存8组KV
用8个query头说明三种结构:MHA为每个query头配置对应的KV头;MQA让全部query头共享一组KV;GQA处在两者之间,例如4个query头共享一组KV,共2组。这里“一个KV头”表示一对同编号的key头与value头,不是把K和V合成一个张量。
| 结构 | 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量化则解释数值表示与误差。结构、存储管理与位宽应分别核对,不要用其中一个名词代替全部内存优化。
资料来源
支持
如果这篇文章对你有帮助,欢迎支持本站。
二维码可点击放大。更多支持方式见支持页面。


