Grouped Query Attention (GQA)
Scopri come la Grouped Query Attention (GQA) riduce la memoria della KV-cache, migliora l'efficienza dell'inferenza e bilancia le prestazioni nei modelli Transformer.
L'attenzione alle query raggruppate (GQA) è un'architettura di attenzione in cui più teste di query condividono un insieme più piccolo di teste di key e value. Mantiene le diverse prospettive delle query multi-head, riducendo al contempo la memoria e lo spostamento dei dati necessari per key e value. GQA è particolarmente utile durante l'inferenza autoregressiva, in cui un modello genera un token alla volta e legge ripetutamente gli stati di attenzione memorizzati.
Come funziona GQA#
All'interno di un meccanismo di attenzione, ogni token viene proiettato in tre rappresentazioni:
- Query: descrive le informazioni che il token corrente sta cercando.
- Key: descrive ciò che rappresenta ogni token disponibile.
- Value: contiene le informazioni recuperate quando una query corrisponde a una key.
Nel self-attention convenzionale, queste proiezioni vengono suddivise in teste affinché il modello possa apprendere relazioni diverse in parallelo. GQA mantiene molte teste di query, ma assegna gruppi di esse a teste di key e value condivise. Ad esempio, un livello di attenzione può utilizzare otto teste di query e due teste di key-value. Ogni testa di key-value serve quindi quattro teste di query.
Il calcolo dell'attenzione rimane un'attenzione con prodotto scalare scalato. Cambiano invece la disposizione delle teste e la quantità di dati key-value prodotti, memorizzati e letti. I framework come l'attenzione con prodotto scalare scalato di PyTorch richiedono quindi che il numero di teste di query sia divisibile per il numero di teste di key-value quando GQA è abilitato. (docs.pytorch.org)
GQA rispetto all'attenzione multi-head e multi-query#
GQA si colloca tra l'attenzione multi-head e l'attenzione multi-query:
- Attenzione multi-head: ogni testa di query ha le proprie teste di key e value. Ciò offre la massima indipendenza tra le teste, ma crea il maggiore fabbisogno di memoria per key-value.
- Attenzione alle query raggruppate: diverse teste di query condividono ciascuna testa di key-value. Bilancia la capacità rappresentazionale con l'efficienza della memoria.
- Attenzione multi-query: tutte le teste di query condividono una testa di key e una testa di value. Riduce al minimo lo spazio di archiviazione key-value, ma offre una minore diversità di key-value.
Se un livello ha 32 teste di query, l'attenzione multi-head può utilizzare anche 32 teste di key-value, GQA potrebbe utilizzarne 8 e l'attenzione multi-query ne utilizza 1. GQA non sostituisce la più ampia architettura Transformer; è una possibile configurazione all'interno di un livello di attenzione Transformer. La documentazione di NVIDIA sull'attenzione multi-head, multi-query e alle query raggruppate descrive queste varianti come differenti principalmente per il numero di teste di key-value che servono le teste di query. (nvidia.github.io)
Perché GQA è importante#
Durante la generazione autoregressiva, un modello memorizza key e value precedenti in una cache KV. Con l'attenzione multi-head standard, ogni testa di attenzione contribuisce con key e value memorizzati nella cache separati. Sequenze lunghe, batch di grandi dimensioni e numerosi livelli del modello possono quindi consumare una quantità considerevole di memoria GPU.
Poiché GQA utilizza meno teste di key-value, la sua cache è più piccola in proporzione alla riduzione di tali teste. Ciò può migliorare:
- Capacità di memoria: prompt più lunghi o un numero maggiore di richieste simultanee possono rientrare nella memoria disponibile.
- Throughput di generazione: è necessario leggere meno dati key-value per ogni token generato.
- Latenza di inferenza: la riduzione del traffico di memoria può abbreviare il tempo di elaborazione tra un token e l'altro.
- Scalabilità della finestra di contesto: la gestione di sequenze lunghe diventa più pratica.
Questi vantaggi dipendono dall'hardware, dalla lunghezza della sequenza, dalle dimensioni del batch e dal supporto del kernel. I runtime di produzione devono comunque gestire un'allocazione efficiente della cache; il sistema di cache KV di TensorRT-LLM di NVIDIA, ad esempio, combina il supporto a GQA con il caching basato su blocchi, il riutilizzo e l'offloading. (nvidia.github.io)
Applicazioni nel mondo reale#
Assistenti con contesto lungo: un assistente alla programmazione può elaborare migliaia di token di codice sorgente prima di generare una risposta. GQA riduce lo stato key-value memorizzato nella cache associato a tale cronologia, consentendo a un server di gestire file più lunghi o un numero maggiore di utenti simultanei senza aumentare proporzionalmente la memoria GPU.
Assistenti multimodali per immagini e video: un modello visione-linguaggio può rappresentare patch di immagini o fotogrammi video come lunghe sequenze di token. GQA riduce la pressione sulla cache dell'attenzione dopo che questi token visivi entrano nel decoder linguistico, lasciando potenzialmente più memoria per immagini, fotogrammi o richieste aggiuntivi. Ciò differisce da un Vision Transformer standard, in cui l'attenzione può elaborare tutte le patch dell'immagine in parallelo senza il caching KV autoregressivo.
Implementazione pratica e compromessi#
PyTorch espone GQA tramite enable_gqa=True. Questo esempio CUDA documentato utilizza otto teste di query e due teste di key-value condivise:
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)L'output mantiene otto teste di query, mentre key e value ne utilizzano solo due. Il selettore del backend di attenzione di PyTorch può controllare quale kernel supportato esegue l'operazione, mentre la guida agli elementi costitutivi di Transformer di PyTorch fornisce un contesto di implementazione più ampio. (docs.pytorch.org)
GQA deve essere scelto durante la progettazione o l'adattamento di un modello; in genere non è un'opzione attivabile in fase di inferenza per un checkpoint esistente. Gli sviluppatori devono verificare il supporto del framework, la divisibilità delle teste, il comportamento numerico, l'utilizzo della memoria e la qualità del task. Per il deployment di computer vision, workflow complementari come l'integrazione Ultralytics TensorRT e la modalità benchmark di Ultralytics aiutano a misurare se le ottimizzazioni architetturali e di runtime producono miglioramenti significativi sull'hardware di destinazione.









