Skip to content

About

A minimal correct implementation of flash attention with CUDA.

Resources

Stars

1 star

Watchers

0 watching

Forks

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Mini Flash Attention

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.

Basic Benchmark

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

Usage

source .venv/bin/activate 
uv pip install -e . --no-build-isolation
python bench.py

TODOs

  • The Backward pass
  • Tiling the dimension as well as the sequence length
  • Float16
  • Vector operations

About

A minimal correct implementation of flash attention with CUDA.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages