FFPA: Fast and Memory-Efficient Exact Attention for Large Headdim, achieving O(1) SRAM complexity (w/ Split-D) and O(d/4) register complexity, 1.5x~15x speedup over PyTorch SDPA. FFPA extends the headdim support beyond D > 256 (up to 1024) without any precision loss.
| Self Attn | GQA/MQA | Cross Attn | Causal/Mask | Dropout | Headdim | Fwd/Bwd |
|---|---|---|---|---|---|---|
βοΈ(Nq=Nkv) |
βοΈ(Hq!=Hkv) |
βοΈ(Nq!=Nkv) |
βοΈ(attn_mask) |
βοΈ(p>0) |
320~1024 | 1.5x~15xβ |
- [2026-08] π Add FFPA FP8/FP4 benchmark results (compare with Sage-2/3) for NVIDIA RTX 5090, PRO 5000 and 6000, achieving significant speedup for FP4 Attention (~1000 TOPS for D=128 on PRO 6000). ππ
- [2026-08] π Cache-DiT x FFPA (FP8/FP4) is ready! Feel free to take a try for your Diffusion models. ππ
- [2026-08] πͺ FFPA now experimental supports FP4 Attention for headdims [64,1024] (sm_120, forward only), achieving 850-980π TOPS (D=128-256) on NVIDIA RTX 5090, 3.8x~4.4xπ speedup over PyTorch SDPA (FlashAttention-2 backend), the performance of large headdims is stay tuned for updates. ππ
- [2026-08] π¦ FFPA now supports D=512 for NVIDIA B200 via CuTe-DSL tcgen05 2-CTA, 1517 TFLOPS forward and 763 TFLOPS backward, achieving 6x~15xπ speedup over standard PyTorch SDPA. ππ
- [2026-07] π― FFPA now supports FP8 Attention for headdims [64,1024] (sm_120, forward only) and achieving 3x~6xπ speedup over PyTorch SDPA for large headdim (D>256). ππ
- [2026-06] FFPA now supports AMD ROCm/HIP GPUs via the TritonBackend, check #268 for more details. π
- [2026-06] π¦ NVIDIA-Nemo/AutoModel x FFPA achieving 1.4x~1.5xπ End2End training throughput speedup for Gemma4-31B (8xH200, FSDP2 + AC) with FFPA accelerating the 10/60 (D=512) full-attention layers. ππ
- [2026-06] π FFPA now supports TritonBackend and CuTeDSLBackend for both forward and backward pass, achieving 1.5x~5xπ speedup over standard PyTorch SDPA across many devices. ππ
- [2026-05] πͺ FFPA now supports GQA, MQA, cross-attn, causal, attn-mask and dropout with CUDABackend for large headdims (D>256, forward only), achieving 1.3x~2xπ speedup over PyTorch SDPA. ππ
First, install the prebuilt package from PyPI or build ffpa-attn from source:
# First, install the prebuilt package from PyPI
pip3 install -U ffpa-attn # CUDA 13.0+, PyTorch 2.11+
# Or, build ffpa-attn from source, just follow the cmds
git clone https://github.com/xlite-dev/ffpa-attn.git
# Then, build the wheel package (Triton + CuTe-DSL backends)
cd ffpa-attn && pip3 install -e . --no-build-isolation
# Optional: install ffpa-attn w/ CUDA backend (forward only)
# ext all: build all kernels, include fp8/fp4 attention kernels
bash ./build.sh --arch sm_120f --ext all --headdim allThen, try to accelerate the attention for large headdim with just one-line of code:
>>> import torch.nn.functional as F
>>> from ffpa_attn import ffpa_attn_func
>>> # Monkey-patch SDPA to point to FFPA. Every thing that FFPA
>>> # does not support will auto fallback to SDPA: N < 512, etc.
>>> F.scaled_dot_product_attention = ffpa_attn_funcOr, try the minimal BF16 usage example β Self-Attention (B=1, H=32, N=8192, D=512):
import torch
import torch.nn.functional as F
from ffpa_attn import ffpa_attn_func
# D: 64, 128, ..., 320, ..., 1024 (FA-2 <= 256, FFPA supports up to 1024).
B, H, N, D = 1, 32, 8192, 512 # batch_size, num_heads, seq_len, head_dim
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device="cuda")
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device="cuda")
# FFPA self attention; layout follows SDPA: (B, H, N, D).
out = ffpa_attn_func(q, k, v) # -> torch.Tensor of shape (B, H, N, D)
ref = F.scaled_dot_product_attention(q, k, v)
print(f"FFPA vs SDPA max_abs_err={(out - ref).abs().max().item():.4e}")Or, try the minimal FP8/FP4 usage example with CUDABackend (sm_120, forward only):
import torch
import torch.nn.functional as F
from ffpa_attn import CUDABackend, ffpa_attn_func
from functools import partial
# D: 64, 128, ..., 320, ..., 1024 (FA-2 <= 256, FFPA supports up to 1024).
B, H, N, D = 1, 32, 8192, 128 # batch_size, num_heads, seq_len, head_dim
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device="cuda")
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device="cuda")
# Currenly, fp8/fp4 attention are only supported on sm_120, forward only.
fp8_backend = CUDABackend(backward=False, forward=True, enable_fp8=True)
fp4_backend = CUDABackend(backward=False, forward=True, enable_fp4=True)
ffpa_attn_func_fp8 = partial(ffpa_attn_func, forward_backend=fp8_backend)
ffpa_attn_func_fp4 = partial(ffpa_attn_func, forward_backend=fp4_backend)
# FFPA self attention; layout follows SDPA: (B, H, N, D).
out_fp8 = ffpa_attn_func_fp8(q, k, v) # -> torch.Tensor of shape (B, H, N, D)
out_fp4 = ffpa_attn_func_fp4(q, k, v) # -> torch.Tensor of shape (B, H, N, D)
ref = F.scaled_dot_product_attention(q, k, v)
print(f"FFPA FP8 vs SDPA max_abs_err={(out_fp8 - ref).abs().max().item():.4e}")
print(f"FFPA FP4 vs SDPA max_abs_err={(out_fp4 - ref).abs().max().item():.4e}")For more advanced features, please refer to our online docs at πffpa-attn.io.
We extend FlashAttention to support large headdim (
Split-D: The tiling of the
TiledMMA: The M4N2 layout breaks the register bottleneck. The
Dispatch: M8N1 for
Runnable benchmark are provided under bench. The performance benchmarks for the NVIDIA L20 (Ada), NVIDIA Geforce RTX 5090 (Blackwell), NVIDIA H800 PCIE (Hopper), NVIDIA H200 SXM (Hopper, CuTe-DSL backend, up to 535 TFLOPS!), B200 (Blackwell, CuTe-DSL tcgen05 2-CTA D=512 backend, up to 1517 TFLOPS forward and 763 TFLOPS backward!) with large headdims can be found at bench.
FFPA supports multiple backends for the forward and backward pass, including: SDPA (baseline), CUDA (forward only), Triton, and CuTe-DSL. The CuTe-DSL backend is currently in early stage, stay tuned for future updates. The Triton backend (forward + backward) also runs on AMD GPUs.
| Backend | Arch | Fwd | Bwd | Headdim | Autotune | Speedup | Recommend |
|---|---|---|---|---|---|---|---|
| SDPA | sm>=75 | β | β | All | βοΈ | 1.0x | sm>=75 |
| CUDA | sm>=80 | β | βοΈ | 320~1024 | βοΈ | 1.5x~3x | sm_80~89,120{a,f} |
| CUDA FP8 | sm_120{a,f} | β | βοΈ | 64~1024 | βοΈ | 3x~6x | sm_120{a,f} |
| CUDA FP4 | sm_120{a,f} | β | βοΈ | 64~512 | βοΈ | 4x~7x | sm_120{a,f} |
| Triton | sm>=80 | β | β | 320~1024 | β | 1.5x~5x | sm>=80 |
| CuTe-DSL | sm>=80 | β | β | 320~1024 | βοΈ | 1.5x~2x | sm_80~89,120{a,f} |
| CuTe-DSL | sm_90a | β | β | 320~512 | βοΈ | 3x~6x | sm_90a |
| CuTe-DSL | sm_100a | β | β | 512 | βοΈ | 6x~15x | sm_100a |
How to use different backends for your own scenario? Users can simply pass the Backend configs (SDPABackend, CUDABackend, TritonBackend or CuTeDSLBackend) to ffpa_attn_func, for example:
>>> from ffpa_attn import ffpa_attn_func, CuTeDSLBackend
>>> # CuTe-DSL backend, D=512 scenario, fastest on H200!
>>> o = ffpa_attn_func(q, k, v, backend=CuTeDSLBackend())Generate device-specific tuned configs for production deployment (currently, Triton only), avoiding per-process autotune cost. The generated JSON is saved under configs dir and automatically loaded when runtime autotune is disabled (the default). See the docs of Triton Autotune for details.
python -m ffpa_attn.autotune --mode max --full-tasks --overwrite # 1 GPU
export CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7 # Multi-GPU (`pip install ray`)
python -m ffpa_attn.autotune --mode max --full-tasks --num-gpus 8 --overwriteNVIDIA-NeMo Automodel PR #2436 shows that on Gemma4-31B training (L=8192, 8xH200, FSDP2 + Activation Checkpointing), accelerating the 10/60 (D=512) full-attention layers with FFPA delivers about 1.4x~1.5x higher throughput (E2E) than SDPA at similar memory footprint, with loss aligned within normal bf16 noise.
The FFPA (FP8/FP4) attention has fully integrated into Cache-DiT. Currently, the FP8/FP4 attention supports most of the attention headdims range from 64 to 1024 (Sage-2/3 only supports D<=128), including any headdims that can be div by 8 (e.g, 120), covering self-attention, cross, causal and GQA/MQA (Sage-3 does not support).
The kernel benchmark results show that FFPA FP8 is comparable or slightly better than SageAttention-2 at D=128, and FFPA FP4 is significantly better than SageAttention-3 at D=128 on NVIDIA RTX PRO 5000/6000/5090. Please check π§± How to Reproduce for more details. Feel free to take a try for your Diffusion models.
python3 -m cache_dit.generate flux --attn native --seed 42 --height 1024 --width 1024
python3 -m cache_dit.generate flux --attn ffpa_fp8 --seed 42 --height 1024 --width 1024
python3 -m cache_dit.generate flux --attn ffpa_fp4 --seed 42 --height 1024 --width 1024FLUX.1-dev, seed=42, 28 steps, 1024 x 1024, NVIDIA RTX PRO 5000
| SDPA (17.19s) | FFPA-FP8 (16.08s) | Sage-2 (FP8, 16.21s) | FFPA-FP4 (15.99s) |
|---|---|---|---|
![]() |
![]() |
![]() |
![]() |
FLUX.1-dev, seed=42, 28 steps, 2048 x 2048, NVIDIA RTX PRO 5000
| SDPA (91.25s) | FFPA-FP8 (79.91s) | Sage-3 (FP4, 80.39s) | FFPA-FP4 (75.73s) |
|---|---|---|---|
![]() |
![]() |
![]() |
![]() |
The performance and precision of FFPA (FP8/FP4) is still under active development, stay tuned for future updates. Please note that the FP8/FP4 attention is not suitable for all scenarios (e.g., small models or short seqlen), and we recommend users to evaluate the precision and performance of FFPA (FP8/FP4) for their own use cases.
Apache License 2.0
@misc{deftruth2026ffpa,
author = {DefTruth and Butterfingrz},
title = {FFPA: Fast and Memory-Efficient Exact Attention for Large Headdim},
year = {2026},
publisher = {Zenodo},
version = {v1.0},
doi = {10.5281/zenodo.20638547},
url = {https://doi.org/10.5281/zenodo.20638547}
}



























