O que FlashAttention muda na inferência de transformers
Tradução automática
Este artigo foi traduzido automaticamente a partir da versão original em inglês.
FlashAttention calcula a atenção exata sem guardar as matrizes completas de pontuações e probabilidades de atenção na memória principal da GPU. Processa blocos mais pequenos e combina os seus resultados com um softmax incremental. Isto reduz o tráfego de memória e o armazenamento intermédio, mas mantém o custo aritmético quadrático da atenção densa em função do comprimento da sequência.
Um menor consumo de memória não faz com que o cálculo da atenção densa cresça de forma linear.
Separar o armazenamento do cálculo
A atenção materializada padrão calcula pontuações para cada par visível de consulta e chave, normaliza essas pontuações e multiplica-as pelos valores. Uma matriz de atenção quadrada sem máscara contém N² entradas para N tokens.
O artigo original de FlashAttention reorganiza esse cálculo em blocos que cabem no chip. Mantém estatísticas de normalização e acumuladores de saída atualizados durante o cálculo, em vez de escrever todas as probabilidades em HBM. Com dimensões de cabeça fixas, o armazenamento necessário para os resultados intermédios de atenção cresce linearmente com o comprimento da sequência. Continua a calcular as interações densas entre consultas e chaves exigidas pela máscara de atenção.
O algoritmo calcula atenção exata; ordens de execução diferentes das operações em vírgula flutuante podem causar pequenas diferenças numéricas. Não é um método aproximado de atenção esparsa que descarta determinados pares de tokens.
A implementação e a carga de trabalho continuam a importar
| Pergunta | Porquê verificar |
|---|---|
| O backend suporta esta GPU e esta precisão? | As implementações dos kernels destinam-se a hardware específico |
| A dimensão das cabeças, a máscara e a disposição da cache são suportadas? | Formas não suportadas podem levar à seleção de outro backend |
| A atenção representa uma grande parte do tempo do pedido? | Um grande ganho no kernel pode ter um efeito global pequeno |
| A carga de trabalho é prefill ou descodificação de um único token? | O trabalho que podem executar em paralelo é diferente |
FlashAttention-2 também paraleliza blocos da sequência de consultas, além dos lotes e das cabeças, e reduz outros custos adicionais de execução. FlashAttention-3 destina-se a funcionalidades de execução específicas de Hopper. Os nomes das versões descrevem algoritmos e implementações; não indicam o backend que o servidor realmente selecionou.
Verifique o backend selecionado no runtime de serving com a versão fixada. Meça a latência do prompt, o pico de memória e o débito de ponta a ponta com comprimentos e tamanhos de lote representativos. Compare máscaras e precisão idênticas. Um benchmark que apenas apresenta os TFLOPS do kernel não demonstra uma melhoria na latência por token que o utilizador observa.
Guia de engenharia: FlashAttention aborda as versões sucessivas e as suas medições nos respetivos âmbitos.