How to Use FlashAttention-2 in PyTorch

FlashAttention-2 is a rewritten attention algorithm that computes the exact same result as standard attention but 2–4x faster and with memory usage that scales linearly with sequence length instead of quadratically. It achieves this by tiling the attention computation to stay within GPU SRAM rather than reading from the slower HBM memory repeatedly. For anyone training or fine-tuning transformer models, enabling FlashAttention-2 is one of the highest-leverage changes you can make — it directly extends the sequence lengths you can afford and reduces training cost with no accuracy tradeoff.

Installation

# Requires CUDA 11.6+ and PyTorch 2.0+
pip install flash-attn --no-build-isolation

# Verify
python -c "import flash_attn; print(flash_attn.__version__)"

# For Ampere+ GPUs (A100, RTX 3090+) — fastest path
# For older GPUs (V100, T4) — works but slower; less beneficial

# Hugging Face Transformers support (no manual install needed if using HF):
pip install transformers accelerate

Build time is 10–20 minutes — it compiles CUDA kernels from source. Use a pre-built wheel from the flash-attention GitHub releases page if you want to skip the build: wheels are provided for common CUDA/PyTorch/Python combinations.

Drop-in Usage with Hugging Face Transformers

For models supported in Transformers (LLaMA, Mistral, Phi, Falcon, GPT-NeoX and many others), FlashAttention-2 is a single argument:

from transformers import AutoModelForCausalLM, AutoTokenizer
import torch

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-8B",
    torch_dtype=torch.bfloat16,
    attn_implementation="flash_attention_2",  # ← this is all you need
    device_map="auto"
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")

inputs = tokenizer("FlashAttention makes this fast:", return_tensors="pt").to("cuda")
with torch.no_grad():
    out = model.generate(**inputs, max_new_tokens=50)
print(tokenizer.decode(out[0], skip_special_tokens=True))

Direct Flash Attention API

For custom attention implementations, use the flash_attn_func directly:

from flash_attn import flash_attn_func, flash_attn_qkvpacked_func
import torch

batch, seqlen, nheads, headdim = 2, 4096, 32, 128

# Q, K, V must be float16 or bfloat16 — flash attn does not support float32
q = torch.randn(batch, seqlen, nheads, headdim, device='cuda', dtype=torch.float16)
k = torch.randn(batch, seqlen, nheads, headdim, device='cuda', dtype=torch.float16)
v = torch.randn(batch, seqlen, nheads, headdim, device='cuda', dtype=torch.float16)

# Equivalent to scaled dot-product attention
out = flash_attn_func(q, k, v, dropout_p=0.0, causal=True)
# out shape: (batch, seqlen, nheads, headdim)

# Packed QKV variant (more memory efficient)
qkv = torch.randn(batch, seqlen, 3, nheads, headdim, device='cuda', dtype=torch.float16)
out = flash_attn_qkvpacked_func(qkv, dropout_p=0.0, causal=True)

Integrating into a Custom Transformer

import torch
import torch.nn as nn
from flash_attn import flash_attn_func

class FlashAttentionLayer(nn.Module):
    def __init__(self, d_model, n_heads, dropout=0.0):
        super().__init__()
        self.n_heads = n_heads
        self.head_dim = d_model // n_heads
        self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
        self.out_proj = nn.Linear(d_model, d_model, bias=False)
        self.dropout = dropout

    def forward(self, x, causal=True):
        B, T, C = x.shape
        qkv = self.qkv(x)
        # Reshape to (batch, seqlen, 3, nheads, headdim)
        qkv = qkv.reshape(B, T, 3, self.n_heads, self.head_dim)
        q, k, v = qkv.unbind(dim=2)  # each: (B, T, nheads, headdim)

        # FlashAttention expects bfloat16 or float16
        q, k, v = q.to(torch.bfloat16), k.to(torch.bfloat16), v.to(torch.bfloat16)

        out = flash_attn_func(q, k, v, dropout_p=self.dropout if self.training else 0.0, causal=causal)
        out = out.reshape(B, T, C).to(x.dtype)
        return self.out_proj(out)

Memory and Speed Comparison

FlashAttention-2’s advantage grows with sequence length. Standard attention’s memory usage is O(N²) in sequence length — at 8k tokens on a 40GB A100, standard attention consumes roughly 16GB for the attention matrix alone; FlashAttention-2 uses about 2GB for the same computation. This is why long-context models (32k, 128k tokens) became practical once FlashAttention was available — without it, the attention matrix does not fit in GPU memory at all for long sequences.

Figure 1 — FlashAttention-2 memory usage vs standard attention by sequence length

Memory usage (GB) on A100 80GB — LLaMA-7B, batch=1 2k 4k 8k 16k Standard O(N²) FlashAttn O(N) OOM →

Benchmarking the Speedup

import torch, time
from flash_attn import flash_attn_func

def benchmark_attn(use_flash, batch=4, seqlen=2048, nheads=32, headdim=128, n_iters=100):
    dtype = torch.float16
    q = torch.randn(batch, seqlen, nheads, headdim, device='cuda', dtype=dtype)
    k = torch.randn(batch, seqlen, nheads, headdim, device='cuda', dtype=dtype)
    v = torch.randn(batch, seqlen, nheads, headdim, device='cuda', dtype=dtype)

    # Warmup
    for _ in range(10):
        if use_flash:
            flash_attn_func(q, k, v, causal=True)
        else:
            scale = headdim ** -0.5
            attn = torch.einsum('bshd,bthd->bsht', q*scale, k).softmax(-1)
            torch.einsum('bsht,bthd->bshd', attn, v)
    torch.cuda.synchronize()

    t0 = time.perf_counter()
    for _ in range(n_iters):
        if use_flash:
            flash_attn_func(q, k, v, causal=True)
        else:
            scale = headdim ** -0.5
            attn = torch.einsum('bshd,bthd->bsht', q*scale, k).softmax(-1)
            torch.einsum('bsht,bthd->bshd', attn, v)
    torch.cuda.synchronize()
    return (time.perf_counter() - t0) / n_iters * 1000

for seqlen in [512, 1024, 2048, 4096]:
    std_ms = benchmark_attn(False, seqlen=seqlen)
    fa2_ms = benchmark_attn(True,  seqlen=seqlen)
    print(f"seqlen={seqlen:5d}: standard={std_ms:.2f}ms  flash={fa2_ms:.2f}ms  speedup={std_ms/fa2_ms:.1f}x")

FlashAttention-2 in Fine-Tuning with PEFT

When fine-tuning with LoRA or QLoRA, FlashAttention-2 is fully compatible and dramatically reduces memory pressure, allowing larger batch sizes or longer sequences during training:

from transformers import AutoModelForCausalLM, BitsAndBytesConfig
from peft import get_peft_model, LoraConfig
import torch

bnb_config = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16)

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-8B",
    quantization_config=bnb_config,
    attn_implementation="flash_attention_2",  # FA2 + 4-bit quantization
    device_map="auto"
)
lora_config = LoraConfig(r=16, lora_alpha=32, target_modules=["q_proj","v_proj"])
model = get_peft_model(model, lora_config)

Requirements and Limitations

FlashAttention-2 requires float16 or bfloat16 — it does not work with float32. Sequence lengths must be multiples of 8 (padding may be needed for variable-length inputs). It requires a GPU with Ampere architecture or newer (A100, H100, RTX 3000/4000 series) for full performance — it works on older GPUs (V100, T4) but with smaller speedups. Head dimension must be a power of 2, up to 256. These constraints are minor in practice — most modern LLM training already uses bfloat16 and Ampere-class hardware.

FlashAttention-2 is one of the few techniques in deep learning that gives you a free performance improvement with zero accuracy tradeoff — it computes exactly the same result as standard attention, just with a smarter memory access pattern. For any transformer training or inference workload on NVIDIA GPU, enabling it via the attn_implementation="flash_attention_2" argument in Transformers, or the direct flash_attn_func in custom code, should be an automatic first step before exploring other optimisations.

Figure 1 — Memory scaling: FlashAttention-2 (linear) vs standard attention (quadratic)

Attention memory (GB) — A100 80GB, LLaMA-7B, batch=1 2k 4k 8k 16k Standard O(N²) OOM boundary FA-2 O(N)

How FlashAttention-2 Works

Standard attention computes the full N×N attention matrix, which must be materialised in GPU high-bandwidth memory (HBM). For a sequence length of 8,192 tokens with 32 heads and float16, this matrix is roughly 16GB — and it must be written to HBM after the softmax and read back for the weighted sum. The memory bandwidth bottleneck is the dominant cost, not the floating point operations. FlashAttention-2 avoids materialising the full attention matrix by tiling the computation: it processes the attention in blocks that fit in the fast SRAM (on-chip memory), computes the softmax incrementally using the online softmax algorithm, and accumulates the output directly — the intermediate N×N matrix is never written to HBM. The result is exactly equivalent to standard attention but with O(N) memory and 2–4x fewer HBM reads/writes. FlashAttention-2 specifically improves on the original by better parallelism across the sequence dimension and less redundant computation, making it faster on non-power-of-two sequence lengths and on H100 tensor cores.

FlashAttention-2 vs torch.nn.functional.scaled_dot_product_attention

PyTorch 2.0+ includes torch.nn.functional.scaled_dot_product_attention (SDPA) which automatically uses FlashAttention-2 when it is installed and the inputs are compatible. This means many models that use SDPA internally already benefit from FlashAttention-2 without explicit configuration:

import torch
import torch.nn.functional as F

q = torch.randn(2, 32, 4096, 128, device='cuda', dtype=torch.float16)
k = torch.randn(2, 32, 4096, 128, device='cuda', dtype=torch.float16)
v = torch.randn(2, 32, 4096, 128, device='cuda', dtype=torch.float16)

# PyTorch will use FlashAttention-2 automatically when available
# for compatible inputs (float16/bf16, appropriate head dim)
with torch.backends.cuda.sdp_kernel(
    enable_flash=True,    # use FlashAttention-2
    enable_math=False,    # disable slow fallback
    enable_mem_efficient=False
):
    out = F.scaled_dot_product_attention(q, k, v, is_causal=True)

# Check which backend was used
print(torch.backends.cuda.flash_sdp_enabled())

The SDPA route is the most portable option — it falls back to a memory-efficient attention implementation or standard attention depending on what the hardware and input format support, with no code changes required. For maximum control and slightly better performance, use the flash_attn package directly.

Practical Impact on Training Cost

The real-world impact depends on sequence length and model size. For a standard 2,048-token fine-tuning run, FlashAttention-2 typically reduces per-step time by 20–30% and memory usage by 30–40% compared to standard attention — not transformative, but meaningful. The impact becomes dramatic at longer sequence lengths: at 32k tokens, standard attention on a 7B model exceeds a single A100’s 80GB memory and requires model sharding; FlashAttention-2 keeps it within a single GPU. This makes FlashAttention-2 the enabling technology for long-context fine-tuning, allowing practitioners to train models on full-document context that was previously impossible without expensive multi-GPU setups. The combination of FlashAttention-2 with 4-bit quantisation (bitsandbytes) and LoRA (PEFT) is the standard recipe for fine-tuning large language models on consumer or single-datacenter-GPU hardware.

Variable-Length Sequences (Padding-Free)

FlashAttention-2 supports padding-free (varlen) batching, which packs sequences of different lengths into a single batch without wasting compute on padding tokens. This gives significant throughput improvements for datasets with variable-length inputs like instruction fine-tuning data:

from flash_attn import flash_attn_varlen_func
import torch

# Pack sequences without padding
# cu_seqlens: cumulative sequence lengths, shape (batch+1,)
# max_seqlen: length of the longest sequence in the batch

cu_seqlens_q = torch.tensor([0, 512, 1024, 2048], dtype=torch.int32, device='cuda')
cu_seqlens_k = cu_seqlens_q.clone()
max_seqlen_q = 1024   # longest sequence
max_seqlen_k = 1024

total_tokens = 2048   # sum of all sequence lengths
nheads, headdim = 32, 128
q = torch.randn(total_tokens, nheads, headdim, device='cuda', dtype=torch.float16)
k = torch.randn(total_tokens, nheads, headdim, device='cuda', dtype=torch.float16)
v = torch.randn(total_tokens, nheads, headdim, device='cuda', dtype=torch.float16)

out = flash_attn_varlen_func(
    q, k, v,
    cu_seqlens_q, cu_seqlens_k,
    max_seqlen_q, max_seqlen_k,
    causal=True
)

Hugging Face Transformers handles varlen packing automatically when you pass packing=True to the SFTTrainer from TRL — it is compatible with attn_implementation="flash_attention_2" and can increase throughput by 20–40% on instruction fine-tuning datasets where sequences vary widely in length.

Troubleshooting FlashAttention-2

The most common issues when enabling FlashAttention-2. ImportError after installation — flash-attn is compiled against a specific CUDA and PyTorch version; verify the installed versions match with python -c "import torch; print(torch.version.cuda)" and reinstall if needed. “Input tensor must be on CUDA” — FlashAttention-2 is GPU-only; CPU tensors raise this error. “Input tensor must be half precision” — switch to float16 or bfloat16 before the attention call. Slower than expected — on GPUs older than Turing (pre-2018), FlashAttention-2 falls back to a slower implementation; check flash_attn.utils.benchmark.benchmark_forward against standard attention to verify you are getting a speedup. Incorrect outputs with causal=False — verify your masking logic; FlashAttention-2’s bidirectional attention does not apply any mask by default, so if your use case requires custom masking, pass it explicitly. Most issues stem from dtype or device mismatches, which are immediately raised as clear errors rather than silent failures.

FlashAttention-3 and Future Developments

FlashAttention-3 (released in 2024) adds specific optimisations for H100 hardware: asynchronous execution using Hopper’s warp-group MMA instructions, FP8 support for additional speed at reduced precision, and improved handling of variable-length sequences. The performance gains are primarily on H100 — on A100 hardware, FlashAttention-2 and FlashAttention-3 deliver similar throughput. Both versions are available from the same flash-attn package; the appropriate kernel is selected based on the detected GPU at runtime. If you are on A100 or earlier hardware, FlashAttention-2 is the practical ceiling; on H100, installing the latest flash-attn package automatically gives you FlashAttention-3 kernels where beneficial.

FlashAttention-2 is one of the few techniques in deep learning that delivers a free performance improvement with zero accuracy tradeoff — the same result, just computed with a smarter memory access pattern. For any transformer workload on NVIDIA Ampere or newer, enabling it is an automatic first step before exploring any other optimisation. The speedup compounds with sequence length, making it especially valuable for long-context models where standard attention would otherwise be the dominant memory and compute cost.

The combination of FlashAttention-2, 4-bit quantisation, and LoRA is now the standard recipe for fine-tuning large language models on limited hardware — each technique addresses a different constraint (compute time, memory capacity, and trainable parameter count respectively), and together they make 7B–13B model fine-tuning accessible on a single consumer GPU.

For anyone building on transformers today, knowing when to enable FlashAttention-2, how to verify it is actually being used, and how to fall back gracefully when it is not available is a practical skill worth internalising early — it touches every serious training and inference workflow.

Leave a Comment