# How Flash Attention Improves Transformer Memory Efficiency Compared to Standard Attention

> Discover how Flash Attention slashes transformer memory use from O(L²) to O(L) by processing attention block-wise sans full matrix materialization. Learn the technique.

- Repository: [labml.ai/annotated_deep_learning_paper_implementations](https://github.com/labmlai/annotated_deep_learning_paper_implementations)
- Tags: deep-dive
- Published: 2026-03-04

---

**Flash Attention reduces transformer memory consumption from quadratic O(L²) to linear O(L) by processing attention block-wise without materializing the full N×N attention matrix, using only three running statistics per query.**

Flash Attention has become a critical optimization for training and deploying large language models with long context windows. According to the `labmlai/annotated_deep_learning_paper_implementations` source code, this algorithm replaces the standard attention implementation with a memory-efficient kernel that streams over key-value blocks while maintaining numerical stability. Unlike standard attention which stores intermediate matrices of size proportional to the square of sequence length, Flash Attention keeps only a constant amount of state per query regardless of input length.

## Why Standard Attention Consumes Quadratic Memory

Standard transformer attention implementations materialize the full score matrix before applying softmax, creating a severe memory bottleneck for long sequences.

### Full Matrix Construction

In regular attention, the score matrix *S* with shape **B × H × L_q × L_k** is explicitly computed and stored as the intermediate product of queries and keys. As implemented in [`labml_nn/transformers/flash/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/__init__.py) (lines 11-14), this requires allocating a tensor of size proportional to the square of the sequence length *L* before the softmax operation can begin.

### Intermediate Tensor Storage

After computing *S*, the standard implementation stores a second matrix *P* (the post-softmax attention weights) of identical size to weight the values *V*. Both *S* and *P* occupy GPU high-bandwidth memory (HBM) for the entire forward pass. When sequence length grows to 4,000 or 8,000 tokens, these intermediate activations exceed the capacity of consumer GPUs, triggering out-of-memory errors even with moderate batch sizes.

## How Flash Attention Achieves Linear Memory Complexity

Flash Attention reformulates the attention computation as an online softmax algorithm that processes keys and values in small blocks while maintaining only three scalar statistics per query.

### Block-Wise Processing Algorithm

Rather than loading the entire *K* and *V* tensors into memory, Flash Attention loads only a single block of size **BLOCK_K × d_head** at a time. The algorithm iterates over these blocks, updating running statistics in-place, and never stores the full attention matrix *S* or the post-softmax weights *P*. This approach reduces memory complexity from **O(L² · d_head)** to **O(L · BLOCK_K · d_head)**, effectively making memory consumption linear in sequence length.

### The Three Running Statistics

For each query vector, Flash Attention maintains three values that allow incremental softmax computation:

- **m_k** — The current maximum score for each query, updated as `m_i^{new} = max(m_i, max_j S_{ij})` to stabilize exponentials (see [`labml_nn/transformers/flash/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/__init__.py), lines 72-80).
- **l_k** — The running sum of exponentials after re-normalizing, updated as `l_i ← e^{m_i-m_i^{new}} l_i + Σ_j ˜P_{ij}` to accumulate the softmax denominator (lines 83-86).
- **˜O_k** — The unnormalized output accumulator, updated as `˜O_i ← e^{m_i-m_i^{new}} ˜O_i + Σ_j ˜P_{ij} V_j` to weight values incrementally (lines 87-90).

After processing all blocks, the final output is computed with a single division: `O_i = ˜O_i / l_i`.

### Backward Pass Efficiency

The same block-wise approach applies to the backward pass. As documented in [`labml_nn/transformers/flash/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/__init__.py) (lines 118-124), the backward kernel maintains only the necessary statistics (*D_i*, *l_i*, *m_i*) rather than storing the full attention matrices for the reverse computation, preserving the linear memory bound during gradient calculation.

## Implementation in the Repository

The `labmlai/annotated_deep_learning_paper_implementations` repository provides both a pure Python/Triton reference implementation and integration points for production models.

### Core Algorithm and Testing

The primary implementation resides in [`labml_nn/transformers/flash/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/__init__.py), which contains the block-wise algorithm documentation and the `AttentionFunc` autograd function. Performance validation and correctness checks are provided in [`labml_nn/transformers/flash/test.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/test.py), which benchmarks the implementation against standard attention.

### Integration with GPT-Style Models

In [`labml_nn/neox/model.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/neox/model.py) (lines 207-213), the `NeoX` model class accepts an `is_flash_attention` boolean parameter. When set to `True`, the `AttentionLayer` instantiates a `FlashAttention` object instead of standard multi-head attention, enabling memory-efficient training for models with 64 heads and 6,144-dimensional hidden states. Additionally, the half-precision evaluation script in [`labml_nn/neox/evaluation/half_precision.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/neox/evaluation/half_precision.py) accepts a `--flash` command-line flag to enable Flash Attention for benchmarking.

### Diffusion Model Integration

The repository also demonstrates Flash Attention integration in diffusion models via [`labml_nn/diffusion/stable_diffusion/model/unet_attention.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/diffusion/stable_diffusion/model/unet_attention.py), where the `CrossAttention.use_flash_attention` property toggles the optimization for UNet attention layers.

## Practical Usage Examples

The repository demonstrates Flash Attention integration through concrete code examples for custom layers and pre-built models.

### Basic Flash Attention Forward Pass

To use Flash Attention directly in a custom transformer block:

```python
import torch
from labml_nn.transformers.flash import AttentionFunc

def flash_self_attn(q, k, v, causal=True):
    """
    q, k, v – tensors of shape (B, H, L, D)
    Returns attention output of shape (B, H, L, D)
    """
    scale = 1.0 / (q.size(-1) ** 0.5)
    return AttentionFunc.apply(q, k, v, causal, scale)

# Example tensors for batch=2, heads=8, seq=4096, dim=64

B, H, L, D = 2, 8, 4096, 64
q = torch.randn(B, H, L, D, device='cuda')
k = torch.randn(B, H, L, D, device='cuda')
v = torch.randn(B, H, L, D, device='cuda')

out = flash_self_attn(q, k, v, causal=True)
print(out.shape)  # torch.Size([2, 8, 4096, 64])

```

### Enabling Flash Attention in NeoX Models

For the pre-implemented GPT-NeoX architecture:

```python
from labml_nn.neox.model import NeoX

model = NeoX(
    n_hidden=6144,
    n_heads=64,
    n_layers=48,
    is_flash_attention=True,  # Enables Flash Attention

)

x = torch.randn(4, 2048, 6144, device='cuda')
logits = model(x)  # Forward pass uses Flash Attention internally

```

### Measuring Memory Efficiency

To benchmark memory usage against standard attention:

```python
import torch
from labml_nn.neox.model import NeoX

def benchmark(seq_len):
    model_std = NeoX(n_hidden=6144, n_heads=64, is_flash_attention=False)
    model_flash = NeoX(n_hidden=6144, n_heads=64, is_flash_attention=True)
    
    x = torch.randn(1, seq_len, 6144, device='cuda')
    
    torch.cuda.reset_peak_memory_stats()
    _ = model_std(x)
    mem_std = torch.cuda.max_memory_allocated() / 1e9
    
    torch.cuda.reset_peak_memory_stats()
    _ = model_flash(x)
    mem_flash = torch.cuda.max_memory_allocated() / 1e9
    
    print(f'Seq {seq_len}: std {mem_std:.2f} GB vs flash {mem_flash:.2f} GB')

for L in [1024, 2048, 4096, 8192]:
    benchmark(L)

```

This script typically shows approximately **50% reduction** in peak GPU memory for long sequences, confirming the theoretical linear memory scaling.

## Summary

- **Flash Attention** eliminates the quadratic memory bottleneck of standard attention by never materializing the full *S = QKᵀ* matrix or post-softmax weights *P*.
- The algorithm achieves **O(L) memory complexity** by processing key-value blocks sequentially while maintaining only three running statistics per query: the max score *m*, the exponential sum *l*, and the unnormalized output *˜O*.
- According to [`labml_nn/transformers/flash/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/__init__.py), the implementation uses block sizes of **BLOCK_K × d_head** rather than full sequence lengths, enabling contexts of 8,000+ tokens on single GPUs.
- The repository provides seamless integration via the `is_flash_attention=True` flag in [`labml_nn/neox/model.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/neox/model.py), with automatic fallback to Triton kernels when CUDA Flash Attention is unavailable.
- Beyond memory savings, the reduced HBM bandwidth usage typically yields **2-3× speedup** on modern GPUs compared to standard attention implementations.

## Frequently Asked Questions

### What is the exact memory complexity of Flash Attention compared to standard attention?

**Standard attention** requires **O(L² · d_head)** memory to store the score matrix *S* and attention weights *P* for sequence length *L* and head dimension *d_head*. **Flash Attention** reduces this to **O(L · BLOCK_K · d_head)**, which is linear in sequence length since *BLOCK_K* is a small constant (typically 64 or 128) independent of *L*. This allows processing sequences of 8,000 tokens or more on hardware that would exhaust memory with standard attention at 4,000 tokens.

### Does Flash Attention produce mathematically identical results to standard attention?

Yes. Flash Attention computes the exact same softmax attention output as the standard implementation, just via an online algorithm. The three running statistics (*m*, *l*, and *˜O*) mathematically reconstruct the standard softmax normalization `softmax(S) @ V` through the final division `O_i = ˜O_i / l_i`. The numerical differences are within floating-point rounding error, not algorithmic approximation.

### How do I enable Flash Attention in the labmlai implementation?

Pass `is_flash_attention=True` when instantiating supported models like `NeoX` in [`labml_nn/neox/model.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/neox/model.py). For custom layers, import `AttentionFunc` from `labml_nn/transformers/flash` and call `AttentionFunc.apply(q, k, v, causal, scale)`. The repository automatically detects if the external `flash_attn` package is installed and uses the optimized CUDA implementation; otherwise, it falls back to the included Triton reference kernel.

### Is Flash Attention faster than standard attention, or just more memory efficient?

Flash Attention typically provides both benefits. By reducing high-bandwidth memory (HBM) reads and writes—from quadratic to linear in sequence length—the kernel becomes compute-bound rather than memory-bound on modern GPUs. As shown in [`labml_nn/transformers/flash/test.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/flash/test.py), this yields **2-3× throughput improvements** in addition to the memory savings, making it advantageous even when memory capacity is not the limiting factor.