USAAIO Lesson 91, from Phase 3. It covers Flash Attention's IO complexity, which is O(N) in HBM traffic rather than O(N²), tiled softmax with an online running maximum, gradient checkpointing, the ALiBi position bias, and Multi-Query Attention. All the numbers were verified with torch 2.7.1+cpu: the tiled output matches standard attention to a maximum absolute error of 6e-8, and the MQA KV cache is eight times smaller than MHA at N=2048. The lesson runs to 32 slides.
Subject: Machine Learning · 56 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 91 · Phase 3
Why standard attention is IO-bound, not compute-bound — and four surgical fixes: Flash Attention tiled SRAM kernels, gradient checkpointing, ALiBi position bias, and Multi-Query Attention.
Objectives
−m·|i−j| for given slopes and explain why it extrapolates beyond training lengthWarm-up
Discussion prompt
Before we open Lesson 91: Flash Attention & Memory-Efficient Transformers: without looking back, what was the main idea of KV Cache & Attention Efficiency, and what could you do by the end of it that you could not do before?
Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.
Answer:
KV cache for autoregressive inference (O(n·d) memory, O(1) per step), Flash Attention block-wise computation (O(n) SRAM), sliding window attention (O(n·w) vs O(n²)), sparse attention patterns (BigBird/Longformer), and Group Query Attention (GQA/MQA).
Section
Part 1 of 4
Concept
Modern GPUs have fast on-chip SRAM (≈ 20 MB, ~19 TB/s) and slow off-chip HBM (≈ 80 GB, ~2 TB/s). Standard attention is bottlenecked by how often it writes and reads the N×N score matrix to/from HBM — not by multiply-adds.
| memory tier | capacity | bandwidth | role |
|---|---|---|---|
| SRAM (on-chip) | ~20 MB | ~19 TB/s | registers, shared mem — fast but tiny |
| HBM (off-chip) | ~80 GB | ~2 TB/s | activations, weights — slow but large |
Arithmetic intensity of attention is low — the bottleneck is data movement, not floating-point ops. Flash Attention eliminates HBM round-trips by never materializing the N×N matrix.
Comparison
Comparison matrix
From Where the time goes: HBM vs SRAM: refill the capacity column from what you know. The rest of the table is as it appeared.
| memory tier | capacity | bandwidth | role |
|---|---|---|---|
| SRAM (on-chip) | ~20 MB | ~19 TB/s | registers, shared mem — fast but tiny |
| HBM (off-chip) | ~80 GB | ~2 TB/s | activations, weights — slow but large |
Concept
Standard attention writes the N×N score matrix S to HBM, reads it back for softmax, writes the probability matrix P, reads it back for P·V. That is four passes over an N² buffer.
\[ \text{HBM accesses}_{\text{standard}} \approx 3Nd + 4N^2 = O(N^2) \]
| N (seq len) | d=64, standard HBM (elems) | Flash approx (elems) | ratio |
|---|---|---|---|
| 512 | 1,146,880 | 131,072 | 8.8× |
| 1024 | 4,390,912 | 262,144 | 16.8× |
| 2048 | 17,170,432 | 524,288 | 32.8× |
| 4096 | 67,895,296 | 1,048,576 | 64.8× |
Trade off
Comparison matrix
From Standard attention IO: O(N²) HBM accesses: every row here is a choice with a cost. Fill the ratio column, then say which row you would actually pick and what you give up for it.
| N (seq len) | d=64, standard HBM (elems) | Flash approx (elems) | ratio |
|---|---|---|---|
| 512 | 1,146,880 | 131,072 | 8.8× |
| 1024 | 4,390,912 | 262,144 | 16.8× |
| 2048 | 17,170,432 | 524,288 | 32.8× |
| 4096 | 67,895,296 | 1,048,576 | 64.8× |
Intuition
Flash Attention cuts the matrix into tiles small enough to fit in SRAM. It computes a tile of Q·Kᵀ, immediately applies softmax and multiplies by V — all on-chip. The N×N matrix is never written to HBM.
The trick is an online softmax: you can fuse the normalization with the tile pass using a running maximum m and a running sum l. The final output is correct even though you never saw the full row at once.
Backward pass: Flash Attention 2 recomputes attention tiles during the backward instead of storing them — gradient checkpointing at the op level. Net: O(N) HBM accesses, same FLOPS as standard attention.
Estimation
Predict first
For tile (i,j), compute the block score S_ij = Q_i · K_j^T · scale. Update the running max, sum, and partial output, then rescale to merge the new tile.
Commit before you compute: what does Online softmax: tiled update equations come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: max abs diff = 5.96e-08 — numerically identical to standard attention
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The online softmax rescales past partial sums each time a new tile yields a larger maximum, preserving exact (float32) equivalence.
Worked example
For tile (i,j), compute the block score S_ij = Q_i · K_j^T · scale. Update the running max, sum, and partial output, then rescale to merge the new tile.
import torch, torch.nn.functional as F
torch.manual_seed(42)
N, d = 8, 4
Q = torch.randn(N, d); K = torch.randn(N, d); V = torch.randn(N, d)
scale = d ** -0.5
# Reference standard attention
O_std = F.softmax((Q @ K.T) * scale, dim=-1) @ V
# Tiled Flash Attention (BLOCK=4)
BLOCK = 4
O = torch.zeros_like(Q)
l = torch.zeros(N) # running normalizer
m = torch.full((N,), float('-inf')) # running max
for j in range(0, N, BLOCK):
Kj, Vj = K[j:j+BLOCK], V[j:j+BLOCK]
for i in range(0, N, BLOCK):
Sij = (Q[i:i+BLOCK] @ Kj.T) * scale # (BLOCK,BLOCK)
mij = Sij.max(dim=-1).values
m_new = torch.maximum(m[i:i+BLOCK], mij)
e = torch.exp(Sij - m_new.unsqueeze(1)) # re-center
l_new = torch.exp(m[i:i+BLOCK] - m_new)*l[i:i+BLOCK] + e.sum(dim=-1)
O[i:i+BLOCK] = (torch.diag(torch.exp(m[i:i+BLOCK]-m_new)) @ O[i:i+BLOCK]
+ e @ Vj)
m[i:i+BLOCK] = m_new; l[i:i+BLOCK] = l_new
O = O / l.unsqueeze(1) # final normalize
print('max abs diff:', (O_std - O).abs().max().item())max abs diff = 5.96e-08 — numerically identical to standard attention
Why: The online softmax rescales past partial sums each time a new tile yields a larger maximum, preserving exact (float32) equivalence. No approximation.
| variable | shape | what it tracks |
|---|---|---|
| m[i:i+B] | (BLOCK,) | per-query running max of logits seen so far |
| l[i:i+B] | (BLOCK,) | per-query running sum of exp(logit − max) |
| O[i:i+B] | (BLOCK,d) | accumulated unnormalized output (rescaled each tile) |
| e | (BLOCK,BLOCK) | exp(S_ij − m_new) for this tile |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
max abs diff = 5.96e-08 — numerically identical to standard attention
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
For tile (i,j), compute the block score S_ij = Q_i · K_j^T · scale. Update the running max, sum, and partial output, then rescale to merge the new tile.
Anomaly
Predict first
A student writes this, and it looks reasonable:
Flash Attention makes attention faster by reducing the number of floating-point multiply-adds.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Flash Attention computes the same O(N²d) multiply-adds.
Flash Attention reduces HBM accesses from O(N²) to O(N) — FLOPS are unchanged.
Why: Flash Attention computes the same O(N²d) multiply-adds. It cannot skip any — every Q·K pair is needed for the exact softmax. The speedup is entirely from eliminating HBM traffic (IO), not from reducing arithmetic.
Trap
Flash Attention makes attention faster by reducing the number of floating-point multiply-adds.
Claim: standard attention does O(N²d) FLOPs; Flash Attention does fewer because it skips parts of S
Why: This is wrong. Flash Attention computes the same O(N²d) multiply-adds. It cannot skip any — every Q·K pair is needed for the exact softmax. The speedup is entirely from eliminating HBM traffic (IO), not from reducing arithmetic.
Flash Attention reduces HBM accesses from O(N²) to O(N) — FLOPS are unchanged.
Both standard and Flash Attention do O(N²d) multiply-adds; Flash is IO-optimal, not compute-optimal
Why: The N×N score matrix still exists conceptually — it is just never written to HBM. Tiles are computed on-chip and consumed immediately. Speedup comes from GPU bandwidth, not FLOP count.
Section
Part 2 of 4
Concept
Backpropagation requires all intermediate activations to be available. For a transformer with L layers and sequence length N, storing every activation during the forward pass consumes O(L·N·d) memory — the dominant cost at scale.
\[ \text{activation memory} = O(L \cdot N \cdot d) \quad \text{(naive)} \]
Gradient checkpointing keeps only a subset of activations (checkpoints). During backward, it re-runs the forward for each segment to recompute the dropped activations. The trade-off: ≈33% extra FLOPs to save O(√L) or O(1) memory per layer.
Missing information
Discussion prompt
Wrap a transformer layer's forward with checkpoint(). Verify the gradients are numerically identical to the non-checkpointed version.
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
PyTorch discards intermediate activations after the checkpointed block's forward pass and replays it during backward. The computation graph is rebuilt on-the-fly, so gradients match exactly.
Worked example
Wrap a transformer layer's forward with checkpoint(). Verify the gradients are numerically identical to the non-checkpointed version.
import torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
class ToyTransformerLayer(nn.Module):
def __init__(self, d=64):
super().__init__()
self.Wq = nn.Linear(d, d)
self.Wk = nn.Linear(d, d)
self.Wv = nn.Linear(d, d)
self.ff = nn.Linear(d, d)
def forward(self, x):
q=self.Wq(x); k=self.Wk(x); v=self.Wv(x)
scale = x.shape[-1] ** -0.5
a = F.softmax(q @ k.transpose(-2,-1) * scale, dim=-1) @ v
return F.relu(self.ff(a + x))
torch.manual_seed(0)
N_tok, d_model = 16, 64
layer = ToyTransformerLayer(d_model)
# Without checkpointing
x1 = torch.randn(1, N_tok, d_model, requires_grad=True)
out1 = layer(x1); out1.sum().backward()
print('Normal: output shape', out1.shape)
# With gradient checkpointing
x2 = x1.detach().clone().requires_grad_(True)
out2 = checkpoint(layer, x2, use_reentrant=False)
out2.sum().backward()
print('Ckpt: output shape', out2.shape)
print('Grad match:', torch.allclose(x1.grad, x2.grad, atol=1e-5))Grad match = True; checkpointing recomputes the forward but produces identical gradients
Why: PyTorch discards intermediate activations after the checkpointed block's forward pass and replays it during backward. The computation graph is rebuilt on-the-fly, so gradients match exactly.
| strategy | memory (activations) | extra FLOPs | typical use |
|---|---|---|---|
| no checkpointing | O(L·N·d) | 0% | small models / short sequences |
| per-layer checkpointing | O(N·d) per layer saved | ~33% extra | LLM pretraining (e.g. GPT-3) |
| Flash Attention backward | O(N) (tiles recomputed) | ~33% extra | attention op specifically |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Grad match = True; checkpointing recomputes the forward but produces identical gradients
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Wrap a transformer layer's forward with checkpoint(). Verify the gradients are numerically identical to the non-checkpointed version.
Section
Part 3 of 4
Concept
Learned absolute position embeddings (BERT) and sinusoidal encodings (Lesson 82: standard Transformer) both degrade at inference sequences longer than training length — the model has never seen those positions.
ALiBi (Attention with Linear Biases) sidesteps this by adding no position embedding to the token vectors at all. Instead it modifies the attention logits directly with a bias that penalizes distant tokens.
\[ \text{logit}(i,j) \;=\; \frac{q_i \cdot k_j}{\sqrt{d}} \;-\; m_h \cdot |i - j| \]
Estimation
Predict first
For H heads, slopes are a geometric sequence m_h = 2^(−h·8/H). Each head uses a different slope — sharper (larger m_h) heads focus locally; flatter heads attend globally.
Commit before you compute: what does ALiBi slopes and bias matrix come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: slopes = [0.25, 0.0625, 0.0156, 0.0039]; head0 row 0 = [0, -0.25, -0.5, -0.75, -1.0, -1.25]
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Slope 0.25 means each token of distance costs 0.25 logit units.
Worked example
For H heads, slopes are a geometric sequence m_h = 2^(−h·8/H). Each head uses a different slope — sharper (larger m_h) heads focus locally; flatter heads attend globally.
import torch
N, H = 6, 4
slopes = torch.tensor([2**(-h*8/H) for h in range(1, H+1)])
print('slopes:', slopes.numpy().round(4))
i_idx = torch.arange(N).unsqueeze(1) # (N,1)
j_idx = torch.arange(N).unsqueeze(0) # (1,N)
dist = (i_idx - j_idx).abs().float() # (N,N)
alibi_h0 = -slopes[0] * dist # (N,N) for head 0
print('head0 bias row 0:', alibi_h0[0].numpy().round(3))
print('head0 bias row 2:', alibi_h0[2].numpy().round(3))slopes = [0.25, 0.0625, 0.0156, 0.0039]; head0 row 0 = [0, -0.25, -0.5, -0.75, -1.0, -1.25]
Why: Slope 0.25 means each token of distance costs 0.25 logit units. Head 1 (slope 0.0625) decays 4× slower — it sees farther context. Head 4 (slope 0.0039) is nearly flat — global attention.
| head | slope m_h | bias at dist=1 | bias at dist=10 | effective reach |
|---|---|---|---|---|
| 1 | 0.2500 | -0.250 | -2.500 | short-range |
| 2 | 0.0625 | -0.063 | -0.625 | medium |
| 3 | 0.0156 | -0.016 | -0.156 | long-range |
| 4 | 0.0039 | -0.004 | -0.039 | near-global |
Pattern
Step through it
Step through ALiBi slopes and bias matrix one row at a time. What is driving the change, and what would the row after the last one be?
Concept
ALiBi adds no new trainable parameters and imposes no upper bound on sequence length. At inference with sequence length > training length, the bias formula −m·|i−j| naturally penalizes the new long-distance pairs — behavior that generalizes from the training distribution.
Empirically (Press et al., 2022): a GPT-style model trained on length 1024 with ALiBi generalizes to length 2048 with lower perplexity than sinusoidal or learned encodings trained directly at 2048.
ALiBi is used in production models including MPT and BLOOM. It is a drop-in replacement for any position encoding — just add the bias matrix to the attention logits before softmax (Lesson 82 callback: same QK^T, different additive term).
Sorting
Sort into buckets
These are the pieces of Lesson 91: Flash Attention & Memory-Efficient Transformers, out of order. Put each one back under the part of the lesson it belongs to.
Anomaly
Predict first
A student writes this, and it looks reasonable:
The ALiBi paper says 'add the bias to attention weights,' so add it after softmax(QK^T/√d) to get the final probabilities.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Adding after softmax produces values outside [0,1] — the result is no longer a probability distribution.
Add ALiBi before softmax — as an additive logit bias.
Why: Adding after softmax produces values outside [0,1] — the result is no longer a probability distribution. Softmax normalization would have to be redone.
Trap
The ALiBi paper says 'add the bias to attention weights,' so add it after softmax(QK^T/√d) to get the final probabilities.
P = softmax(QK^T * scale) + alibi_bias
Why: Wrong. Adding after softmax produces values outside [0,1] — the result is no longer a probability distribution. Softmax normalization would have to be redone.
Add ALiBi before softmax — as an additive logit bias.
P = softmax(QK^T * scale + alibi_bias, dim=-1)
Why: The bias shifts logits before normalization. Negative bias for distant positions reduces their probability after softmax — the further the token, the smaller its weight. The result is still a valid probability distribution.
Break the constraint
Discussion prompt
The rule this trap just fixed:
Add ALiBi before softmax — as an additive logit bias.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
Adding after softmax produces values outside [0,1] — the result is no longer a probability distribution. Softmax normalization would have to be redone.
Section
Part 4 of 4
Concept
During autoregressive decoding, every token generation step re-reads the full KV cache (keys and values for all past tokens). With H heads, the cache holds H×N×d_k×2 values per layer — it dominates memory bandwidth at long sequences.
\[ \text{KV cache (MHA)} = 2 \cdot H \cdot N \cdot d_k \cdot \text{bytes} \quad \text{per layer} \]
| attention type | KV heads | KV cache at N=2048, H=8, d_k=8, float32 | reduction |
|---|---|---|---|
| MHA | H=8 per layer | 1024.0 KB / layer | 1× |
| MQA | 1 per layer | 128.0 KB / layer | 8× |
| GQA (G=2) | 2 per layer | 256.0 KB / layer | 4× |
Concept
Multi-Query Attention (Shazeer, 2019) keeps H separate Q projections but collapses K and V to a single head each. All Q heads attend to the same K, V. Memory usage for KV cache drops by H× with minimal quality loss on most tasks.
\[ Q \in \mathbb{R}^{H \times N \times d_k}, \quad K,V \in \mathbb{R}^{1 \times N \times d_k} \;\;(\text{shared across heads}) \]
Grouped-Query Attention (GQA) is a middle ground: G groups of query heads share one K, V. MQA = GQA with G=1; MHA = GQA with G=H. Used in LLaMA-2 (G=8) and Mistral.
Estimation
Predict first
Build MQA with H=4 query heads, d_model=16, d_k=4. One Wk and one Wv projection (not H of each). Verify shapes.
Commit before you compute: what does MQA forward pass in PyTorch come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Q:(1,4,6,4) K:(1,1,6,4) V:(1,1,6,4) — K,V broadcast to all 4 Q heads via PyTorch broadcasting
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The (1,1,N,dk) K and V tensors are broadcast along the H dimension when matmul'd with the (1,H,N,dk) Q.
Worked example
Build MQA with H=4 query heads, d_model=16, d_k=4. One Wk and one Wv projection (not H of each). Verify shapes.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(3)
H, d, dk, N = 4, 16, 4, 6
x = torch.randn(1, N, d)
Wq = nn.Linear(d, H*dk, bias=False) # H query heads
Wk = nn.Linear(d, dk, bias=False) # single K head
Wv = nn.Linear(d, dk, bias=False) # single V head
Wo = nn.Linear(H*dk, d, bias=False)
Q = Wq(x).view(1, N, H, dk).transpose(1,2) # (1,H,N,dk)
K = Wk(x).view(1, N, 1, dk).transpose(1,2) # (1,1,N,dk) shared
V = Wv(x).view(1, N, 1, dk).transpose(1,2) # (1,1,N,dk) shared
scores = (Q @ K.transpose(-2,-1)) * (dk**-0.5) # (1,H,N,N) broadcast
attn = F.softmax(scores, dim=-1) @ V # (1,H,N,dk)
out = Wo(attn.transpose(1,2).contiguous().view(1, N, -1))
print(f'Q:{Q.shape} K:{K.shape} V:{V.shape}')
print(f'scores:{scores.shape} attn:{attn.shape} out:{out.shape}')Q:(1,4,6,4) K:(1,1,6,4) V:(1,1,6,4) — K,V broadcast to all 4 Q heads via PyTorch broadcasting
Why: The (1,1,N,dk) K and V tensors are broadcast along the H dimension when matmul'd with the (1,H,N,dk) Q. No explicit tiling needed — PyTorch handles it.
| projection | MHA shape | MQA shape | param reduction |
|---|---|---|---|
| Wq | (d, H·dk) = (16,16) | (d, H·dk) = (16,16) | none |
| Wk | (d, H·dk) = (16,16) | (d, dk) = (16, 4) | H× fewer |
| Wv | (d, H·dk) = (16,16) | (d, dk) = (16, 4) | H× fewer |
| KV cache (N=2048) | H·N·dk×2·bytes | N·dk×2·bytes | H× = 4× here |
Comparison
Comparison matrix
From MQA forward pass in PyTorch: refill the MHA shape column from what you know. The rest of the table is as it appeared.
| projection | MHA shape | MQA shape | param reduction |
|---|---|---|---|
| Wq | (d, H·dk) = (16,16) | (d, H·dk) = (16,16) | none |
| Wk | (d, H·dk) = (16,16) | (d, dk) = (16, 4) | H× fewer |
| Wv | (d, H·dk) = (16,16) | (d, dk) = (16, 4) | H× fewer |
| KV cache (N=2048) | H·N·dk×2·bytes | N·dk×2·bytes | H× = 4× here |
Constraint
Discussion prompt
Run Memory-efficient attention recipe with this step confiscated:
Gradient checkpointing: checkpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per checkpointed block
Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.
Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.
Answer:
m (max), l (normalizer), O (partial output); rescale each tile with exp(m_old − m_new) before adding new tilecheckpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per…−m_h·|i−j| to logits before softmax; slopes 2^(−h·8/H) differ per head; no trainable params, extrapolates past training lengthPattern
m (max), l (normalizer), O (partial output); rescale each tile with exp(m_old − m_new) before adding new tilecheckpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per checkpointed block−m_h·|i−j| to logits before softmax; slopes 2^(−h·8/H) differ per head; no trainable params, extrapolates past training lengthEdge cases
Discussion prompt
Memory-efficient attention recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
m (max), l (normalizer), O (partial output); rescale each tile with exp(m_old − m_new) before adding new tilecheckpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per…−m_h·|i−j| to logits before softmax; slopes 2^(−h·8/H) differ per head; no trainable params, extrapolates past training lengthElimination
Eliminate the wrong options
For sequence length N=2048, d=64, approximately how much more HBM does standard attention access than Flash Attention?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: Standard: 3·N·d + 4·N² = 3·2048·64 + 4·2048² = 393,216 + 16,777,216 = 17,170,432. Flash: ≈4·N·d = 4·2048·64 = 524,288. Ratio = 32.8×. The N² term dominates at long sequences.
Check
Work it out before clicking.
Check your understanding
For sequence length N=2048, d=64, approximately how much more HBM does standard attention access than Flash Attention?
Answer: A
Why: Standard: 3·N·d + 4·N² = 3·2048·64 + 4·2048² = 393,216 + 16,777,216 = 17,170,432. Flash: ≈4·N·d = 4·2048·64 = 524,288. Ratio = 32.8×. The N² term dominates at long sequences.
Prediction
Predict first
In ALiBi with 4 heads and slopes [0.25, 0.0625, 0.0156, 0.0039], what is the bias added to logit(query position 2, key position 5) for head 1 (slope 0.25)?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: −0.75 (= −0.25 × |2−5| = −0.25 × 3)
Why: ALiBi bias = −m_h × |i − j| = −0.25 × |2 − 5| = −0.25 × 3 = −0.75. Verified: alibi_bias_head0[2,5] = −0.7500. Negative because ALiBi always discourages attending to distant tokens.
Check
Apply the formula.
Check your understanding
In ALiBi with 4 heads and slopes [0.25, 0.0625, 0.0156, 0.0039], what is the bias added to logit(query position 2, key position 5) for head 1 (slope 0.25)?
Answer: A
Why: ALiBi bias = −m_h × |i − j| = −0.25 × |2 − 5| = −0.25 × 3 = −0.75. Verified: alibi_bias_head0[2,5] = −0.7500. Negative because ALiBi always discourages attending to distant tokens.
Elimination
Eliminate the wrong options
A transformer uses MQA (1 K,V head) vs MHA (8 heads), d_k=8, N=2048, float32. By what factor does MQA reduce the KV cache per layer?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: MHA KV cache = 2·H·N·d_k·4 = 2·8·2048·8·4 = 1,048,576 bytes = 1024 KB. MQA = 2·1·2048·8·4 = 131,072 bytes = 128 KB. Ratio = 1024/128 = 8×. Verified in the parameter table.
Check
Use the formula.
Check your understanding
A transformer uses MQA (1 K,V head) vs MHA (8 heads), d_k=8, N=2048, float32. By what factor does MQA reduce the KV cache per layer?
Answer: A
Why: MHA KV cache = 2·H·N·d_k·4 = 2·8·2048·8·4 = 1,048,576 bytes = 1024 KB. MQA = 2·1·2048·8·4 = 131,072 bytes = 128 KB. Ratio = 1024/128 = 8×. Verified in the parameter table.
Section
Project
Concept
Three milestones: tiled Flash Attention matching F.scaled_dot_product_attention, a checkpointed transformer layer, and a complete MQA forward pass with KV cache simulation.
| # | milestone | key tool |
|---|---|---|
| 1 | Tiled Flash Attention: match standard output to 1e-6 | online softmax, torch tensors |
| 2 | Checkpointed transformer layer: verify grad match | torch.utils.checkpoint |
| 3 | MQA forward pass: check K,V broadcasting and KV-cache size | nn.Linear, view, transpose |
Build rules: type every line, use fixed seeds (torch.manual_seed(42)), and print shapes at every intermediate step.
Counterexample
Discussion prompt
Three milestones: tiled Flash Attention matching F.scaled_dot_product_attention, a checkpointed transformer layer, and a complete MQA forward pass with KV cache simulation.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
Build rules: type every line, use fixed seeds (torch.manual_seed(42)), and print shapes at every intermediate step.
Worked example
Your turn: implement the tiled online-softmax loop. Predict: will the output match F.scaled_dot_product_attention to < 1e-6 max abs error?
Hint: for each (i,j) tile: m_new = max(m_old, tile_max), rescale O and l by exp(m_old − m_new), accumulate exp(Sij − m_new) @ Vj.
import torch, torch.nn.functional as F
torch.manual_seed(42)
N, d, BLOCK = 8, 4, 4
Q = torch.randn(N, d); K = torch.randn(N, d); V = torch.randn(N, d)
scale = d**-0.5
O_ref = F.scaled_dot_product_attention(
Q.unsqueeze(0).unsqueeze(0),
K.unsqueeze(0).unsqueeze(0),
V.unsqueeze(0).unsqueeze(0)).squeeze()
O = torch.zeros_like(Q)
l = torch.zeros(N)
m = torch.full((N,), float('-inf'))
for j in range(0, N, BLOCK):
Kj, Vj = K[j:j+BLOCK], V[j:j+BLOCK]
for i in range(0, N, BLOCK):
Sij = (Q[i:i+BLOCK] @ Kj.T) * scale
mij = Sij.max(dim=-1).values
m_new = torch.maximum(m[i:i+BLOCK], mij)
e = torch.exp(Sij - m_new.unsqueeze(1))
l_new = torch.exp(m[i:i+BLOCK]-m_new)*l[i:i+BLOCK]+e.sum(dim=-1)
O[i:i+BLOCK] = torch.diag(torch.exp(m[i:i+BLOCK]-m_new))@O[i:i+BLOCK]+e@Vj
m[i:i+BLOCK]=m_new; l[i:i+BLOCK]=l_new
O = O / l.unsqueeze(1)
print('max abs diff:', (O_ref - O).abs().max().item())| method | O[0] | max abs diff vs SDPA |
|---|---|---|
| F.scaled_dot_product_attention | [ 0.2733, -0.0495, -0.3578, 0.3780] | 0 (reference) |
| tiled online-softmax | [ 0.2733, -0.0495, -0.3578, 0.3780] | 5.96e-08 |
| standard manual | [ 0.2733, -0.0495, -0.3578, 0.3780] | 1.19e-07 |
Worked example
Your turn: wrap your transformer layer with checkpoint(). Predict: will the loss and gradients match the non-checkpointed forward?
Hint: from torch.utils.checkpoint import checkpoint; out = checkpoint(layer, x, use_reentrant=False). Same five backward steps as Lesson 40.
import torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
class ToyLayer(nn.Module):
def __init__(self, d=32):
super().__init__()
self.ff = nn.Linear(d, d)
def forward(self, x):
return F.relu(self.ff(x))
torch.manual_seed(0)
d = 32; N = 16
layer = ToyLayer(d)
x1 = torch.randn(1, N, d, requires_grad=True)
out1 = layer(x1); out1.sum().backward()
print('Normal grad norm:', x1.grad.norm().item())
x2 = x1.detach().clone().requires_grad_(True)
out2 = checkpoint(layer, x2, use_reentrant=False)
out2.sum().backward()
print('Ckpt grad norm:', x2.grad.norm().item())| approach | grad norm | activation memory saved |
|---|---|---|
| normal forward | 1.9487 (example) | 0% |
| checkpoint forward | 1.9487 (identical) | all intermediate activations freed after forward |
| trade-off | same gradients | ~33% extra FLOPs from recompute |
Worked example
Your turn: build the MQA forward. Before running, predict the shape of K after view and transpose, and predict the KV cache size at N=2048.
Hint: K = Wk(x).view(1, N, 1, dk).transpose(1,2) — shape (1, 1, N, dk). It broadcasts against (1, H, N, dk) Q during matmul.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(3)
H, d, dk, N = 4, 16, 4, 6
x = torch.randn(1, N, d)
Wq=nn.Linear(d,H*dk,bias=False); Wk=nn.Linear(d,dk,bias=False)
Wv=nn.Linear(d,dk,bias=False); Wo=nn.Linear(H*dk,d,bias=False)
Q=Wq(x).view(1,N,H,dk).transpose(1,2)
K=Wk(x).view(1,N,1,dk).transpose(1,2)
V=Wv(x).view(1,N,1,dk).transpose(1,2)
scores=(Q@K.transpose(-2,-1))*(dk**-0.5)
attn=F.softmax(scores,dim=-1)@V
out=Wo(attn.transpose(1,2).contiguous().view(1,N,-1))
print(f'Q:{Q.shape} K:{K.shape} V:{V.shape}')
kv_mha = 2*H*2048*dk*4; kv_mqa = 2*1*2048*dk*4
print(f'KV MHA N=2048: {kv_mha/1024:.1f} KB MQA: {kv_mqa/1024:.1f} KB ratio: {kv_mha//kv_mqa}x')| tensor | shape | HBM at N=2048 (float32) |
|---|---|---|
| MHA K+V (H=4) | (1, 4, 2048, 4) × 2 | 512.0 KB |
| MQA K+V (1 head) | (1, 1, 2048, 4) × 2 | 128.0 KB |
| reduction | — | 4× (= H) |
Trade off
Comparison matrix
From Milestone 3 — MQA with KV cache sizing: every row here is a choice with a cost. Fill the HBM at N=2048 (float32) column, then say which row you would actually pick and what you give up for it.
| tensor | shape | HBM at N=2048 (float32) |
|---|---|---|
| MHA K+V (H=4) | (1, 4, 2048, 4) × 2 | 512.0 KB |
| MQA K+V (1 head) | (1, 1, 2048, 4) × 2 | 128.0 KB |
| reduction | — | 4× (= H) |
Concept
import torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
# --- Flash Attention (tiled) ---
torch.manual_seed(42)
N, d, BLOCK = 8, 4, 4
Q=torch.randn(N,d); K=torch.randn(N,d); V=torch.randn(N,d); scale=d**-0.5
O=torch.zeros_like(Q); l=torch.zeros(N); m=torch.full((N,),float('-inf'))
for j in range(0,N,BLOCK):
Kj,Vj=K[j:j+BLOCK],V[j:j+BLOCK]
for i in range(0,N,BLOCK):
Sij=(Q[i:i+BLOCK]@Kj.T)*scale; mij=Sij.max(dim=-1).values
m_new=torch.maximum(m[i:i+BLOCK],mij)
e=torch.exp(Sij-m_new.unsqueeze(1))
l_new=torch.exp(m[i:i+BLOCK]-m_new)*l[i:i+BLOCK]+e.sum(dim=-1)
O[i:i+BLOCK]=torch.diag(torch.exp(m[i:i+BLOCK]-m_new))@O[i:i+BLOCK]+e@Vj
m[i:i+BLOCK]=m_new; l[i:i+BLOCK]=l_new
O=O/l.unsqueeze(1)
ref=F.scaled_dot_product_attention(Q.view(1,1,N,d),K.view(1,1,N,d),V.view(1,1,N,d)).squeeze()
print('flash diff:', (O-ref).abs().max().item())
# --- MQA ---
torch.manual_seed(3)
H,d,dk,Ns=4,16,4,6; x=torch.randn(1,Ns,d)
Wq=nn.Linear(d,H*dk,bias=False); Wk=nn.Linear(d,dk,bias=False)
Wv=nn.Linear(d,dk,bias=False); Wo=nn.Linear(H*dk,d,bias=False)
Q_=Wq(x).view(1,Ns,H,dk).transpose(1,2)
K_=Wk(x).view(1,Ns,1,dk).transpose(1,2)
V_=Wv(x).view(1,Ns,1,dk).transpose(1,2)
out_=Wo((F.softmax((Q_@K_.transpose(-2,-1))*(dk**-0.5),dim=-1)@V_).transpose(1,2).contiguous().view(1,Ns,-1))
print('MQA out:', out_.shape)| component | result | notes |
|---|---|---|
| Flash tiled diff | 5.96e-08 | numerically identical to SDPA |
| ALiBi slope h1 | 0.2500 | 2^(−1·8/4); decays 0.25 per token |
| MQA KV cache (H=8, N=2048, d_k=8) | 128 KB vs 1024 KB MHA | 8× reduction verified |
| Checkpoint grad match | True (atol=1e-5) | recompute = correct backward |
All four techniques target the same bottleneck: memory bandwidth. Sequence lengths double when you halve the HBM footprint — that is why GPT-4, LLaMA-2, and Gemini all use Flash Attention + MQA/GQA.
Comparison
Comparison matrix
From The full program: refill the notes column from what you know. The rest of the table is as it appeared.
| component | result | notes |
|---|---|---|
| Flash tiled diff | 5.96e-08 | numerically identical to SDPA |
| ALiBi slope h1 | 0.2500 | 2^(−1·8/4); decays 0.25 per token |
| MQA KV cache (H=8, N=2048, d_k=8) | 128 KB vs 1024 KB MHA | 8× reduction verified |
| Checkpoint grad match | True (atol=1e-5) | recompute = correct backward |
Concept
Out loud, slides closed: (1) explain why standard attention is IO-bound and how Flash Attention's tile loop eliminates HBM round-trips; (2) state the three variables maintained in the online softmax and their update rule; (3) contrast ALiBi with sinusoidal position encoding on the extrapolation question; (4) compute the KV-cache size for MHA vs MQA for your choice of H, d_k, N.
Stretch (homework): implement F.scaled_dot_product_attention profiling against your manual standard attention for N in [512, 1024, 2048] and measure wall-clock time; add ALiBi bias to the toy attention block; implement Grouped-Query Attention (GQA) with G=2 heads per group and verify it lies between MHA and MQA in KV cache size. Next up: Lesson 92 — sparse attention, linear attention, and long-context efficiency.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Standard attention is IO-bound · Gradient checkpointing in a transformer · ALiBi position bias · Multi-Query Attention (MQA) · Your turn: implement Flash + MQA. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
m, sum l, partial output O; rescale with exp(m_old−m_new) per tiletorch.utils.checkpoint to a transformer layer — 33% FLOP overhead saves full activation memory−m_h·|i−j| before softmax; slopes per head; generalizes to unseen sequence lengths| technique | what it optimizes | the one number |
|---|---|---|
| Flash Attention | HBM reads/writes (IO) | O(N) vs O(N²): 32.8× at N=2048, d=64 |
| Gradient checkpointing | activation memory | O(1) per block; +33% FLOPs |
| ALiBi | position encoding + extrapolation | bias = −m_h·|i−j|; no trainable params |
| MQA | KV cache bandwidth | 8× cache reduction for H=8 |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.