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 tile and a 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 , then we do row by row, and finally the output .
The thing is, in standard attention they had to move the entire and to SRAM from HBM, and also send back to HBM from SRAM; and then send the entire to SRAM from HBM again, and send back to HBM from SRAM; and also send the entire and to SRAM from HBM, and also send back 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, and are sent to SRAM from HBM using tiling, and when doing softmax, is sent to SRAM row by row, and is also sent to SRAM by tiling.)

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 and to SRAM from HBM, does the work on-chip, and only sends the final back, so the entire and 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 row sitting there at once. (Here, all three, , , , are sent to SRAM from HBM using tiling.)

Online Softmax

How the Read/Write Complexity Actually Got Improved
Normal Attention
Input is , so are each , and and are .
| Step | Read/Write complexity | Comment |
|---|---|---|
| move (HBM → SRAM) | ||
| move (HBM → SRAM) | and do matrix multiplication | |
| move (SRAM → HBM) | ||
| move (HBM → SRAM) | and calculate softmax | |
| move (SRAM → HBM) | ||
| move (HBM → SRAM) | ||
| move (HBM → SRAM) | and do matrix multiplication | |
| move (SRAM → HBM) |
Total: , where the terms from writing/reading the big matrices dominate.
Flash Attention
| Step | Read/Write complexity | Comment |
|---|---|---|
| move (HBM → SRAM) | ||
| move (HBM → SRAM) | do on-chip, and online softmax | |
| move (HBM → SRAM) | multiply with | |
| move (SRAM → HBM) |
Total: , since the matrices never touch HBM, so the terms are gone.