This repo implements the most prominent features of the forward pass for flash attention. The official repo is very complex, so this one tries to provide a minimal implementation for those learning Cuda.
Note: flash-attention-minimal has the same goal (this repo was inspired from that), however that implementation is slightly incorrect since it stores Sij in the smem of a Cuda block which is not how we properly avoid N*N complexity. This implementation is more true to the algorithm found in the paper.
For Q, K, V matrices of shape [1024, 32] (sequence length 32, dimension 32 per vec) we see that we significantly outperform the base pytorch implementation.
=== PyTorch ===
Self CPU time total: 172.849ms
Self CUDA time total: 242.809us
=== Flash ===
Self CPU time total: 10.107ms
Self CUDA time total: 6.702ms
source .venv/bin/activate
uv pip install -e . --no-build-isolation
python bench.py
- The Backward pass
- Tiling the dimension as well as the sequence length
- Float16
- Vector operations