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

PerguntaPorquê 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.