AI-modeller og platforme
Flash Attention: Revolutionering af Transformer Effektivitet

Af
Aayush Mittal Mittal
Da transformer-modeller vokser i størrelse og kompleksitet, står de over for betydelige udfordringer i forhold til beregnings-effektivitet og hukommelsesbrug, især når det handler om lange sekvenser. Flash Attention er en optimeringsteknik, der lover at revolutionere måden, vi implementerer og skalerer opmærksomheds-mekanismer i Transformer-modeller.
I denne omfattende vejledning, dykker vi dybt ind i Flash Attention, hvor vi udforsker dets kernebegreber, implementationsdetaljer og den dybe indvirkning, det har på feltet machine learning.
Problemet: Opmærksomhed Er Dyrt
Før vi dykker ind i løsningen, lad os først forstå problemet, som Flash Attention sigter mod at løse. Opmærksomheds-mekanismen, selvom den er kraftfuld, kommer med en betydelig beregnings-omkostning, især for lange sekvenser.
Standard Opmærksomhed: En Hurtig Gennemgang
Den standard opmærksomheds-mekanisme i Transformer-modeller kan sammenfattes med følgende ligning:
Opmærksomhed(Q, K, V) = softmax(QK^T / √d) VHvor Q, K og V er henholdsvis Query-, Key- og Value-matricer, og d er dimensionen af nøgle-vektorerne.
Selvom denne formulering er elegant, fører dens implementering til flere ineffektiver:
- Hukommelses-Flaskehals: Den intermediate opmærksomheds-matrix (QK^T) har en størrelse på N x N, hvor N er sekvens-længden. For lange sekvenser kan dette hurtigt udtømme tilgængelig GPU-hukommelse.
- Redundant Hukommelses-adgang: I standard-implementeringer beregnes opmærksomheds-matrixen, gemmes i high-bandwidth-hukommelse (HBM) og læses derefter tilbage for softmax-operationen. Denne redundante hukommelses-adgang er en større flaskehals.
- Under-udnyttelse af GPU-Beregning: Moderne GPU’er har betydeligt mere beregnings-kapacitet (FLOPS) end hukommelses-båndbredde. Den standard opmærksomheds-implementering er hukommelses-bunden, hvilket efterlader meget af GPU’ens beregnings-potentiale utilgængeligt.
Lad os illustrere dette med et simpelt Python-kode-udsnit, der viser den standard opmærksomheds-implementering:
</code> import torch <p>def standard_opmærksomhed(Q, K, V):</p> <p># Q, K, V form: (batch_størrelse, sekvens_længde, d_model)</p> <p>d_k = K.size(-1)</p> <p>score = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k))</p> <p>opmærksomheds_vægt = torch.softmax(score, dim=-1)</p> <p>return torch.matmul(opmærksomheds_vægt, V)</p>
Denne implementering, selvom den er retlinet, lider under de ineffektiver, der er nævnt ovenfor. score-tensoren, der har formen (batch_størrelse, sekvens_længde, sekvens_længde), kan blive forbudt stor for lange sekvenser.
Introduktion til Flash Opmærksomhed
Flash Opmærksomhed, introduceret af Tri Dao og kolleger i deres artikel fra 2022, er en tilgang til beregning af opmærksomhed, der dramatisk reducerer hukommelses-brug og forbedrer beregnings-effektivitet. De nøgle-idéer bag Flash Opmærksomhed er:
- Tilføjelse: Opdel den store opmærksomheds-matrix i mindre blokke, der passer i hurtig on-chip SRAM.
- Gen-beregning: I stedet for at gemme den hele opmærksomheds-matrix, gen-beregner dele af den under baglæns-passagen.
- IO-Bevidst Implementering: Optimer algoritmen for at minimere data-bevægelse mellem forskellige niveauer af GPU-hukommelses-hierarkiet.
Flash Opmærksomheds Algoritmen
I sin kerne gen-tænker Flash Opmærksomhed, hvordan vi beregner opmærksomheds-mekanismen. I stedet for at beregne den hele opmærksomheds-matrix på én gang, behandler den den i blokke, udnyttende GPU’ens hukommelses-hierarki.
Her er en overordnet gennemgang af algoritmen:
- Input: Matricer Q, K, V i HBM (High Bandwidth Memory) og on-chip SRAM af størrelse M.
- Block-størrelser beregnes baseret på tilgængelig SRAM.
- Initialisering af output-matrix O og hjælpe-vektorer l og m.
- Algoritmen opdeler input-matricer i blokke for at passe i SRAM.
- To indbyrdes løkker behandler disse blokke:
- Ydre løkke indlæser K og V-blokke
- Indre løkke indlæser Q-blokke og udfører beregninger
- On-chip-beregninger inkluderer matrix-multiplication, softmax og output-beregning.
- Resultaterne skrives tilbage til HBM efter behandling af hvert block.
Denne block-vis beregning tillader Flash Opmærksomhed at opretholde en meget mindre hukommelses-fodaftryk, samtidig med at den stadig beregner den nøjagtige opmærksomhed.
Matematikken Bag Flash Opmærksomhed
Nøglen til at gøre Flash Opmærksomhed til at virke er en matematisk trick, der tillader os at beregne softmax på en block-vis måde. Artiklen introducerer to nøgle-formler:
- Softmax-dekomposition:
softmax(x) = exp(x - m) / Σexp(x - m)hvor m er den maksimale værdi i x.
- Softmax-sammenlægning:
softmax(x ∪ y) = softmax(softmax(x) * e^(m_x - m), softmax(y) * e^(m_y - m))hvor m = max(m_x, m_y)
Disse formler tillader Flash Opmærksomhed at beregne delvise softmax-resultater for hvert block og derefter kombinere dem korrekt for at få det endelige resultat.
Implementeringsdetaljer
Lad os dykke ind i en simplificeret implementering af Flash Opmærksomhed for at illustrere dets kernebegreber:
import torch
<p>def flash_opmærksomhed(Q, K, V, block_størrelse=256):</p>
<p># Initialisering af output og kørende statistik</p>
<p>O = torch.zeros_like(Q)</p>
<p>L = torch.zeros((batch_størrelse, sekvens_længde, 1))</p>
<p>M = torch.full((batch_størrelse, sekvens_længde, 1), float('-inf'))</p>
<p>for i in range(0, sekvens_længde, block_størrelse):</p>
<p>Q_block = Q[:, i:i+block_størrelse, :]</p>
<p>for j in range(0, sekvens_længde, block_størrelse):</p>
<p>K_block = K[:, j:j+block_størrelse, :]</p>
<p>V_block = V[:, j:j+block_størrelse, :]</p>
<p># Beregn opmærksomheds-score for dette block</p>
<p>S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p>
<p># Opdaterer kørende maks</p>
<p>M_new = torch.maximum(M[:, i:i+block_størrelse], S_block.max(dim=-1, keepdim=True).values)</p>
<p># Beregn eksponentialer</p>
<p>exp_S = torch.exp(S_block - M_new)</p>
<p>exp_M_diff = torch.exp(M[:, i:i+block_størrelse] - M_new)</p>
<p># Opdaterer kørende sum</p>
<p>L_new = exp_M_diff * L[:, i:i+block_størrelse] + exp_S.sum(dim=-1, keepdim=True)</p>
<p># Beregn output for dette block</p>
<p>O[:, i:i+block_størrelse] = (</p>
<p>exp_M_diff * O[:, i:i+block_størrelse] +</p>
<p>torch.matmul(exp_S, V_block)</p>
<p>) / L_new</p>
<p># Opdaterer kørende statistik</p>
<p>L[:, i:i+block_størrelse] = L_new</p>
<p>M[:, i:i+block_størrelse] = M_new</p>
<p>return O</p>
Denne implementering, selvom den er simplificeret, fanger essensen af Flash Opmærksomhed. Den behandler input i blokke, mens den opretholder kørende statistik (M og L) for at korrekt beregne softmax over alle blokke.
Flash Opmærksomheds Indvirkning
Introduktionen af Flash Opmærksomhed har haft en dyb indvirkning på feltet machine learning, især for store sprog-modeller og lange-kontekst-anvendelser. Nogle af de nøgle-fordele inkluderer:
- Reduceret Hukommelses-Brug: Flash Opmærksomhed reducerer hukommelses-kompleksiteten fra O(N^2) til O(N), hvor N er sekvens-længden. Dette tillader behandling af langt længere sekvenser med samme hardware.
- Forbedret Hastighed: Ved at minimere data-bevægelse og bedre udnytte GPU-beregningsevne, opnår Flash Opmærksomhed betydelige hastighedsforbedringer. Forfatterne rapporterer op til 3 gange hurtigere træning for GPT-2 i forhold til standard-implementeringer.
- Nøjagtig Beregning: I modsætning til andre opmærksomheds-optimeringsteknikker beregner Flash Opmærksomhed den nøjagtige opmærksomhed, ikke en approksimation.
- Skalérbarhed: Den reducerede hukommelses-fodaftryk tillader skalering til langt længere sekvenser, potentielt op til millioner af tokens.
Reel Verden-Indvirkning
Flash Opmærksomheds indvirkning strækker sig langt ud over akademisk forskning. Den er hurtigt blevet adopteret i mange populære machine learning-biblioteker og -modeller:
- Hugging Face Transformers: Det populære Transformers-bibliotek har integreret Flash Opmærksomhed, hvilket tillader brugere at let udnytte dets fordele.
- GPT-4 og derefter: Selvom det ikke er bekræftet, er der spekulationer om, at avancerede sprog-modeller som GPT-4 muligvis kan bruge teknikker lignende Flash Opmærksomhed til at håndtere lange kontekster.
- Lange-Kontekst-Modeller: Flash Opmærksomhed har muliggjort en ny generation af modeller, der kan håndtere ekstremt lange kontekster, såsom modeller, der kan behandle hele bøger eller lange videoer.
FlashOpmærksomhed: Seneste Udviklinger
FlashOpmærksomhed-2
Bygget på succesen af den originale Flash Opmærksomhed, introducerede det samme team FlashOpmærksomhed-2 i 2023. Denne opdaterede version bringer flere forbedringer:
- Yderligere Optimering: FlashOpmærksomhed-2 opnår endnu bedre GPU-udnyttelse, op til 70% af den teoretiske peak FLOPS på A100 GPU’er.
- Forbedret Baglæns-Passering: Baglæns-passagen er optimeret til at være næsten lige så hurtig som forløbs-passagen, hvilket fører til betydelige hastighedsforbedringer.
- Understøttelse af Forskellige Opmærksomheds-Variationer: FlashOpmærksomhed-2 udvider understøttelsen til forskellige opmærksomheds-variationer, herunder gruppe-forespørgsels-opmærksomhed og multi-forespørgsels-opmærksomhed.
FlashOpmærksomhed-3
Udgivet i 2024, FlashOpmærksomhed-3 repræsenterer den seneste fremgang i denne forskningslinje. Den introducerer flere nye teknikker til at yderligere forbedre ydelsen:
- Asynkron Beregning: Udvinder den asynkrone natur af nye GPU-instruktioner til at overlappe forskellige beregninger.
- FP8-Understøttelse: Udnytter lav-præcisions FP8-beregning for endnu hurtigere behandling.
- Ukoherent Behandling: En teknik til at reducere kvantiserings-fejl, når lav-præcisions-formater bruges.
Her er et forenklet eksempel på, hvordan FlashOpmærksomhed-3 muligvis kan udnytte asynkron beregning:
import torch from torch.cuda.amp import autocast <p>def flash_opmærksomhed_3(Q, K, V, block_størrelse=256):</p> <p>med autocast(dtype=torch.float8): # Brug FP8 til beregning</p> <p># ... (lignende med tidligere implementering)</p> <p># Asynkron beregningseksempel</p> <p>med torch.cuda.stream(torch.cuda.Stream()):</p> <p># Beregn GEMM asynkront</p> <p>S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p> <p># I mellemtiden, på standard-streamen:</p> <p># Forbered til softmax-beregning</p> <p># Synkroniser streams</p> <p>torch.cuda.synchronize()</p> <p># Fortsæt med softmax og output-beregning</p> <p># ...</p> return O
Dette kode-udsnit illustrerer, hvordan FlashOpmærksomhed-3 muligvis kan udnytte asynkron beregning og FP8-præcision. Bemærk, at dette er et forenklet eksempel, og den faktiske implementering ville være langt mere kompleks og hardware-specifik.
Implementering af Flash Opmærksomhed i Dine Projekter
Hvis du er begejstret for at udnytte Flash Opmærksomhed i dine egne projekter, har du flere muligheder:
- Brug Eksisterende Biblioteker: Mange populære biblioteker som Hugging Face Transformers inkluderer allerede Flash Opmærksomhed-implementeringer. En simpel opdatering til den seneste version og aktivering af de relevante flags kan være tilstrækkeligt.
- Tilpasset Implementering: For mere kontrol eller specialiserede anvendelser kan du selv implementere Flash Opmærksomhed. xformers-biblioteket giver en god reference-implementering.
- Hardware-Specifikke Optimeringer: Hvis du arbejder med specifik hardware (f.eks. NVIDIA H100 GPU’er), kan du udnytte hardware-specifikke funktioner for maksimal ydelse.
Her er et eksempel på, hvordan du kan bruge Flash Opmærksomhed med Hugging Face Transformers-biblioteket:
from transformers import AutoModel, AutoConfig
<p># Aktiver Flash Opmærksomhed</p>
<p>config = AutoConfig.from_pretrained("bert-base-uncased")</p>
<p>config.use_flash_opmærksomhed = True</p>
<p># Indlæs model med Flash Opmærksomhed</p>
<p>model = AutoModel.from_pretrained("bert-base-uncased", config=config)</p>
<p># Brug modellen som normalt</p>
<p># ...
Udfordringer og Fremtidige Retninger
Selvom Flash Opmærksomhed har gjort betydelige fremskridt i forbedring af opmærksomheds-mekanismens effektivitet, er der stadig udfordringer og områder for fremtidig forskning:
- Hardware-Specifik: Nuverende implementeringer er ofte optimeret for specifikke GPU-arkitekturer. At generalisere disse optimeringer på tværs af forskellige hardware er en udfordring.
- Integration med Andre Teknikker: At kombinere Flash Opmærksomhed med andre optimeringsteknikker som pruning, kvantificering og model-kompression er et aktivt forskningsområde.
- Udvidelse til Andre Domæner: Selvom Flash Opmærksomhed har vist stor succes i NLP, er udvidelse af dens fordele til andre domæner som computer-vision og multimodale modeller en pågående bestræbelse.
- Teoretisk Forståelse: At dybe vores teoretiske forståelse af, hvorfor Flash Opmærksomhed virker så godt, kunne føre til endnu mere kraftfulde optimeringer.
Konklusion
Ved at udnytte GPU-hukommelses-hierarkier og anvende matematiske tricks, opnår Flash Opmærksomhed betydelige forbedringer i både hastighed og hukommelses-brug uden at gå på kompromis med nøjagtigheden.
Som vi har udforsket i denne artikel, strækker Flash Opmærksomheds indvirkning sig langt ud over en simpel optimeringsteknik. Den har muliggjort udviklingen af mere kraftfulde og effektive modeller.
Jeg har brugt de sidste fem år på at dykke ned i den fascinerende verden af Machine Learning og Deep Learning. Min passion og ekspertise har ført mig til at bidrage til over 50 forskellige software-ingeniørprojekter, med en særlig fokus på AI/ML. Min fortsatte nysgerrighed har også ført mig mod Natural Language Processing, et felt jeg er ivrig efter at udforske yderligere.
Opdag mere


Din KV-cache har ikke et bit-problem. Den har et geometri-problem.


Den Intensiverede AI-Våbenkapløb: AMD’s Strategiske Partnerskab med OpenAI


Hemmeligheden bag hurtigere AI er ikke flere GPU’er, men smartere netværk


Hvorfor store sprogmodeller glemmer midten: Afsløring af AI’s skjulte blindplet


NVIDIA Udsender Hotfix til GPU-Drivers Overophedingsproblemer


Sapiens: Gennembrud i Menneskelige Visionmodeller

