Grouped Query Attention (GQA)
了解分组查询注意力 (GQA) 如何减少 KV 缓存内存、提高推理效率并在 Transformer 模型中平衡性能。
分组查询注意力机制 (GQA) 是一种注意力设计,其中多个查询头共享较小的一组键头和值头。它保留了多头查询的多样化视角,同时减少了键和值所需的内存与数据移动。GQA 在自回归推理中尤为重要,因为在这种推理中,模型每次生成一个 token 并重复读取存储的注意力状态。
GQA 的工作原理#
在注意力机制中,每个 token 会被投影到三个表示中:
- Query: 描述当前 token 正在寻找的信息。
- Key: 描述每个可用 token 代表的内容。
- Value: 包含当 query 匹配 key 时检索到的信息。
在传统的自注意力中,这些投影被划分为多个头,以便模型并行学习不同的关系。GQA 保留了许多查询头,但将它们分组指派给共享的键头和值头。例如,一个注意力层可能会使用八个查询头和两个键值头。此时,每个键值头服务于四个查询头。
注意力计算本身仍然是缩放点积注意力。改变的是头的排列方式以及生成、存储和读取的键值数据量。因此,当启用 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 内存。
Multimodal image and video assistants: A vision-language model can represent image patches or video frames as long token sequences. GQA reduces attention-cache pressure after these visual tokens enter the language decoder, potentially leaving more memory for additional images, frames, or requests. This differs from a standard Vision Transformer, where attention may process all image patches in parallel without autoregressive KV caching.
实际实现与权衡#
PyTorch 通过 enable_gqa=True 公开 GQA。此文档记录的 CUDA 示例使用了八个查询头和两个共享的键值头:
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)输出保留了八个查询头,而键和值仅使用两个头。PyTorch 注意力后端选择器可以控制哪个受支持的内核执行该操作,而PyTorch Transformer 构建块指南提供了更广泛的实现上下文。(docs.pytorch.org)
在设计或适配模型时必须选择 GQA;它通常不是现有检查点的推理时开关。开发人员应验证框架支持、头的可分性、数值行为、内存使用情况和任务质量。对于计算机视觉部署,诸如Ultralytics TensorRT 集成和Ultralytics 基准测试模式等配套工作流有助于衡量架构和运行时优化是否能在目标硬件上带来显着的性能提升。






