Grouped Query Attention (GQA)
了解分组查询注意力(GQA)如何减少 KV-cache 内存占用、提升推理效率,并平衡 Transformer 模型的性能。
分组查询注意力 (GQA) 是一种注意力设计,其中多个查询头共享较少的键头和值头。它在减少键和值所需的内存及数据传输的同时,保留了多头查询的多样化视角。GQA 在自回归推理期间尤其有价值,因为模型每次生成一个 token,并反复读取存储的注意力状态。
GQA 的工作原理#
在注意力机制中,每个 token 会被投影为三种表示:
- 查询(Query): 描述当前 token 正在寻找的信息。
- 键(Key): 描述每个可用 token 所代表的内容。
- 值(Value): 包含查询与键匹配时检索到的信息。
在传统的自注意力中,这些投影会被划分为多个头,使模型能够并行学习不同的关系。GQA 保留多个查询头,但将它们分组并分配给共享的键头和值头。例如,一个注意力层可以使用 8 个查询头和 2 个键值头。这样,每个键值头就服务于 4 个查询头。
注意力计算本身仍然是缩放点积注意力。变化之处在于头的排列方式,以及生成、存储和读取的键值数据量。因此,启用 GQA 时,诸如 PyTorch 缩放点积注意力之类的框架要求查询头数量必须能被键值头数量整除。(docs.pytorch.org)
GQA 与多头注意力和多查询注意力的比较#
GQA 介于多头注意力和多查询注意力之间:
- 多头注意力: 每个查询头都有自己的键头和值头。这提供了最大的头独立性,但也带来了最大的键值内存需求。
- 分组查询注意力: 多个查询头共享每个键值头。它在表示能力与内存效率之间实现了平衡。
- 多查询注意力: 所有查询头共享一个键头和一个值头。这最大限度地减少了键值存储,但键值多样性较低。
如果一个层有 32 个查询头,那么多头注意力也可能使用 32 个键值头,GQA 可能使用 8 个,而多查询注意力使用 1 个。GQA 并不是对更广义的 Transformer 架构的替代;它是 Transformer 注意力层中的一种可选配置。NVIDIA 的多头、多查询和分组查询注意力文档指出,这些变体的主要区别在于为查询头提供服务的键值头数量不同。(nvidia.github.io)
GQA 的重要性#
在自回归生成期间,模型会将之前的键和值存储在 KV 缓存中。使用标准多头注意力时,每个注意力头都会贡献独立缓存的键和值。因此,长序列、大批量以及大量模型层可能会消耗大量 GPU 内存。
由于 GQA 使用更少的键值头,其缓存大小会按键值头数量的减少比例缩小。这可以改善:
- 内存容量: 更长的提示词或更多并发请求可以容纳在可用内存中。
- 生成吞吐量: 每生成一个 token,需要读取的键值数据更少。
- 推理延迟: 减少内存流量可以缩短 token 到 token 的处理时间。
- 上下文窗口可扩展性: 长序列服务变得更加可行。
这些收益取决于硬件、序列长度、批量大小和内核支持。生产环境运行时仍需要高效的缓存分配;例如,NVIDIA 的 TensorRT-LLM KV 缓存系统将 GQA 支持与基于块的缓存、复用和卸载相结合。(nvidia.github.io)
实际应用#
长上下文助手: 编程助手可能会在生成答案前处理数千个源代码 token。GQA 可减少与这些历史记录相关的缓存键值状态,使服务器能够处理更长的文件或更多并发用户,而无需按比例增加 GPU 内存。
多模态图像和视频助手: 视觉语言模型可以将图像块或视频帧表示为很长的 token 序列。当这些视觉 token 进入语言解码器后,GQA 可以减轻注意力缓存压力,从而可能为更多图像、帧或请求留出内存。这不同于标准的 Vision Transformer,后者的注意力可以并行处理所有图像块,而无需进行自回归 KV 缓存。
实际实现与权衡#
PyTorch 通过 enable_gqa=True 提供 GQA。这个文档中的 CUDA 示例使用 8 个查询头和 2 个共享键值头:
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
assert torch.cuda.is_available(), "A CUDA device is required."
query = torch.randn(1, 8, 16, 32, device="cuda")
key = torch.randn(1, 2, 16, 32, device="cuda")
value = torch.randn(1, 2, 16, 32, device="cuda")
# Four query heads share each key-value head.
with sdpa_kernel(SDPBackend.MATH):
output = F.scaled_dot_product_attention(
query,
key,
value,
enable_gqa=True,
)
print(output.shape)输出保留 8 个查询头,而键和值只使用 2 个头。PyTorch 注意力后端选择器可以控制由哪个受支持的内核执行该操作,而 PyTorch Transformer 构建模块指南则提供更广泛的实现背景。(docs.pytorch.org)
在设计或调整模型时必须选择 GQA;对于现有检查点,它通常不是推理时的开关。开发者应验证框架支持、头数量的整除关系、数值行为、内存使用情况和任务质量。对于计算机视觉部署,Ultralytics TensorRT 集成和 Ultralytics 基准测试模式等配套工作流,有助于衡量架构和运行时优化是否能在目标硬件上带来有意义的改进。









