Grouped Query Attention (GQA)
Aprende cómo la atención de consultas agrupadas (GQA) reduce la memoria de la caché KV, mejora la eficiencia de la inferencia y equilibra el rendimiento en los modelos Transformer.
El Grouped Query Attention (GQA) es un diseño de atención en el cual múltiples cabezas de consulta comparten un conjunto más pequeño de cabezas de clave y valor. Preserva las diversas perspectivas de las consultas de múltiples cabezas al tiempo que reduce la memoria y el movimiento de datos necesarios para las claves y los valores. El GQA es especialmente valioso durante la inferencia autorregresiva, donde un modelo genera un token a la vez y lee repetidamente los estados de atención almacenados.
Cómo funciona el GQA#
Dentro de un mecanismo de atención, cada token se proyecta en tres representaciones:
- Query: Describe la información que busca el token actual.
- Key: Describe lo que representa cada token disponible.
- Value: Contiene la información recuperada cuando una consulta coincide con una clave.
En la autoatención convencional, estas proyecciones se dividen en cabezas para que el modelo pueda aprender diferentes relaciones en paralelo. El GQA mantiene muchas cabezas de consulta, pero asigna grupos de ellas a cabezas compartidas de clave y valor. Por ejemplo, una capa de atención puede usar ocho cabezas de consulta y dos cabezas de clave-valor. Cada cabeza de clave-valor sirve entonces a cuatro cabezas de consulta.
El cálculo de atención en sí sigue siendo atención de producto escalar escalado. Lo que cambia es la disposición de las cabezas y la cantidad de datos de clave-valor producidos, almacenados y leídos. Marcos como la atención de producto escalar escalado de PyTorch requieren por tanto que el número de cabezas de consulta sea divisible por el número de cabezas de clave-valor cuando el GQA está habilitado. (docs.pytorch.org)
GQA frente a atención de múltiples cabezas y de consulta múltiple#
El GQA se sitúa entre la atención de múltiples cabezas y la atención de consulta múltiple:
- Atención de múltiples cabezas: Cada cabeza de consulta tiene sus propias cabezas de clave y valor. Esto ofrece la máxima independencia de las cabezas, pero crea el mayor requisito de memoria para clave-valor.
- Grouped query attention: Varias cabezas de consulta comparten cada cabeza de clave-valor. Equilibra la capacidad de representación con la eficiencia de memoria.
- Atención de consulta múltiple: Todas las cabezas de consulta comparten una cabeza de clave y una cabeza de valor. Esto minimiza el almacenamiento de clave-valor, pero proporciona menor diversidad de clave-valor.
Si una capa tiene 32 cabezas de consulta, la atención de múltiples cabezas puede usar también 32 cabezas de clave-valor, el GQA podría usar 8 y la atención de consulta múltiple usa 1. El GQA no sustituye a la arquitectura Transformer en general; es una configuración posible dentro de una capa de atención de Transformer. La documentación sobre atención de múltiples cabezas, de consulta múltiple y de consultas agrupadas de NVIDIA describe estas variantes indicando que difieren principalmente en cuántas cabezas de clave-valor sirven a las cabezas de consulta. (nvidia.github.io)
Por qué es importante el GQA#
Durante la generación autorregresiva, un modelo almacena las claves y los valores anteriores en una caché KV. Con la atención de múltiples cabezas estándar, cada cabeza de atención aporta claves y valores en caché separados. Por lo tanto, secuencias largas, lotes grandes y muchas capas de modelos pueden consumir una cantidad sustancial de memoria de la GPU.
Debido a que el GQA utiliza menos cabezas de clave-valor, su caché es más pequeña en proporción a la reducción de dichas cabezas. Esto puede mejorar:
- Capacidad de memoria: Se pueden ajustar en la memoria disponible prompts más largos o un mayor número de solicitudes simultáneas.
- Rendimiento de generación: Se deben leer menos datos de clave-valor por cada token generado.
- Latencia de inferencia: El tráfico de memoria reducido puede acortar el tiempo de procesamiento de token a token.
- Escalabilidad de la ventana de contexto: Gestionar secuencias largas resulta más práctico.
Estas ganancias dependen del hardware, la longitud de la secuencia, el tamaño del lote y la compatibilidad con los núcleos. Los entornos de ejecución en producción siguen necesitando una asignación eficiente de caché; el sistema de caché KV de TensorRT-LLM de NVIDIA, por ejemplo, combina la compatibilidad con GQA con el almacenamiento en caché basado en bloques, la reutilización y la descarga. (nvidia.github.io)
Aplicaciones en el mundo real#
Asistentes de contexto largo: Un asistente de programación puede procesar miles de tokens de código fuente antes de generar una respuesta. El GQA reduce el estado de clave-valor en caché asociado con dicho historial, lo que permite que un servidor gestione archivos más largos o más usuarios concurrentes sin aumentar proporcionalmente la memoria de la GPU.
Asistentes multimodales de imagen y vídeo: Un modelo de visión-lenguaje puede representar fragmentos de imagen o fotogramas de vídeo como secuencias largas de tokens. El GQA reduce la presión sobre la caché de atención después de que estos tokens visuales entran en el decodificador de lenguaje, lo que potencialmente deja más memoria para imágenes, fotogramas o solicitudes adicionales. Esto difiere de un Vision Transformer estándar, donde la atención puede procesar todos los fragmentos de imagen en paralelo sin almacenamiento en caché KV autorregresivo.
Implementación práctica y compensaciones#
PyTorch expone el GQA a través de enable_gqa=True. Este ejemplo documentado de CUDA utiliza ocho cabezas de consulta y dos cabezas de clave-valor compartidas:
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)La salida conserva ocho cabezas de consulta, mientras que las claves y los valores utilizan solo dos cabezas. El selector de backend de atención de PyTorch puede controlar qué núcleo compatible realiza la operación, mientras que la guía de bloques de construcción de Transformer de PyTorch proporciona un contexto de implementación más amplio. (docs.pytorch.org)
El GQA debe elegirse al diseñar o adaptar un modelo; por lo general, no es un interruptor aplicable en tiempo de inferencia para un punto de control existente. Los desarrolladores deben verificar la compatibilidad del framework, la divisibilidad de las cabezas, el comportamiento numérico, el uso de memoria y la calidad de la tarea. Para la implementación en visión artificial, flujos de trabajo complementarios como la integración de TensorRT en Ultralytics y el modo de benchmark de Ultralytics ayudan a medir si las optimizaciones de arquitectura y de tiempo de ejecución ofrecen mejoras significativas en el hardware de destino.






