Lesson 91: Flash Attention & Memory-Efficient Transformers

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

What this lesson covers

The lesson, slide by slide

1. Flash Attention & Memory-Efficient Transformers

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.

2. By the end of this lesson you can

Objectives

  1. State Flash Attention's IO complexity — O(N) HBM reads vs standard O(N²) — and explain why IO is the bottleneck
  2. Trace the tiled online-softmax loop (running max m, sum l, output accumulator O) and reproduce the update equations
  3. Implement gradient checkpointing for a transformer layer and measure the activation-memory trade-off
  4. Compute ALiBi position bias −m·|i−j| for given slopes and explain why it extrapolates beyond training length
  5. Build a Multi-Query Attention (MQA) forward pass and quantify the KV-cache reduction vs standard MHA

3. What survived from KV Cache & Attention Efficiency?

Warm-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).

4. Standard attention is IO-bound

Section

Part 1 of 4

5. Where the time goes: HBM vs SRAM

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 tiercapacitybandwidthrole
SRAM (on-chip)~20 MB~19 TB/sregisters, shared mem — fast but tiny
HBM (off-chip)~80 GB~2 TB/sactivations, 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.

6. Fill in: capacity for Where the time goes: HBM vs SRAM

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 tiercapacitybandwidthrole
SRAM (on-chip)~20 MB~19 TB/sregisters, shared mem — fast but tiny
HBM (off-chip)~80 GB~2 TB/sactivations, weights — slow but large

7. Standard attention IO: O(N²) HBM accesses

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
5121,146,880131,0728.8×
10244,390,912262,14416.8×
204817,170,432524,28832.8×
409667,895,2961,048,57664.8×

8. What each one costs: Standard attention IO: O(N²) HBM accesses

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
5121,146,880131,0728.8×
10244,390,912262,14416.8×
204817,170,432524,28832.8×
409667,895,2961,048,57664.8×

9. Intuition: tiles fit in SRAM

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.

10. Guess the shape of the answer: Online softmax: tiled update equations

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.

11. Online softmax: tiled update equations

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.

variableshapewhat 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

12. Work backwards from the answer: Online softmax: tiled update equations

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.

13. Something is wrong here: Flash Attention reduces FLOPS

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.

14. Trap: Flash Attention reduces FLOPS

Trap

The 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.

The fix

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.

15. Gradient checkpointing in a transformer

Section

Part 2 of 4

16. Activation memory in deep transformers

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.

17. What has to be given first: torch.utils.checkpoint in a transformer 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.

18. torch.utils.checkpoint in a transformer layer

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.

strategymemory (activations)extra FLOPstypical use
no checkpointingO(L·N·d)0%small models / short sequences
per-layer checkpointingO(N·d) per layer saved~33% extraLLM pretraining (e.g. GPT-3)
Flash Attention backwardO(N) (tiles recomputed)~33% extraattention op specifically

19. Work backwards from the answer: torch.utils.checkpoint in a transformer layer

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.

20. ALiBi position bias

Section

Part 3 of 4

21. Rotary and learned position encodings: the extrapolation problem

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| \]

22. Guess the shape of the answer: ALiBi slopes and bias matrix

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.

23. ALiBi slopes and bias matrix

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.

headslope m_hbias at dist=1bias at dist=10effective reach
10.2500-0.250-2.500short-range
20.0625-0.063-0.625medium
30.0156-0.016-0.156long-range
40.0039-0.004-0.039near-global

24. Watch it run: ALiBi slopes and bias matrix

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?

  1. Step 1: head is 1
  2. Step 2: head is 2
  3. Step 3: head is 3
  4. Step 4: head is 4

25. Why ALiBi extrapolates to longer sequences

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).

26. Where does each piece belong: Lesson 91: Flash Attention &…

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.

Standard attention is IO-bound
Where the time goes: HBM vs SRAM; Standard attention IO: O(N²) HBM accesses; Intuition: tiles fit in SRAM
Gradient checkpointing in a transformer
Activation memory in deep transformers; torch.utils.checkpoint in a transformer layer
ALiBi position bias
Rotary and learned position encodings: the extrapolation problem; ALiBi slopes and bias matrix; Why ALiBi extrapolates to longer sequences
s1
Standard attention is IO-bound is where Lesson 91: Flash Attention & Memory-Efficient Transformers puts Where the time goes: HBM vs SRAM, Standard attention IO: O(N²) HBM accesses, Intuition: tiles fit in SRAM. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
Gradient checkpointing in a transformer is where Lesson 91: Flash Attention & Memory-Efficient Transformers puts Activation memory in deep transformers, torch.utils.checkpoint in a transformer layer. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
ALiBi position bias is where Lesson 91: Flash Attention & Memory-Efficient Transformers puts Rotary and learned position encodings: the extrapolation problem, ALiBi slopes and bias matrix, Why ALiBi extrapolates to longer sequences. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

27. Something is wrong here: ALiBi bias added after softmax

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.

28. Trap: ALiBi bias added after softmax

Trap

The 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.

The fix

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.

29. Break it on purpose: ALiBi bias added after softmax

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.

30. Multi-Query Attention (MQA)

Section

Part 4 of 4

31. MHA KV cache is the inference bottleneck

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 typeKV headsKV cache at N=2048, H=8, d_k=8, float32reduction
MHAH=8 per layer1024.0 KB / layer1×
MQA1 per layer128.0 KB / layer8×
GQA (G=2)2 per layer256.0 KB / layer4×

32. MQA: one shared K, V for all query heads

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.

33. Guess the shape of the answer: MQA forward pass in PyTorch

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.

34. MQA forward pass in PyTorch

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.

projectionMHA shapeMQA shapeparam 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·bytesN·dk×2·bytesH× = 4× here

35. Fill in: MHA shape for MQA forward pass in PyTorch

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.

projectionMHA shapeMQA shapeparam 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·bytesN·dk×2·bytesH× = 4× here

36. Without one step: Memory-efficient attention recipe

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:

  1. IO complexity first: standard attention = O(N²) HBM; Flash Attention = O(N) HBM via SRAM tiling — same FLOPs, faster because bandwidth-bound
  2. Tiled softmax: maintain running m (max), l (normalizer), O (partial output); rescale each tile with exp(m_old − m_new) before adding new tile
  3. Gradient checkpointing: checkpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per…
  4. ALiBi: add −m_h·|i−j| to logits before softmax; slopes 2^(−h·8/H) differ per head; no trainable params, extrapolates past training length
  5. MQA / GQA: H Q heads share 1 (or G) K,V heads — reduces KV-cache by H/G×; critical for long-context autoregressive inference (LLaMA-2, Mistral)

37. Memory-efficient attention recipe

Pattern

  1. IO complexity first: standard attention = O(N²) HBM; Flash Attention = O(N) HBM via SRAM tiling — same FLOPs, faster because bandwidth-bound
  2. Tiled softmax: maintain running m (max), l (normalizer), O (partial output); rescale each tile with exp(m_old − m_new) before adding new tile
  3. Gradient checkpointing: checkpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per checkpointed block
  4. ALiBi: add −m_h·|i−j| to logits before softmax; slopes 2^(−h·8/H) differ per head; no trainable params, extrapolates past training length
  5. MQA / GQA: H Q heads share 1 (or G) K,V heads — reduces KV-cache by H/G×; critical for long-context autoregressive inference (LLaMA-2, Mistral)

38. Where does it stop working: Memory-efficient attention recipe

Edge 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:

  1. IO complexity first: standard attention = O(N²) HBM; Flash Attention = O(N) HBM via SRAM tiling — same FLOPs, faster because bandwidth-bound
  2. Tiled softmax: maintain running m (max), l (normalizer), O (partial output); rescale each tile with exp(m_old − m_new) before adding new tile
  3. Gradient checkpointing: checkpoint(fn, x) discards activations, recomputes them during backward — 33% FLOP overhead, O(1) activation memory per…
  4. ALiBi: add −m_h·|i−j| to logits before softmax; slopes 2^(−h·8/H) differ per head; no trainable params, extrapolates past training length
  5. MQA / GQA: H Q heads share 1 (or G) K,V heads — reduces KV-cache by H/G×; critical for long-context autoregressive inference (LLaMA-2, Mistral)

39. Rule out three: Check yourself — Flash Attention IO

Elimination

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.

  • A. ~32.8× more (verified: 17,170,432 vs 524,288 elements)
  • B. ~2× more — the N² matrix is stored once, not twice
  • C. ~N× more — O(N²) vs O(N) means exactly N-fold
  • D. No difference — Flash Attention only reduces backward-pass memory

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.

40. Check yourself — Flash Attention IO

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?

  • A. ~32.8× more (verified: 17,170,432 vs 524,288 elements) (correct)
  • B. ~2× more — the N² matrix is stored once, not twice
  • C. ~N× more — O(N²) vs O(N) means exactly N-fold
  • D. No difference — Flash Attention only reduces backward-pass memory

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.

Why B tempts people
The 4N² comes from S being written once, read for softmax, P written, P read for P·V — four N² passes, not two. And the 3Nd load terms are negligible compared to N².
Why C tempts people
O(N²)/O(N) = O(N) asymptotically, but for N=2048, d=64, the constant factor (4N² vs 4Nd) gives 32.8×, not exactly 2048×. Asymptotic ratio ≈ N/(4d) = 2048/256 = 8 for the dominant terms.
Why D tempts people
Flash Attention reduces both forward and backward HBM. In the forward pass it never writes S to HBM; in the backward it recomputes S tiles rather than loading a stored N×N P matrix.

41. Answer it before you see the options: Check yourself — ALiBi

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.

42. Check yourself — ALiBi

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)?

  • A. −0.75 (= −0.25 × |2−5| = −0.25 × 3) (correct)
  • B. −0.50 (= −0.25 × 2)
  • C. +0.75 (bias favors distant tokens)
  • D. 0 (same position = no bias)

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.

Why B tempts people
−0.50 = −0.25×2 uses |i−j|=2 (the distance from position 2 to 0), not from position 2 to 5 (distance 3). The direction matters — |2−5|=3, not 2.
Why C tempts people
ALiBi bias is always ≤ 0 (distance ≥ 0, slope > 0). The intent is to penalize distant attention, not reward it. A positive bias would encourage attending further — the opposite of the design.
Why D tempts people
Zero bias only applies when i=j (|i−j|=0). Position 2 and position 5 differ by 3 positions, so the penalty is nonzero.

43. Rule out three: Check yourself — MQA KV cache

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.

  • A. 8× (MHA: 1024 KB/layer; MQA: 128 KB/layer)
  • B. 2× (only K or only V is shared, not both)
  • C. 16× (MQA removes both K and V, each was 8 heads)
  • D. No reduction — the cache still holds N past states regardless

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.

44. Check yourself — MQA KV cache

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?

  • A. 8× (MHA: 1024 KB/layer; MQA: 128 KB/layer) (correct)
  • B. 2× (only K or only V is shared, not both)
  • C. 16× (MQA removes both K and V, each was 8 heads)
  • D. No reduction — the cache still holds N past states regardless

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.

Why B tempts people
In MQA both K and V collapse to 1 head each. The reduction is full H× (from 8 to 1 K-head, 8 to 1 V-head), not a partial H/2 × from sharing only one of them.
Why C tempts people
16× would require eliminating the KV cache entirely or reducing to 0.5 heads. MQA still needs 1 K and 1 V head — the cache shrinks by H (heads, from 8 to 1), not by 2H.
Why D tempts people
The cache size is proportional to the number of KV heads. MHA has H=8 K,V tensors per layer; MQA has 1. The N-dependent factor is the same, but the H factor drops from 8 to 1.

45. Your turn: implement Flash + MQA

Section

Project

46. Project: memory-efficient attention from scratch

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.

#milestonekey tool
1Tiled Flash Attention: match standard output to 1e-6online softmax, torch tensors
2Checkpointed transformer layer: verify grad matchtorch.utils.checkpoint
3MQA forward pass: check K,V broadcasting and KV-cache sizenn.Linear, view, transpose

Build rules: type every line, use fixed seeds (torch.manual_seed(42)), and print shapes at every intermediate step.

47. Break it if you can: Project: memory-efficient attention from scratch

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.

48. Milestone 1 — tiled Flash Attention

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())
methodO[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

49. Milestone 2 — gradient checkpointing

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())
approachgrad normactivation memory saved
normal forward1.9487 (example)0%
checkpoint forward1.9487 (identical)all intermediate activations freed after forward
trade-offsame gradients~33% extra FLOPs from recompute

50. Milestone 3 — MQA with KV cache sizing

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')
tensorshapeHBM at N=2048 (float32)
MHA K+V (H=4)(1, 4, 2048, 4) × 2512.0 KB
MQA K+V (1 head)(1, 1, 2048, 4) × 2128.0 KB
reduction—4× (= H)

51. What each one costs: Milestone 3 — MQA with KV cache sizing

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.

tensorshapeHBM at N=2048 (float32)
MHA K+V (H=4)(1, 4, 2048, 4) × 2512.0 KB
MQA K+V (1 head)(1, 1, 2048, 4) × 2128.0 KB
reduction—4× (= H)

52. The full program

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)
componentresultnotes
Flash tiled diff5.96e-08numerically identical to SDPA
ALiBi slope h10.25002^(−1·8/4); decays 0.25 per token
MQA KV cache (H=8, N=2048, d_k=8)128 KB vs 1024 KB MHA8× reduction verified
Checkpoint grad matchTrue (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.

53. Fill in: notes for The full program

Comparison

Comparison matrix

From The full program: refill the notes column from what you know. The rest of the table is as it appeared.

componentresultnotes
Flash tiled diff5.96e-08numerically identical to SDPA
ALiBi slope h10.25002^(−1·8/4); decays 0.25 per token
MQA KV cache (H=8, N=2048, d_k=8)128 KB vs 1024 KB MHA8× reduction verified
Checkpoint grad matchTrue (atol=1e-5)recompute = correct backward

54. Show it off

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.

55. Connect it up: Lesson 91: Flash Attention & Memory-Efficient Transformers

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.

56. What you can do now

Recap

techniquewhat it optimizesthe one number
Flash AttentionHBM reads/writes (IO)O(N) vs O(N²): 32.8× at N=2048, d=64
Gradient checkpointingactivation memoryO(1) per block; +33% FLOPs
ALiBiposition encoding + extrapolationbias = −m_h·|i−j|; no trainable params
MQAKV cache bandwidth8× cache reduction for H=8

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 91 — Flash Attention, Gradient Checkpointing, ALiBi, MQA — Barron · USAAIO Round 2 Preparation, 2026
  2. Flash Attention: Fast and Memory-Efficient Exact Attention with IO-Awareness (Dao et al., 2022)
  3. Tiled softmax, ALiBi slopes, MQA KV-cache counts, gradient checkpointing — verified with torch 2.7.1+cpu and numpy 2.2.6, June 2026 — Real execution, verified

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108