Flash attention och IO-medveten attention
Kunna förklara varför flash attention är snabbare utan att ändra matematiken.
Förkunskaper
Intuition
Naiv attention gör så här:
1. Beräkna S = QKᵀ → en N×N-matris, SKRIV till HBM
2. Läs S, beräkna softmax → ännu en N×N-matris, SKRIV till HBM
3. Läs P, beräkna PV → utdata
För N = 8 192 är varje N×N-matris 67 miljoner element — 134 MB i bf16. Den skrivs och läses flera gånger, och det är där tiden går, inte i multiplikationerna.
Flash attention gör exakt samma matematik men flyttar aldrig den stora matrisen till HBM. I stället delas Q, K och V i block som ryms i SM:ens snabba SRAM, och hela beräkningen görs där.
Resultatet är identiskt — bit för bit, inte bara ungefär. Det är ingen approximation. Det som ändras är enbart i vilken ordning och i vilket minne operationerna utförs.
| Naiv | Flash | |
|---|---|---|
| Minne | ||
| HBM-trafik | där är SRAM-storleken | |
| Matematik | exakt | exakt |
| Hastighet vid N=8k | 1× | 2–4× |
Härledning
Problemet att lösa: softmax kräver normalt att man ser alla värden innan man kan normalisera. Hur beräknar man den blockvis?
Online softmax. Anta att du redan behandlat ett block och har det löpande maxvärdet , den löpande summan och den löpande utdatan . Ett nytt block med värden kommer:
- Nytt max:
- Korrigera den gamla summan för att den beräknades med fel max:
- Korrigera utdatan på samma sätt:
Faktorn är omskalningen: allt som beräknats med det gamla maxvärdet räknas om till det nya. Eftersom softmax är translationsinvariant är resultatet exakt detsamma som om allt beräknats på en gång.
Bakåtpasset har samma problem omvänt: att spara (den stora matrisen) för bakåtpasset skulle återinföra -minnet. Lösningen är rematerialisering — spara bara och per rad, och räkna om blockvis i bakåtpasset. Det kostar extra FLOPs men sparar enormt mycket minnestrafik, och eftersom operationen är minnesbunden blir det ändå snabbare.
Vad flash attention INTE gör:
| Missuppfattning | Verkligheten |
|---|---|
| «Det är en approximation» | nej, exakt samma resultat |
| «Det ändrar komplexiteten till O(N)» | nej, FLOPs är fortfarande |
| «Det gör långa kontexter gratis» | nej, bara billigare i minne |
| «Det gör alla modeller snabbare» | bara de där attention är flaskhalsen |
Den sista är viktig: för korta sekvenser dominerar MLP-blocken, och flash attention gör då knappt någon skillnad. Vinsten växer med sekvenslängden.
Versionerna: FlashAttention-2 förbättrade parallelliseringen över sekvenslängden och nådde ungefär 70 % av teoretiskt max på A100. FlashAttention-3 utnyttjar H100:s asynkrona kopiering och fp8. I PyTorch finns det inbyggt som scaled_dot_product_attention, som väljer backend automatiskt.
Kod
import torch, torch.nn.functional as F, time
# I praktiken: använd det inbyggda anropet, som väljer flash-backend automatiskt
def attention(q, k, v, kausal=True):
return F.scaled_dot_product_attention(q, k, v, is_causal=kausal)
# Online softmax — kärnan i algoritmen, i ren Python
def online_softmax_attention(Q, K, V, blockstorlek=256):
"""Samma resultat som vanlig attention, men utan att bilda hela N×N-matrisen."""
N, d = Q.shape
O = torch.zeros_like(Q)
for i in range(0, N, blockstorlek):
q = Q[i:i + blockstorlek]
m = torch.full((len(q), 1), float("-inf")) # löpande max
l = torch.zeros((len(q), 1)) # löpande summa
o = torch.zeros((len(q), d)) # löpande utdata
for j in range(0, N, blockstorlek):
kb, vb = K[j:j + blockstorlek], V[j:j + blockstorlek]
s = q @ kb.T / d ** 0.5
m_ny = torch.maximum(m, s.max(dim=-1, keepdim=True).values)
korr = torch.exp(m - m_ny) # skala om det gamla
p = torch.exp(s - m_ny)
l = korr * l + p.sum(dim=-1, keepdim=True)
o = korr * o + p @ vb
m = m_ny
O[i:i + blockstorlek] = o / l
return O
# Verifiera att resultatet är IDENTISKT, inte ungefär
torch.manual_seed(0)
N, d = 512, 64
Q, K, V = (torch.randn(N, d) for _ in range(3))
naiv = torch.softmax(Q @ K.T / d ** 0.5, dim=-1) @ V
block = online_softmax_attention(Q, K, V, blockstorlek=128)
print("största avvikelse:", float((naiv - block).abs().max())) # ~1e-6, ren flyttalsbrus
# Minnesåtgång: N×N mot O(N)
for N in (2048, 8192, 32768):
stor_matris_mb = N * N * 2 / 1024**2
print(f"N={N:>6}: attentionmatrisen {stor_matris_mb:>9.1f} MB i bf16")
# N= 2048: attentionmatrisen 8.0 MB i bf16
# N= 8192: attentionmatrisen 128.0 MB i bf16
# N= 32768: attentionmatrisen 2048.0 MB i bf16 ← per huvud, per lager
# Vinsten växer med sekvenslängden — mät på din hårdvara
def matt(N, d=64, huvuden=8, upprepningar=20):
dev = "cuda" if torch.cuda.is_available() else "cpu"
q, k, v = (torch.randn(1, huvuden, N, d, device=dev, dtype=torch.bfloat16)
for _ in range(3))
t0 = time.perf_counter()
for _ in range(upprepningar):
F.scaled_dot_product_attention(q, k, v, is_causal=True)
if dev == "cuda":
torch.cuda.synchronize()
return round((time.perf_counter() - t0) / upprepningar * 1000, 2)
Raden N=32768 visar varför detta blev nödvändigt: två gigabyte per huvud och lager, bara för en mellanliggande matris som direkt kastas bort.
Behärskning innebär
- Förklarar varför naiv attention är minnesbunden
- Beskriver tiling och online softmax
- Vet vad flash attention inte ändrar
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — arXiv (öppen åtkomst; licens per artikel)
- arXiv — FlashAttention-2 — arXiv (öppen åtkomst; licens per artikel)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause