All concepts

Flash Attention

Compute exact attention while reducing expensive GPU memory traffic.

Transformers & LLMs · Advanced · ~8 min

In plain English

Same maths as attention, better housekeeping: never write the giant score matrix to slow memory, compute it in tiles inside the fast memory instead.

Why it's worth your time

It made long contexts affordable without changing a single result — the rare optimization with no accuracy trade-off.

If you remember three things

  • Exact, not approximate — identical output to standard attention
  • It's an IO optimization, not a maths one
  • Memory drops from quadratic to linear in sequence length

Overview

Exact attention computed without ever materializing the n×n score matrix in slow HBM. It tiles Q, K, and V into blocks that fit in on-chip SRAM and uses an online softmax to accumulate results, so memory scales O(n) instead of O(n²). This unlocks longer context and larger batches.

How it works

  1. Start: Q K V Blocks Attention normally materializes a huge n-by-n score matrix.
  2. Q K V Blocks -> Tiled SRAM FlashAttention processes blocks that fit in fast on-chip memory.
  3. Tiled SRAM -> Online Softmax It computes softmax statistics incrementally without storing the whole matrix.
  4. Online Softmax -> Exact Output The result is mathematically exact attention, but faster and more memory efficient.
  5. Exact Output -> Long Context Lower memory pressure makes longer sequences and larger batches practical.

In an interview

FlashAttention is an IO-aware exact attention algorithm. Instead of writing the full n×n scores to GPU HBM, it streams Q, K, and V in tiles through fast SRAM and keeps running softmax statistics (max and sum). It's bound by memory bandwidth, not capacity, giving the identical output with far less traffic.

Production defaults

Turn it on
always, if your hardware and library support it. There is no downside
Expect
2–4× speedup on long sequences and a large memory reduction
Check
that your attention variant (ALiBi, sliding window) is supported by the kernel you're using

What breaks

  • Results differ slightly from the reference — Floating-point accumulation order, not a bug. Differences should be at the noise floor — if they're larger, check the mask handling.
  • No speedup — Short sequences are compute-bound, not memory-bound. The win scales with sequence length.

Watch it explained

Flash Attention: The Fastest Attention Mechanism? — Tales Of Tensors, 8:43

Related