MHA, MQA e GQA: como a atenção altera o consumo de memória KV

Tradução automática

Este artigo foi traduzido automaticamente a partir da versão original em inglês.

MHA atribui a cada cabeça de consulta a sua própria cabeça de chave e a sua própria cabeça de valor. MQA partilha uma única cabeça de chave/valor entre todas as cabeças de consulta. GQA partilha uma cabeça de chave/valor dentro de cada grupo de cabeças de consulta. Menos cabeças KV reduzem a cache de atenção convencional e a quantidade de dados que a descodificação lê dessa cache.

Ao estimar o tamanho da cache, verifique num_key_value_heads. Usar apenas o número de cabeças de consulta pode sobrestimar bastante o seu tamanho.

Comparar diretamente a partilha de cabeças

Considere oito cabeças de consulta e a mesma dimensão por cabeça, número de camadas, comprimento de sequência e precisão da cache:

AtençãoCabeças de consultaCabeças KVCache em relação a MHA
Atenção com múltiplas cabeças, MHA88100%
Atenção com múltiplas consultas, MQA8112.5%
Atenção com consultas agrupadas, GQA8225%

Uma cabeça KV inclui uma cabeça de chave e uma cabeça de valor. As projeções de consulta continuam separadas. O artigo de Shazeer sobre MQA e o artigo de Ainslie et al. sobre GQA definem estas alterações.

No Llama 3 70B, 64 cabeças de consulta partilham oito cabeças KV: isto representa oito vezes menos dados KV brutos do que num modelo idêntico nos restantes aspetos com 64 cabeças KV. O Llama 3.1 405B tem 128 cabeças de consulta e oito cabeças KV, o que corresponde a uma redução de dezasseis vezes na mesma comparação. A tabela de arquitetura da Meta fornece estes números.

O que a redução não garante

Uma KV cache oito vezes mais pequena não significa oito vezes menos memória total de GPU nem uma geração oito vezes mais rápida. Os pesos, as ativações, a execução de kernels e a comunicação continuam presentes. As implementações distribuídas também podem replicar cabeças KV entre partições.

O artigo sobre GQA encontrou um compromisso útil entre qualidade e velocidade nos modelos da família T5 que testou. Converter um checkpoint MHA treinado para GQA é uma alteração de arquitetura que exige adaptação e avaliação; não é uma definição de atribuição de memória.

Verifique a configuração de atenção do checkpoint, use o número de cabeças KV na fórmula da cache e meça o desempenho do servidor com os comprimentos de contexto e o número de pedidos simultâneos de que precisa. Inclua a qualidade nas suas tarefas ao comparar checkpoints.

Guia de engenharia: GQA e MQA ilustra os padrões de partilha de cabeças.