Skip to content

NTU Hung-yi Lee ML 2026 Guide: Faster Generation, Part 1: Flash Attention and Why Moving Data Is the Bottleneck

Sep 30, 20261 min
TL;DRIn week 3 of ML 2026, Hung-yi Lee spends the first half of the inference lecture on one technique: Flash Attention. A GPU's execution units are fast, but their workbench (on-chip SRAM) is tiny, so data has to be carried to and from the warehouse (HBM). The carrying is the bottleneck. A naive softmax makes several round trips to the warehouse. Flash Attention assumes the current maximum is Amax, then multiplies by a correction factor when a larger value shows up. That lets it find the maximum, build the denominator, and compute the weighted sum in one pass, without ever materializing the attention weights. The output is identical to standard attention, no retraining is needed, and the cost is a little extra compute and a little brain strain.

🌏 中文版

This guide is based on the 3/20 materials of NTU Hung-yi Lee's Machine Learning 2026 Spring (taught in Mandarin). It is part 6 of the Reading NTU Hung-yi Lee Machine Learning 2026 Spring series. The previous post is HW2: An AI Agent as an AI Engineer. The earlier posts were about agents: OpenClaw prepends a long system prompt to every message you send, and Context Engineering deals with context that no longer fits. This post moves inside the model: when inputs run to tens or hundreds of thousands of tokens, why does generation slow down, and how do you speed it up?

Official materials used: the slides inference.pdf, pages 1–28 (55 pages total; the second half is the next post on KV Cache), the lecture video 加快語言模型生成速度 (1/2):Flash Attention (in Mandarin), and the demo Colab linked on slide 28. Access level is A3: slides (pdf/pptx), recording, and demo code are all public. There is no quiz for this lecture; the matching exercises are in HW3.

Prerequisite: you are expected to know Transformers

Slide 2 has one prerequisite link: Lecture 3 of Intro to Generative AI & ML 2025: Dissecting LLMs (in Mandarin). Lee opens by saying the lecture assumes you already understand how a Transformer works inside, and that it is about inference, not training.

Slides 3–4 review self-attention in two figures. Inputs x1…x5 are each multiplied by three matrices to get q, k, v. The output at position 4 comes from dotting q4 with k1…k4 to get a1…a4, applying softmax to get â1…â4, and taking the weighted sum of v1…v4. The rest of the lecture is about reordering this computation.

Slide 5 splits generation into two phases: Prefill, which ingests the whole prompt at once, and Decode, which emits tokens one at a time. Both terms come back in the KV Cache post.

For any speed-up, ask what it costs

Slide 6 lists three classic techniques: Flash Attention, KV Cache, and Speculative Decoding. Lee sets up a framework that runs through both lectures: if someone tells you they invented a speed-up, ask what they paid for it. The usual costs are:

  1. It changes the attention computation, so the result is an approximation.
  2. It is tied to a model: you have to train or customize a specific model, so it is not plug-and-play.
  3. Even without those two, something else was traded away.

At the end of the second lecture every method goes into one table (slide 55, covered in full in the next post).

Speculative Decoding is not covered this semester. Slide 7 only links the old video, Intro to Generative AI 2024, Lecture 16: Speculative Decoding (in Mandarin). Lee says topics already covered in past years are not repeated, but the homework includes it. One-line version: a small model guesses a few tokens ahead, and the large model verifies them in one parallel pass. Details are in the HW3 guide.

The scene: compute is fast, moving data is slow

Flash Attention comes from the 2022 paper FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (slide 8). Lee first explains why it is impressive. Its output is exactly the same as standard attention, not an approximation. It drops into any Transformer that uses self-attention, with no model lock-in. And its cost is tiny.

The core idea is to respect how a GPU actually computes. Slides 9–10 use an analogy (Lee stresses it is simplified and not GPU-specific):

  • Execution units are a crowd of many-armed sprites: lots of them, and fast.
  • The workbench is on-chip SRAM. It is small and holds only a few values at a time.
  • The warehouse is HBM. It is far bigger than the workbench, but not infinite.

Data has to be carried from the warehouse to the workbench to be processed, then carried back. Once something is on the workbench, you can treat the computation as instant. The carrying is what slows things down. Flash Attention reorders the computation to cut the number of trips, without changing the result.

The demo Colab ran on an A100 80GB. Lee points out that 80GB is the warehouse. The workbench stays small, usually a dozen or so MB.

Intuition: the naive way makes many warehouse trips

Slides 11–17 walk through standard attention. To keep it simple, Lee uses a single query (a real GPU handles many at once).

The key constraint: nothing that scales with sequence length L can sit on the workbench. An agent's input might be ten thousand, a hundred thousand, or a million tokens. Even one number per position is a million numbers, which will not fit. So the keys are split into chunks of N keys each, B = L/N chunks in total.

Under that constraint, turning a_i into â_i takes several passes:

  1. For each chunk, compute the dot products a_i of q and k and write them back to the warehouse.
  2. Read the a_i chunk by chunk, keeping a running maximum. Only after the last chunk do you know Amax. In practice you subtract Amax before taking the exponential to avoid overflow.
  3. Read the a_i again, compute exp(a_i − Amax), and write it back.
  4. Read again and accumulate the denominator S.
  5. Read again, divide by S to get â_i, and write it back.
  6. Finally read â_i and v chunk by chunk and accumulate the weighted sum to get the output O.

Slide 17 asks: going from a_i to â_i takes several reads. Is that really necessary?

Mechanism: run with the wrong answer, then correct it

Slides 18–24 start with a simplified version. The difficulty is that the denominator depends on Amax, and you only know Amax after seeing everything, which seems to force two passes.

The fix is to assume the maximum of the first chunk, d1, is Amax and compute the partial sum s1. When the second chunk reveals a larger d2, you do not reread the first chunk. You multiply s1 by exp(d1 − d2), and it becomes what you would have gotten using d2 from the start. Lee says the correction is "nothing fancy, just one formula," but all of Flash Attention is this trick applied over and over. After the last chunk, d is Amax and s is the correct denominator.

Expand: why multiplying by exp(d1 − d2) fixes it

The first chunk's partial sum is

s1 = Σ_{i=1..N} exp(a_i − d1)

Multiply by exp(d1 − d2):

s1 · exp(d1 − d2) = Σ_{i=1..N} exp(a_i − d1 + d1 − d2) = Σ_{i=1..N} exp(a_i − d2)

That is exactly what the first chunk should contribute if d2 were Amax, so it can be added directly to the second chunk's Σ exp(a_i − d2). For chunk k in general:

d_k = max(d_{k−1}, max of chunk k) s_k = s_{k−1} · exp(d_{k−1} − d_k) + Σ_{i∈chunk k} exp(a_i − d_k)

Now getting from a_i to â_i takes only two reads: one that finds Amax and the denominator together, and one that produces â_i (slide 24).

That is still not the whole of Flash Attention. Slides 25–27 pose what Lee calls the soul-searching question: do you really need the attention weights before computing the weighted sum?

Flash Attention says no. When the first chunk is read, v is read too, and O1 is computed with the "wrong" weights exp(a_i − d1)/s1. When the second chunk arrives, d and s are updated to the more accurate d2 and s2. O1 is multiplied by (s1/s2)·exp(d1 − d2) to erase the traces of d1 and s1, and then the second chunk's contribution is added. After correcting all the way to the last chunk, O is the right output. q, k, and v all go onto the workbench in one pass, and â_i is never written out.

Expand: the correction for the output O

O_k = O_{k−1} · (s_{k−1} / s_k) · exp(d_{k−1} − d_k) + Σ_{i∈chunk k} [exp(a_i − d_k) / s_k] · v_i

Multiplying by s_{k−1} cancels the old denominator and dividing by s_k applies the new one. Multiplying by exp(d_{k−1} − d_k) fixes the exponent. After chunk B, d_B = Amax and s_B is the full denominator, so

O_B = Σ_{i=1..L} â_i · v_i

which in theory matches computing â_i first and then taking the weighted sum.

Lee mentions one practical consequence. Because the attention weights are never computed, if you use Flash Attention in Hugging Face and try to plot an attention matrix for analysis, you get an error saying there are no attention weights to read.

Back to real models: you are probably using it already

The demo Colab calls PyTorch's scaled_dot_product_attention with either SDPBackend.MATH (the naive algorithm) or SDPBackend.FLASH_ATTENTION. Lee explains that without special settings PyTorch usually defaults to Flash Attention, so your everyday Transformer runs likely use it already.

The Colab does three things:

  • Numerical check: random q, k, v (B=4, H=8, L=256, D=64). The maximum difference between the two algorithms is around 1e-7.
  • Speed comparison: sequence lengths from 64 to 4096. On the A100, Flash Attention is about 9× faster at length 4096 (the saved Colab output shows 9.49x).
  • A real model: google/gemma-3-4b-it with attn_implementation set to eager (no Flash Attention) or sdpa, fed a long string repeated many times and asked for one token. In the lecture, short inputs showed no real difference, because the model also spends a lot of time on feed-forward layers and embeddings. Once the string got longer, Flash Attention was clearly faster. Ten times longer again, and the run hit CUDA out of memory.

Lee notes that what ran out was the warehouse, not the workbench. However big it is, the warehouse has a limit. Why long sequences blow up the warehouse is the topic of the next lecture, KV Cache.

His conclusion: Flash Attention gets a several-fold speed-up just by carrying data less often. The only costs are a more complex algorithm and a bit of extra compute for the corrections, which he considers minor next to the gain.

Going deeper

  • Papers: read Algorithm 1 in FlashAttention alongside the correction formulas above. HW3 also assigns FlashAttention-2 and FlashAttention-3; see the HW3 guide.
  • Try it: open the demo Colab, extend SEQ_LENS to 8192, and check whether the speed-up keeps growing. Then shrink B and H and see whether Flash Attention still pays off for short sequences.
  • Related on this site: the Stanford CS336 inference guide covers memory-bound workloads from a roofline angle, and the CME295 LLM systems guide also covers Flash Attention and inference optimization.

What this post could and could not verify

Verified: the structure and text of slides 1–28, the video transcript (from the zh-TW captions on YouTube), the demo Colab's code and saved outputs, and the titles of cited papers (checked on arXiv).

Not verified: the captions name a different demo model; this post follows the Colab code, which uses google/gemma-3-4b-it. The timings in the lecture do not match the saved Colab output (the longest saved run did not hit OOM), so the real-model section reports trends only, not seconds. Slides 11–27 are mostly animated figures, so their content is paraphrased from the transcript.

Series navigation: Series overview | Previous: HW2: An AI Agent as an AI Engineer | Next: Faster Generation, Part 2: KV Cache

References