Trace FlashAttention's evolution through plain PyTorch code.

Explore the updated FlashAttention-PyTorch repository, which now features educational implementations of FA1 through FA4 in plain PyTorch.

3 min readMachine Learning

The FlashAttention series has always carried a reputation for being impenetrable, a stack of CUDA tricks that only the most hardened systems programmers dared to approach. That reputation is undeserved. Shreyansh26's updated repository strips away the hardware-specific noise and lays bare the actual intellectual through-line: the same attention math, re-orchestrated four different ways. That is the story worth telling, and it is one that more people deserve to read.

What makes this approach so effective is its insistence on plain PyTorch as the medium. By avoiding the Hopper and Blackwell internals, the repo forces the algorithmic ideas to stand on their own. FA1 gives you the tiled online softmax baseline, the foundational trick that makes flash attention possible at all. FA2 then introduces split-Q and deferred normalization, a subtle shift in who owns which query tile and when the rescaling happens. FA3 moves into an explicit staged pipeline with ping-pong buffers, and FA4 adds a scheduler with distinct main, softmax, and correction phases, plus conditional rescaling. None of these are new mathematical insights on their own. The insight is that orchestration itself is the innovation, and that is a lesson that gets lost when you are buried in kernel directives.

For the working engineer or the curious student, this is the missing bridge. You can read the FlashAttention papers and nod along to the high-level summaries, but the gap between "we fuse the softmax" and a working implementation is enormous. This repo closes that gap by showing you the code, line by line, version by version. It does not pretend to be production-ready, and it does not need to be. Its purpose is pedagogical, and it succeeds because it isolates the variable that actually changed: not the math, but the scheduling. That is a rare clarity, and it is exactly what most educational materials get wrong.

The practical takeaway here is not that you should go rewrite your attention kernels tomorrow. It is that understanding the progression from FA1 to FA4 is now a tractable task, one that does not require a decade of CUDA experience. If you have ever wanted to know what actually changed, or why each iteration matters, this repo hands you the keys. Start with the FA1 baseline, trace the deferred normalization in FA2, watch the pipeline take shape in FA3, and see the scheduler emerge in FA4. The code is the argument, and it is a convincing one. Go read it.

From Machine Learning

I recently updated my FlashAttention-PyTorch repo so it now includes educational implementations of FA1, FA2, FA3, and FA4 in plain PyTorch.

The main goal is to make the progression across versions easier to understand from code.

Read the original at Machine Learning