Minneshierarki och bandbredd
Kunna resonera om cache, VRAM och minnesbandbredd som flaskhals i inferens och träning.
Förkunskaper
Intuition
Beräkning är billigt. Att flytta data är dyrt. Det är den enskilt viktigaste insikten för prestanda på modern hårdvara.
Minneshierarkin på en GPU, ungefärligt:
| Nivå | Storlek | Bandbredd | Relativ kostnad |
|---|---|---|---|
| Register | ~kB per kärna | extremt hög | 1 |
| SRAM (delat minne) | ~200 kB per SM | ~19 TB/s | ~10 |
| HBM (VRAM) | 40–80 GB | 1–3 TB/s | ~100 |
| Värd-RAM över PCIe | hundratals GB | ~30 GB/s | ~10 000 |
| NVMe | TB | ~7 GB/s | ~100 000 |
Skillnaden mellan SRAM och HBM är ungefär en faktor tio i bandbredd — och det är den skillnaden flash attention utnyttjar.
Konsekvensen: en operation som läser mycket data och räknar lite är begränsad av minnet, inte av GPU:ns beräkningskapacitet. Att köpa ett snabbare kort hjälper då inte.
Formellt
Aritmetisk intensitet är kvoten
Varje hårdvara har en brytpunkt: en A100 klarar cirka 312 TFLOP/s i bf16 och har cirka 2,0 TB/s HBM-bandbredd, vilket ger en brytpunkt runt 156 FLOP/byte. Ligger din operation under det är den bandbreddsbunden.
| Operation | Aritmetisk intensitet | Bunden av |
|---|---|---|
| Elementvis (ReLU, add) | ~0,25 | minne |
| Layer/BatchNorm | ~1 | minne |
| Attention (naiv) | ~1–10 | minne |
| Matrismultiplikation, stor | ~100–1 000 | beräkning |
| Generering, en token | ~1 | minne |
Den sista raden förklarar varför tokengenerering är långsam: varje token kräver att alla modellvikter läses från HBM, men bara en vektor–matris-produkt beräknas. Genomströmningen begränsas därför av
Roofline-modellen sätter ihop det: prestandan är . Att rita sin operation i det diagrammet säger direkt vad man ska optimera.
Tre optimeringar som följer av detta:
| Optimering | Idé |
|---|---|
| Kernel fusion | slå ihop flera elementvisa operationer så att data läses en gång i stället för tre |
| Tiling | dela upp i block som ryms i SRAM och återanvänd dem |
| Batchning | läs vikterna en gång, använd dem för många exempel |
Batchning är den mest dramatiska vid generering: vikterna läses en gång oavsett om batchen är 1 eller 64, så genomströmningen växer nästan linjärt tills beräkningen blir flaskhals.
Detta förklarar också varför kvantisering hjälper så mycket vid inferens. Att gå från bf16 till 4 bitar gör modellen fyra gånger mindre, och eftersom hastigheten är bandbreddsbunden blir genereringen ungefär fyra gånger snabbare — inte för att beräkningen är snabbare, utan för att det är mindre att läsa.
Kod
import torch, time
def aritmetisk_intensitet(flops, bytes_flyttade):
return flops / bytes_flyttade
# Elementvis operation på en stor tensor
n = 100_000_000
flops = n # en operation per element
byte = n * 2 * 2 # läs + skriv, bf16
print(f"ReLU: I = {aritmetisk_intensitet(flops, byte):.2f} FLOP/byte → minnesbunden")
# Matrismultiplikation N×N
for N in (128, 1024, 8192):
flops = 2 * N ** 3
byte = 3 * N ** 2 * 2 # två indata + ett utdata, bf16
print(f"matmul {N:>5}: I = {aritmetisk_intensitet(flops, byte):>8.1f} FLOP/byte")
# matmul 128: I = 42.7 FLOP/byte → minnesbunden
# matmul 1024: I = 341.3 FLOP/byte → beräkningsbunden
# matmul 8192: I = 2730.7 FLOP/byte → tydligt beräkningsbunden
# Brytpunkt för en given hårdvara
def brytpunkt(tflops, tb_per_s):
return tflops * 1e12 / (tb_per_s * 1e12)
for namn, tf, bw in (("A100 bf16", 312, 2.0), ("H100 bf16", 990, 3.35), ("laptop-GPU", 20, 0.3)):
print(f"{namn:<11} brytpunkt {brytpunkt(tf, bw):>6.0f} FLOP/byte")
# A100 bf16 brytpunkt 156 FLOP/byte
# H100 bf16 brytpunkt 296 FLOP/byte
# laptop-GPU brytpunkt 67 FLOP/byte
# Generering: tak från bandbredd, inte från beräkning
def genereringstak(modell_gb, bandbredd_gb_s):
return bandbredd_gb_s / modell_gb
for namn, gb, bw in (("7B bf16", 14, 2000), ("7B int4", 3.5, 2000), ("70B int4", 35, 2000)):
print(f"{namn:<9} tak ~{genereringstak(gb, bw):>6.0f} tokens/s")
# 7B bf16 tak ~ 143 tokens/s
# 7B int4 tak ~ 571 tokens/s ← 4× mindre modell, 4× snabbare
# 70B int4 tak ~ 57 tokens/s
# Kernel fusion: tre operationer, en läsning i stället för tre
x = torch.randn(10_000_000, device="cuda" if torch.cuda.is_available() else "cpu")
def ofuserat(x):
return torch.relu(x * 2.0 + 1.0)
fuserat = torch.compile(ofuserat) # torch.compile fuserar automatiskt
Raden 7B int4 är den mest praktiska konsekvensen: kvantisering ger fyrfaldig hastighet vid generering, inte för att beräkningen går snabbare utan för att det är en fjärdedel så mycket att läsa.
Behärskning innebär
- Beskriver minneshierarkin och dess kostnader
- Räknar ut aritmetisk intensitet
- Avgör om en operation är bandbredds- eller beräkningsbunden
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)
- NVIDIA — Deep Learning Performance Guide — dokumentation, fri läsning
- PyTorch — tutorials (BSD-3) — BSD-3-Clause