Seungjun Lee

NLP Bascis

Flash Attention

Overview

HBM is large but slow, and SRAM is small but fast and good for matrix multiplication. So from standard attention, we have been using tiling to calculate attention, like moving a QQ tile and a KK tile to SRAM, calculating the attention answer, and sending it back to HBM.

Here comes the main bottleneck. Keep moving tensors from HBM to SRAM takes a lot of time, and flash attention managed to reduce the read/write operations substantially, making attention calculation faster.

Standard Attention

The way attention gets calculated is basically three steps: first we get the scores S=QK⊤S = QK^\top, then we do P=softmax(S)P = \text{softmax}(S) row by row, and finally the output O=PVO = PV.

The thing is, in standard attention they had to move the entire QQ and KK to SRAM from HBM, and also send QK⊤QK^\top back to HBM from SRAM; and then send the entire QK⊤QK^\top to SRAM from HBM again, and send softmax(QK⊤)\text{softmax}(QK^\top) back to HBM from SRAM; and also send the entire softmax(QK⊤)\text{softmax}(QK^\top) and VV to SRAM from HBM, and also send back softmax(QK⊤) V\text{softmax}(QK^\top)\,V from SRAM to HBM. It's not moving it all at once of course, but tiling, tiling, tiling, and eventually it becomes like this. (Here, QQ and KK are sent to SRAM from HBM using tiling, and when doing softmax, QK⊤QK^\top is sent to SRAM row by row, and VV is also sent to SRAM by tiling.)

screenshot-2026-09-16-at-5-09-15-pm

Flash Attention

Flash attention gets the same answer, but the trick is it doesn't keep sending those big intermediate matrices back to HBM. It sends the entire QQ and KK to SRAM from HBM, does the work on-chip, and only sends the final softmax(QK⊤) V\text{softmax}(QK^\top)\,V back, so the entire softmax(QK⊤)\text{softmax}(QK^\top) and VV stay in SRAM while it computes, instead of going back and forth. Again, it's not all at once, but tiling, tiling, tiling, and eventually it becomes like this.

This became possible thanks to online softmax, since it lets us do the softmax bit by bit as each tile comes in, so we never need the whole QK⊤QK^\top row sitting there at once. (Here, all three, QQ, KK, VV, are sent to SRAM from HBM using tiling.)

screenshot-2026-09-16-at-5-09-45-pm

Online Softmax

image

How the Read/Write Complexity Actually Got Improved

Normal Attention

Input is (N,L,D)(N, L, D), so Q,K,VQ, K, V are each (N,L,D)(N, L, D), and QK⊤QK^\top and softmax(QK⊤)\text{softmax}(QK^\top) are (N,L,L)(N, L, L).

StepRead/Write complexityComment
move QQ (HBM → SRAM)NLDN L D
move KK (HBM → SRAM)NLDN L Dand do matrix multiplication
move QK⊤QK^\top (SRAM → HBM)NL2N L^2
move QK⊤QK^\top (HBM → SRAM)NL2N L^2and calculate softmax
move softmax(QK⊤)\text{softmax}(QK^\top) (SRAM → HBM)NL2N L^2
move softmax(QK⊤)\text{softmax}(QK^\top) (HBM → SRAM)NL2N L^2
move VV (HBM → SRAM)NLDN L Dand do matrix multiplication
move softmax(QK⊤) V\text{softmax}(QK^\top)\,V (SRAM → HBM)NLDN L D

Total: O(NL2+NLD)O(N L^2 + N L D), where the NL2N L^2 terms from writing/reading the big (N,L,L)(N, L, L) matrices dominate.

Flash Attention

StepRead/Write complexityComment
move QQ (HBM → SRAM)NLDN L D
move KK (HBM → SRAM)NLDN L Ddo QK⊤QK^\top on-chip, and online softmax
move VV (HBM → SRAM)NLDN L Dmultiply with VV
move softmax(QK⊤) V\text{softmax}(QK^\top)\,V (SRAM → HBM)NLDN L D

Total: O(NLD)O(N L D), since the (N,L,L)(N, L, L) matrices never touch HBM, so the NL2N L^2 terms are gone.