FlashAttention (NL)

From Systems analysis Wiki
Jump to navigation Jump to search

FlashAttention — dit is een revolutionair algoritme voor de berekening van het aandachtsmechanisme (attention), ontwikkeld voor een significante versnelling van het trainen en de inferentie van grote taalmodellen (LLM) met behoud van volledige rekennauwkeurigheid. Het algoritme werd voor het eerst gepresenteerd in 2022 door een team van onderzoekers van Stanford University onder leiding van Tri Dao[1].

Het kernidee van FlashAttention is de herorganisatie van berekeningen met inachtneming van de geheugenhiërarchie van de GPU, waardoor het aantal toegangen tot het langzame geheugen wordt geminimaliseerd en het belangrijkste knelpunt van het standaard aandachtsmechanisme wordt weggenomen.

Problemen met standaard attention

Het standaard zelfaandachtsmechanisme in transformers wordt berekend via de formule: Attention(Q,K,V)=softmax(QKTdk)V waar Q, K, V de matrices van queries, keys en values zijn.

Het belangrijkste probleem van deze aanpak is de kwadratische complexiteit in tijd en geheugen (O(N²)) ten opzichte van de sequentielengte N[1]. Bij een naïeve implementatie moet de volledige aandachtsmatrix S van grootte N×N worden berekend en opgeslagen in het GPU-geheugen, wat leidt tot twee kritieke problemen:

  1. Hoog geheugenverbruik: Het opslaan van de N×N-matrix wordt onmogelijk bij het werken met lange contexten.
  2. Invoer-uitvoerbewerkingen (IO): Het belangrijkste knelpunt is niet het aantal rekenkundige bewerkingen, maar de constante toegangen tot het langzame GPU-geheugen.

Geheugenhiërarchie van de GPU

Voor een goed begrip van het probleem is het belangrijk onderscheid te maken tussen twee geheugentypen in een GPU (aan de hand van de NVIDIA A100):

  • SRAM (statisch geheugen): Snel on-chip geheugen met een kleine capaciteit (~20 MB) en een enorme bandbreedte (tot 19 TB/s).
  • HBM (High Bandwidth Memory): Traag geheugen met grote capaciteit (40–80 GB) en een veel lagere bandbreedte (circa 1,5 TB/s)[2].

Deze asymmetrie maakt het standaard aandachtsalgoritme geheugenbandbreedte-gebonden (memory-bound), omdat het voortdurend grote matrices leest en schrijft vanuit de trage HBM, wat de voornaamste bron van vertraging is.

Belangrijkste innovaties van FlashAttention

FlashAttention is een IO-bewust (IO-aware) algoritme dat het probleem oplost door het aantal toegangen tot HBM te minimaliseren. Dit wordt bereikt door middel van drie hoofdtechnieken.

Tiling en blokgewijze verwerking

In plaats van de volledige matrix in één keer te verwerken, verdeelt FlashAttention de invoermatrices Q, K, V in kleine blokken (tiles) die in de snelle SRAM passen. Het algoritme laadt deze blokken sequentieel, voert alle aandachtsberekeningen erop uit en werkt het eindresultaat bij, zonder de volledige aandachtsmatrix in de trage HBM op te slaan[1].

Online berekening van Softmax

Een belangrijke technische doorbraak was de "online" berekening van Softmax. Standaard Softmax vereist kennis van alle elementen van de invoervector voor normalisatie. FlashAttention maakt gebruik van een aangepast algoritme waarmee Softmax in delen kan worden berekend. Het houdt twee tussenliggende waarden bij (het huidige maximum en de som van de exponenten), die worden bijgewerkt naarmate nieuwe blokken worden verwerkt, zodat een nauwkeurig resultaat wordt verkregen zonder toegang tot de volledige matrix tegelijk[2].

Samensmelting van bewerkingen in één CUDA-kernel

Alle aandachtsbewerkingen (matrixvermenigvuldiging QKᵀ, maskering, Softmax, vermenigvuldiging met V) worden samengevoegd in één gefuseerde CUDA-kernel (fused kernel). Dit vermindert drastisch het aantal lees-/schrijfbewerkingen naar HBM: in plaats van meerdere doorgangen over de volledige matrix laadt het algoritme een blok eenmalig in SRAM, voert alle berekeningen uit en schrijft alleen het eindresultaat weg.

Theoretische en praktische efficiëntie

Complexiteit en optimaliteit

FlashAttention verlaagt het geheugenverbruik van O(N²) naar O(N), wat zorgt voor lineaire schaalbaarheid. Er is aangetoond dat de IO-complexiteit van het algoritme theoretisch optimaal is voor het berekenen van attention in een tweeniveaus geheugenhiërarchie, wat betekent dat snellere exacte attention zonder aanpassing van de hardware onmogelijk is[3].

Empirische resultaten

De eerste versie van FlashAttention liet significante verbeteringen zien:

  • Versnelling:
    • BERT-large (lengte 512): 15% versnelling van de training.
    • GPT-2 (lengte 1K): 3-voudige versnelling.
    • Long-Range Arena-taken (1K-4K): 2,4-voudige versnelling[1].
  • Geheugenbesparing: Tot 20-voudige geheugenbesparing vergeleken met exacte basisimplementaties.
  • Verbetering van modelkwaliteit: Dankzij de mogelijkheid om met langere contexten te werken, verliest FlashAttention niet alleen geen kwaliteit, maar verbetert het deze zelfs. Zo verbeterde de perplexiteit van GPT-2 met 0,7 punten en steeg de nauwkeurigheid bij classificatietaken voor lange documenten met 6,4 punten[1].

Evolutie en verdere ontwikkelingen

Het succes van FlashAttention leidde tot een hele reeks hardware-georiënteerde algoritmen.

FlashAttention-2 (2023)

De tweede versie was gericht op een vollediger gebruik van GPU-resources. In het originele FlashAttention bedroeg de efficiëntie op de NVIDIA A100 slechts 25–40% van het maximum. FlashAttention-2 introduceerde verbeteringen in de parallelisatie van berekeningen, waardoor[4]:

  • Een tweevoudige versnelling werd bereikt ten opzichte van de eerste versie.
  • De GPU-benutting werd verhoogd naar 50–73% van het theoretische maximum.
  • Ondersteuning werd uitgebreid naar aandachtshoofden van grootte 256, evenals voor Multi-Query Attention (MQA)-architecturen.

FlashAttention-3 (2024)

De derde versie werd specifiek geoptimaliseerd voor de GPU-architectuur NVIDIA Hopper (H100)[5]. Het maakt gebruik van nieuwe hardwaremogelijkheden, zoals asynchrone Tensor Cores en ondersteuning voor FP8, waardoor:

  • Een verdere 1,5–2-voudige versnelling werd bereikt ten opzichte van FlashAttention-2.
  • Een prestatie tot 740 TFLOPS op FP16 en dicht bij 1,2 PFLOPS op FP8 werd gehaald.

Gespecialiseerde oplossingen

De ideeën van FlashAttention zijn verder uitgewerkt in andere projecten:

  • FlashInfer (2025): Een aanpasbare attention-engine, specifiek geoptimaliseerd voor LLM-inferentietaken. De focus ligt op efficiënte verwerking van de KV-cache in de modus van streaming generatie[6].
  • FlashMLA (2024): Een implementatie van attention met compressie van de contextuele cache (latent attention), waarmee geheugen wordt bespaard bij zeer lange sequenties met minimaal informatieverlies[7].

Impact op de industrie en het ecosysteem

FlashAttention is een fundamentele doorbraak geworden en is snel uitgegroeid tot de industriestandaard voor efficiënt trainen en inferentie van LLM's. Het is geïntegreerd in belangrijke bibliotheken zoals PyTorch en Hugging Face en wordt gebruikt in de meeste grote taalmodellen (LLaMA, MPT, Falcon, Claude en anderen).

Juist FlashAttention en zijn opvolgende versies speelden een doorslaggevende rol bij het vergroten van de contextvensters van taalmodellen: van 2–4 duizend tokens (GPT-3) tot 128 duizend tokens (GPT-4) en zelfs tot miljoenen tokens in experimentele modellen[8]. Het algoritme heeft een van de belangrijkste obstakels voor de schaalvergroting van transformers weggenomen en nieuwe mogelijkheden geopend voor AI-toepassingen, van de analyse van lange documenten tot multimodaal begrip.

Verwijzingen

  • Officiële FlashAttention-repository op GitHub

Literatuur

  • 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.

Noten

  1. 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. 2.0 2.1 Дао, Три, и др. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». OpenReview. [2]
  3. «We're Training AI Twice as Fast This Year as Last». IEEE Spectrum. [3]
  4. Дао, Три. «FlashAttention-2». tridao.me. [4]
  5. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». PyTorch Blog. [5]
  6. «[2501.01005] FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving». arXiv. [6]
  7. «GitHub - deepseek-ai/FlashMLA: FlashMLA: Efficient MLA decoding kernels». GitHub. [7]
  8. «The Evolution of Flash Attention: Revolutionizing Transformer Efficiency». Medium. [8]