Skip to content
Research on Earth

FlashDance: Flash Attention vs Vanilla Attention

How much faster fused attention really is

Status: in progress. The benchmark code is complete. The numeric results have not been recorded yet, so this post covers scope, method and qualitative findings. Tables will be added once the runs are saved.

Question permalink

Flash Attention computes attention without materializing the full N x N score matrix in high-bandwidth memory. The question is how large the gain is in practice, and at what sequence length it starts to matter.

Setup permalink

Two implementations are compared on the same inputs:

  • Vanilla: explicit softmax(QK^T / sqrt(d)) V with a causal mask.
  • SDPA: PyTorch scaled_dot_product_attention, which dispatches to a fused kernel (Flash Attention on CUDA) when the inputs allow it.

Default configuration: batch size 4, 8 heads, head dimension 64. Timing uses the median over repeated runs, with explicit device synchronization on MPS.

What is measured permalink

DimensionRange
Forward pass timesequence length 128 to 4096
Backward pass timesame range
Peak memorysame range
Head dimension32, 64, 128
Batch size1 to 16
Precisionfloat32, float16, bfloat16
Kernel breakdownTorch profiler

Findings so far permalink

  • The speedup of SDPA over vanilla attention grows with sequence length.
  • Below roughly 256 tokens the difference is small.
  • From about 2K tokens, SDPA uses noticeably less time and memory.
  • The backward pass speedup is often larger than the forward pass speedup.
  • On CUDA, the Flash Attention kernel is selected automatically through SDPA.
  • float16 and bfloat16 add further speedup with minimal precision loss.

Extended scope permalink

The repository later grew past the original comparison. It now includes implementations of RoPE, ALiBi, grouped query attention (GQA), multi-query attention (MQA), multi-head latent attention (MLA), sliding window attention, cross-attention and a KV cache, each with tests.

It also contains analyses of IO-bound behavior, memory use, prefill vs decode, latency percentiles (p50 to p999), attention entropy and attention sinks, head specialization, and long-context scaling with NTK and position interpolation. A small decoder combines GQA, RoPE, RMSNorm and SwiGLU.

Pending permalink

  • Record forward and backward results with hardware specified.
  • Add plots for speedup and memory vs sequence length.
  • Repeat on a CUDA GPU, where the Flash Attention kernel is available, to compare against Apple Silicon.

Code: marcelo-earth/flashdance