S = QK^T (N × N) 
P = softmax(S) (N × N) 
O = PV (N × d)

In standard attention, you have to initialize an matrix. For N=8K, that’d be 67M entires that you have to write to GPU memory, read back for softmax, and then read again for the PV multiply. Each of those would require a full trip to HBM. So the self-attention operation here is memory-bandwith bound.

FlashAttention solves this by never materializing the full matrix and instead does:

  1. Tiling: Splits Q,K,V into blocks that fit in SRAM
  2. Fusing the matmulsoftmaxmatmul pipeline into one GPU kernel so that the intermediate results stay in SRAM the whole time
  3. Online Softmax: compute softmax incrementally block by clock, without seeing the whole erow at once.

Online Softmax

Softmax equation:

To calculate this you would need the global max of the vector x, and the sum of the exponentials of the vector x. FlashAttention instead processes blocks of K/V one at a time, and calculates the max and sum of the block.

It keeps two running values per query row:

  • : The maximum value of the current block
  • : The sum of the exponentials of the current block

When a new block arrives with a new local max, everything computed so far is rescaled and the new block’s contribution gets folded in.