FlashAttention (PT)
FlashAttention — é um algoritmo revolucionário para o cálculo do mecanismo de atenção (attention), desenvolvido para acelerar significativamente o treinamento e a inferência de grandes modelos de linguagem (LLMs) preservando a precisão total dos cálculos. O algoritmo foi apresentado pela primeira vez em 2022 por uma equipe de pesquisadores da Universidade de Stanford liderada por Tri Dao[1].
A ideia central do FlashAttention é reorganizar os cálculos levando em consideração a hierarquia de memória da GPU, o que permite minimizar o número de acessos à memória lenta e eliminar o principal gargalo do mecanismo de atenção padrão.
Problemática da Atenção Padrão
O mecanismo de autoatenção padrão nos transformadores é calculado pela fórmula: onde Q, K, V são as matrizes de consulta (query), chaves (key) e valores (value).
O principal problema dessa abordagem é a complexidade quadrática em tempo e memória (O(N²)) em relação ao comprimento da sequência N[1]. Em uma implementação ingênua, é necessário calcular e armazenar na memória da GPU a matriz de atenção completa S de tamanho N×N, o que leva a dois problemas críticos:
- Alto consumo de memória: O armazenamento da matriz N×N se torna inviável ao trabalhar com contextos longos.
- Operações de entrada/saída (IO): O principal gargalo não é o número de operações aritméticas, mas sim os constantes acessos à memória lenta da GPU.
Hierarquia de Memória da GPU
Para entender o problema, é importante distinguir dois tipos de memória em uma GPU (usando a NVIDIA A100 como exemplo):
- SRAM (memória estática): Memória rápida on-chip de pequeno volume (~20 MB) com enorme largura de banda (até 19 TB/s).
- HBM (memória de alta largura de banda): Memória lenta de grande volume (40–80 GB) com uma largura de banda muito menor (cerca de 1.5 TB/s)[2].
Essa assimetria torna o algoritmo de atenção padrão limitado pela largura de banda da memória (memory-bound), pois ele lê e escreve constantemente grandes matrizes da HBM lenta, que é a principal fonte de latência.
Inovações Chave do FlashAttention
O FlashAttention é um algoritmo consciente de I/O (IO-aware) que resolve o problema minimizando os acessos à HBM. Isso é alcançado por meio de três técnicas principais.
Tiling (Divisão em Blocos) e Processamento em Blocos
Em vez de processar a matriz inteira de uma vez, o FlashAttention divide as matrizes de entrada Q, K e V em pequenos blocos (tiles), que cabem na SRAM rápida. O algoritmo carrega sequencialmente esses blocos, realiza todos os cálculos de atenção para eles e atualiza o resultado final, sem salvar a matriz de atenção completa na HBM lenta[1].
Cálculo Online do Softmax
O avanço técnico crucial foi o cálculo "online" do Softmax. O Softmax padrão exige o conhecimento de todos os elementos do vetor de entrada para a normalização. O FlashAttention utiliza um algoritmo modificado que permite calcular o Softmax em partes. Ele mantém dois valores intermediários (o máximo atual e a soma das exponenciais), que são atualizados à medida que novos blocos são processados, permitindo obter o resultado exato sem acessar a matriz inteira de uma vez[2].
Fusão de Operações em um Único Kernel CUDA
Todas as operações de atenção (multiplicação de matrizes QKᵀ, mascaramento, Softmax, multiplicação por V) são combinadas em um único kernel CUDA fundido (fused kernel). Isso reduz drasticamente o número de operações de leitura/escrita na HBM: em vez de múltiplas passagens pela matriz inteira, o algoritmo carrega um bloco na SRAM uma vez, realiza todos os cálculos e escreve apenas o resultado final.
Eficiência Teórica e Prática
Complexidade e Otimidade
O FlashAttention reduz o consumo de memória de O(N²) para O(N), o que garante escalabilidade linear. Foi provado que a complexidade de I/O do algoritmo é teoricamente ótima para o cálculo da atenção em uma hierarquia de memória de dois níveis, o que significa que é impossível executar a atenção exata mais rapidamente sem alterar o hardware[3].
Resultados Empíricos
A primeira versão do FlashAttention demonstrou melhorias significativas:
- Aceleração:
- BERT-large (comprimento de sequência 512): aceleração de 15% no treinamento.
- GPT-2 (comprimento 1K): aceleração de 3 vezes.
- Tarefas do Long-Range Arena (1K-4K): aceleração de 2,4 vezes[1].
- Economia de memória: Economia de memória de até 20 vezes em comparação com implementações base exatas.
- Melhora na qualidade dos modelos: Graças à capacidade de trabalhar com contextos mais longos, o FlashAttention não apenas mantém, mas também melhora a qualidade dos modelos. Por exemplo, a perplexidade do GPT-2 melhorou em 0,7 pontos, e a precisão em tarefas de classificação de documentos longos aumentou em 6,4 pontos[1].
Evolução e Desenvolvimentos Futuros
O sucesso do FlashAttention deu início a uma série de algoritmos orientados a hardware.
FlashAttention-2 (2023)
A segunda versão visava a uma utilização mais completa dos recursos da GPU. No FlashAttention original, a eficiência na NVIDIA A100 era de apenas 25–40% do máximo. O FlashAttention-2 introduziu melhorias na paralelização dos cálculos, o que permitiu[4]:
- Alcançar uma aceleração de duas vezes em comparação com a primeira versão.
- Aumentar a utilização da GPU para 50–73% do máximo teórico.
- Expandir o suporte para cabeças de atenção de tamanho 256, bem como para arquiteturas Multi-Query Attention (MQA).
FlashAttention-3 (2024)
A terceira versão foi otimizada especificamente para a arquitetura de GPU NVIDIA Hopper (H100)[5]. Ela utiliza novos recursos de hardware, como a assincronia dos Tensor Cores e o suporte a FP8, o que permitiu:
- Alcançar uma aceleração adicional de 1,5 a 2 vezes em comparação com o FlashAttention-2.
- Atingir um desempenho de até 740 TFLOPS em FP16 e perto de 1.2 PFLOPS em FP8.
Soluções Especializadas
As ideias do FlashAttention foram desenvolvidas em outros projetos:
- FlashInfer (2025): Um motor de atenção customizável, otimizado especificamente para tarefas de inferência de LLMs. Ele foca na operação eficiente com o cache KV no modo de geração de streaming[6].
- FlashMLA (2024): Uma implementação de atenção com compressão do cache de contexto (latent attention), que permite economizar memória em sequências muito longas com perda mínima de informação[7].
Impacto na Indústria e no Ecossistema
O FlashAttention se tornou um avanço fundamental e rapidamente se transformou no padrão da indústria para o treinamento e a inferência eficientes de LLMs. Ele foi integrado em bibliotecas-chave, como PyTorch e Hugging Face, e é utilizado na maioria dos grandes modelos de linguagem (LLaMA, MPT, Falcon, Claude, etc.).
Foram o FlashAttention e suas versões subsequentes que desempenharam um papel decisivo no aumento das janelas de contexto dos modelos de linguagem: de 2–4 mil tokens (GPT-3) para 128 mil tokens (GPT-4) e até milhões de tokens em modelos experimentais[8]. O algoritmo eliminou um dos principais obstáculos para a escalabilidade dos transformadores, abrindo novas possibilidades para aplicações de IA, desde a análise de documentos longos até a compreensão multimodal.
Ligações externas
Literatura
- Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135.
- Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv:2307.08691.
- Shah, J. et al. (2024). FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision. arXiv:2407.08608.
- Kwon, W. et al. (2023). Efficient Memory Management for Large Language Model Serving with PagedAttention. arXiv:2309.06180.
- Hong, K. et al. (2023). FlashDecoding++: Faster Large Language Model Inference on GPUs. arXiv:2311.01282.
- Ye, Z. et al. (2025). FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving. arXiv:2501.01005.
- Dege, P. et al. (2025). FlashMLA-ETAP: Efficient Transpose Attention Pipeline for Accelerating MLA Inference on NVIDIA H20 GPUs. arXiv:2506.01969.
- Wang, G. et al. (2025). FlashMask: Efficient and Rich Mask Extension of FlashAttention. OpenReview wUtXB43Chi.
- Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (OpenReview version). OpenReview H4DqfPSibmx.
- Gholami, A. et al. (2024). FlashAttention on a Napkin: A Diagrammatic Approach to Deep Learning IO-Awareness. OpenReview pF2ukh7HxA.
Notas
- ↑ 1.0 1.1 1.2 1.3 1.4 Dao, Tri, et al. “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”. arXiv:2205.14135 [cs.LG], 28 de maio de 2022. [1]
- ↑ 2.0 2.1 Dao, Tri, et al. “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”. OpenReview. [2]
- ↑ “We're Training AI Twice as Fast This Year as Last”. IEEE Spectrum. [3]
- ↑ Dao, Tri. “FlashAttention-2”. tridao.me. [4]
- ↑ “FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision”. PyTorch Blog. [5]
- ↑ “[2501.01005] FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving”. arXiv. [6]
- ↑ “GitHub - deepseek-ai/FlashMLA: FlashMLA: Efficient MLA decoding kernels”. GitHub. [7]
- ↑ “The Evolution of Flash Attention: Revolutionizing Transformer Efficiency”. Medium. [8]