Hoppa till innehållet
AI-grafen
F· AI engineeringtransformerarkitektur· ca 90 min· volatil — kontrolleras ofta· verifierad 2026-09-21

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.

NaivFlash
MinneO(N2)O(N^2)O(N)O(N)
HBM-trafikO(N2)O(N^2)O(N2/M)O(N^2/M) där MM är SRAM-storleken
Matematikexaktexakt
Hastighet vid N=8k1×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 mm, den löpande summan ℓ\ell och den löpande utdatan OO. Ett nytt block med värden xx kommer:

  1. Nytt max: m′=max⁡(m,max⁡x)m' = \max(m, \max x)
  2. Korrigera den gamla summan för att den beräknades med fel max:

ℓ′=em−m′ℓ+∑iexi−m′\ell' = e^{m - m'}\ell + \sum_i e^{x_i - m'}

  1. Korrigera utdatan på samma sätt:

O′=em−m′ℓℓ′O+1ℓ′∑iexi−m′viO' = \frac{e^{m-m'}\ell}{\ell'}O + \frac{1}{\ell'}\sum_i e^{x_i - m'}v_i

Faktorn em−m′e^{m-m'} ä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 PP (den stora matrisen) för bakåtpasset skulle återinföra O(N2)O(N^2)-minnet. Lösningen är rematerialisering — spara bara mm och ℓ\ell per rad, och räkna om PP 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:

MissuppfattningVerkligheten
«Det är en approximation»nej, exakt samma resultat
«Det ändrar komplexiteten till O(N)»nej, FLOPs är fortfarande O(N2)O(N^2)
«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

Alla källor och licenser