FlashAttention (IT)
FlashAttention — è un algoritmo rivoluzionario per il calcolo del meccanismo di attention, sviluppato per accelerare significativamente l'addestramento e l'inferenza dei grandi modelli linguistici (LLM) preservando la piena precisione dei calcoli. L'algoritmo è stato presentato per la prima volta nel 2022 da un team di ricercatori dell'Università di Stanford guidato da Tri Dao[1].
L'idea chiave di FlashAttention consiste nella riorganizzazione dei calcoli tenendo conto della gerarchia di memoria della GPU, il che consente di minimizzare il numero di accessi alla memoria lenta ed eliminare il principale collo di bottiglia del meccanismo di attention standard.
Problematica dell'attention standard
Il meccanismo di self-attention standard nei transformer viene calcolato secondo la formula: dove Q, K, V sono le matrici di query, chiavi e valori.
Il problema principale di questo approccio è la complessità quadratica in termini di tempo e memoria (O(N²)) rispetto alla lunghezza della sequenza N[1]. In un'implementazione ingenua è necessario calcolare e mantenere in memoria GPU l'intera matrice di attention S di dimensione N×N, il che porta a due problemi critici:
- Elevato consumo di memoria: Memorizzare la matrice N×N diventa impossibile quando si lavora con contesti lunghi.
- Operazioni di input/output (IO): Il principale collo di bottiglia non è il numero di operazioni aritmetiche, ma i continui accessi alla memoria lenta della GPU.
Gerarchia di memoria della GPU
Per comprendere il problema è importante distinguere due tipi di memoria nella GPU (ad esempio su NVIDIA A100):
- SRAM (memoria statica): Memoria on-chip veloce di piccolo volume (~20 MB) con un'enorme larghezza di banda (fino a 19 TB/s).
- HBM (High Bandwidth Memory): Memoria lenta di grande volume (40–80 GB) con una larghezza di banda molto inferiore (circa 1,5 TB/s)[2].
Questa asimmetria rende l'algoritmo di attention standard limitato dalla larghezza di banda della memoria (memory-bound), poiché legge e scrive costantemente grandi matrici dalla lenta HBM, il che costituisce la principale fonte di latenza.
Innovazioni chiave di FlashAttention
FlashAttention è un algoritmo IO-aware che risolve il problema minimizzando gli accessi all'HBM. Ciò viene ottenuto grazie a tre tecniche principali.
Tiling ed elaborazione a blocchi
Invece di elaborare l'intera matrice in una volta, FlashAttention suddivide le matrici di input Q, K, V in piccoli blocchi (tile) che entrano nella veloce SRAM. L'algoritmo carica questi blocchi in sequenza, esegue tutti i calcoli di attention su di essi e aggiorna il risultato finale, senza salvare l'intera matrice di attention nella lenta HBM[1].
Calcolo online del Softmax
Il principale avanzamento tecnico è stato il calcolo "online" del Softmax. Il Softmax standard richiede la conoscenza di tutti gli elementi del vettore di input per la normalizzazione. FlashAttention utilizza un algoritmo modificato che consente di calcolare il Softmax per parti. Mantiene due valori intermedi (il massimo corrente e la somma degli esponenziali) che vengono aggiornati man mano che vengono elaborati nuovi blocchi, permettendo di ottenere un risultato preciso senza accedere all'intera matrice contemporaneamente[2].
Fusione delle operazioni in un unico kernel CUDA
Tutte le operazioni di attention (moltiplicazione matriciale QKᵀ, mascheramento, Softmax, moltiplicazione per V) sono unite in un unico kernel CUDA fuso (fused kernel). Ciò riduce drasticamente il numero di operazioni di lettura/scrittura nell'HBM: invece di passare più volte sull'intera matrice, l'algoritmo carica il blocco nella SRAM una volta sola, esegue tutti i calcoli e scrive solo il risultato finale.
Efficienza teorica e pratica
Complessità e ottimalità
FlashAttention riduce il consumo di memoria da O(N²) a O(N), garantendo una scalabilità lineare. È stato dimostrato che la complessità IO dell'algoritmo è teoricamente ottimale per il calcolo dell'attention in una gerarchia di memoria a due livelli, ovvero non è possibile eseguire un'attention esatta più velocemente senza modificare l'hardware[3].
Risultati empirici
La prima versione di FlashAttention ha dimostrato miglioramenti significativi:
- Accelerazione:
- BERT-large (lunghezza 512): accelerazione dell'addestramento del 15%.
- GPT-2 (lunghezza 1K): accelerazione 3 volte superiore.
- Task Long-Range Arena (1K-4K): accelerazione 2,4 volte superiore[1].
- Risparmio di memoria: Fino a 20 volte di risparmio di memoria rispetto alle implementazioni di riferimento esatte.
- Miglioramento della qualità dei modelli: Grazie alla possibilità di lavorare con contesti più lunghi, FlashAttention non solo non perde qualità, ma la migliora. Ad esempio, la perplexity di GPT-2 è migliorata di 0,7 punti, mentre la precisione nei task di classificazione di documenti lunghi è aumentata di 6,4 punti[1].
Evoluzione e sviluppi successivi
Il successo di FlashAttention ha dato origine a un'intera serie di algoritmi orientati all'hardware.
FlashAttention-2 (2023)
La seconda versione era mirata a un utilizzo più completo delle risorse della GPU. Nel FlashAttention originale l'efficienza su NVIDIA A100 era solo del 25–40% del massimo. FlashAttention-2 ha introdotto miglioramenti nel parallelismo dei calcoli, consentendo di[4]:
- Raggiungere un'accelerazione doppia rispetto alla prima versione.
- Aumentare l'utilizzo della GPU fino al 50–73% del massimo teorico.
- Estendere il supporto a head di attention di dimensione 256, nonché per le architetture Multi-Query Attention (MQA).
FlashAttention-3 (2024)
La terza versione è stata ottimizzata specificamente per l'architettura GPU NVIDIA Hopper (H100)[5]. Sfrutta nuove capacità hardware, come l'asincronicità dei Tensor Core e il supporto per FP8, il che ha permesso di:
- Raggiungere un'ulteriore accelerazione di 1,5–2 volte rispetto a FlashAttention-2.
- Raggiungere prestazioni fino a 740 TFLOPS in FP16 e vicino a 1,2 PFLOPS in FP8.
Soluzioni specializzate
Le idee di FlashAttention sono state sviluppate in altri progetti:
- FlashInfer (2025): Un motore di attention configurabile, ottimizzato specificamente per i task di inferenza LLM. Si concentra sulla gestione efficiente della KV-cache in modalità di generazione in streaming[6].
- FlashMLA (2024): Un'implementazione dell'attention con compressione della cache contestuale (latent attention), che consente di risparmiare memoria su sequenze molto lunghe con una perdita minima di informazioni[7].
Impatto sull'industria e sull'ecosistema
FlashAttention è diventato un avanzamento fondamentale e si è rapidamente affermato come standard dell'industria per l'addestramento e l'inferenza efficienti degli LLM. È stato integrato nelle principali librerie, come PyTorch e Hugging Face, ed è utilizzato dalla maggior parte dei grandi modelli linguistici (LLaMA, MPT, Falcon, Claude e altri).
Proprio FlashAttention e le sue versioni successive hanno svolto un ruolo decisivo nell'ampliamento delle finestre di contesto dei modelli linguistici: da 2–4 mila token (GPT-3) a 128 mila token (GPT-4) e persino a milioni di token nei modelli sperimentali[8]. L'algoritmo ha eliminato uno dei principali ostacoli alla scalabilità dei transformer, aprendo nuove possibilità per le applicazioni AI, dall'analisi di documenti lunghi alla comprensione multimodale.
Riferimenti
- Repository ufficiale di FlashAttention su GitHub
Bibliografia
- 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 (versione OpenReview). OpenReview H4DqfPSibmx.
- Gholami, A. et al. (2024). FlashAttention on a Napkin: A Diagrammatic Approach to Deep Learning IO-Awareness. OpenReview pF2ukh7HxA.
Note
- ↑ 1.0 1.1 1.2 1.3 1.4 Дао, Три, и др. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». arXiv:2205.14135 [cs.LG], 28 мая 2022 г. [1]
- ↑ 2.0 2.1 Дао, Три, и др. «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]
- ↑ Дао, Три. «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]