Grouped Query Attention (GQA)
Aprende como o Grouped Query Attention (GQA) reduz a memória da KV-cache, melhora a eficiência da inferência e equilibra o desempenho em modelos Transformer.
A Atenção de Consultas Agrupadas (GQA) é um design de atenção no qual várias cabeças de consulta partilham um conjunto menor de cabeças de chave e valor. Preserva as perspetivas diversificadas das consultas multi-head, reduzindo simultaneamente a memória e a movimentação de dados necessárias para chaves e valores. A GQA é especialmente valiosa durante a inferência autorregressiva, na qual um modelo gera um token de cada vez e lê repetidamente os estados de atenção armazenados.
Como funciona a GQA#
Dentro de um mecanismo de atenção, cada token é projetado em três representações:
- Consulta: Descreve a informação que o token atual procura.
- Chave: Descreve o que cada token disponível representa.
- Valor: Contém a informação recuperada quando uma consulta corresponde a uma chave.
Na autoatenção convencional, estas projeções são divididas em cabeças para que o modelo possa aprender diferentes relações em paralelo. A GQA mantém muitas cabeças de consulta, mas atribui grupos dessas cabeças a cabeças de chave e valor partilhadas. Por exemplo, uma camada de atenção pode usar oito cabeças de consulta e duas cabeças de chave-valor. Cada cabeça de chave-valor serve então quatro cabeças de consulta.
O próprio cálculo da atenção continua a ser uma atenção de produto escalar escalado. O que muda é a disposição das cabeças e a quantidade de dados de chave-valor produzidos, armazenados e lidos. Por isso, frameworks como a atenção de produto escalar escalado do PyTorch exigem que o número de cabeças de consulta seja divisível pelo número de cabeças de chave-valor quando a GQA está ativada. (docs.pytorch.org)
GQA vs. atenção multi-head e multi-query#
A GQA situa-se entre a atenção multi-head e a atenção multi-query:
- Atenção multi-head: Cada cabeça de consulta tem as suas próprias cabeças de chave e valor. Isto oferece a máxima independência entre cabeças, mas cria o maior requisito de memória para chave-valor.
- Atenção de consultas agrupadas: Várias cabeças de consulta partilham cada cabeça de chave-valor. Equilibra a capacidade de representação com a eficiência de memória.
- Atenção multi-query: Todas as cabeças de consulta partilham uma cabeça de chave e uma cabeça de valor. Isto minimiza o armazenamento de chave-valor, mas proporciona menos diversidade de chave-valor.
Se uma camada tiver 32 cabeças de consulta, a atenção multi-head também pode usar 32 cabeças de chave-valor, a GQA pode usar 8 e a atenção multi-query usa 1. A GQA não substitui a arquitetura Transformer mais abrangente; é uma possível configuração dentro de uma camada de atenção Transformer. A documentação da NVIDIA sobre atenção multi-head, multi-query e de consultas agrupadas descreve estas variantes como diferindo principalmente no número de cabeças de chave-valor que servem as cabeças de consulta. (nvidia.github.io)
Por que a GQA é importante#
Durante a geração autorregressiva, um modelo armazena chaves e valores anteriores numa cache KV. Com a atenção multi-head padrão, cada cabeça de atenção contribui com chaves e valores em cache separados. Sequências longas, lotes grandes e muitas camadas do modelo podem, por isso, consumir uma quantidade substancial de memória GPU.
Como a GQA usa menos cabeças de chave-valor, a sua cache é menor na proporção da redução dessas cabeças. Isto pode melhorar:
- Capacidade de memória: Prompts mais longos ou mais pedidos simultâneos podem caber na memória disponível.
- Débito de geração: É necessário ler menos dados de chave-valor para cada token gerado.
- Latência de inferência: A redução do tráfego de memória pode diminuir o tempo de processamento entre tokens.
- Escalabilidade da janela de contexto: A disponibilização de sequências longas torna-se mais prática.
Estes ganhos dependem do hardware, do comprimento da sequência, do tamanho do lote e do suporte do kernel. Os runtimes de produção continuam a precisar de uma alocação eficiente da cache; o sistema de cache KV do TensorRT-LLM da NVIDIA, por exemplo, combina suporte para GQA com cache baseada em blocos, reutilização e descarregamento. (nvidia.github.io)
Aplicações no mundo real#
Assistentes com contexto longo: Um assistente de programação pode processar milhares de tokens de código-fonte antes de gerar uma resposta. A GQA reduz o estado de chave-valor em cache associado a esse histórico, permitindo que um servidor processe ficheiros mais longos ou mais utilizadores simultâneos sem aumentar proporcionalmente a memória GPU.
Assistentes multimodais de imagem e vídeo: Um modelo de visão-linguagem pode representar regiões de imagens ou fotogramas de vídeo como sequências longas de tokens. A GQA reduz a pressão sobre a cache de atenção depois de estes tokens visuais entrarem no descodificador de linguagem, deixando potencialmente mais memória para imagens, fotogramas ou pedidos adicionais. Isto difere de um Vision Transformer padrão, no qual a atenção pode processar todas as regiões da imagem em paralelo sem cache KV autorregressiva.
Implementação prática e compromissos#
O PyTorch expõe a GQA através de enable_gqa=True. Este exemplo CUDA documentado usa oito cabeças de consulta e duas cabeças de chave-valor partilhadas:
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)A saída mantém oito cabeças de consulta, enquanto as chaves e os valores usam apenas duas cabeças. O seletor de backend de atenção do PyTorch pode controlar qual kernel suportado executa a operação, enquanto o guia de blocos de construção Transformer do PyTorch fornece um contexto de implementação mais abrangente. (docs.pytorch.org)
A GQA deve ser escolhida ao conceber ou adaptar um modelo; geralmente não é uma opção de inferência para um checkpoint existente. Os programadores devem verificar o suporte do framework, a divisibilidade das cabeças, o comportamento numérico, o uso de memória e a qualidade da tarefa. Para a implementação em visão computacional, fluxos de trabalho complementares, como a integração Ultralytics TensorRT e o modo de benchmark Ultralytics, ajudam a medir se as otimizações arquiteturais e de runtime proporcionam melhorias significativas no hardware-alvo.









