Skip to main content

Fast and Memory-Efficient Exact Attention for Large Headdim


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โ†‘

Latest News

  • [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. ๐ŸŽ‰๐ŸŽ‰

Quick Start

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 all

Then, 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_func

Or, 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.

Split-D and TiledMMA

We extend FlashAttention to support large headdim ($D>256$) via fine-grained tiling at the MMA level for $QK^\top$ and $PV$ matrix multiplication. Two orthogonal $O(D)$ bottlenecks โ€” SRAM footprint and register pressure โ€” are broken by Split-D and TiledMMA<4,2,1> respectively.

Split-D: The tiling of the $D$ axis breaks the SRAM bottleneck. A persist-D layout keeps $Q$ resident in SRAM at $O(D)$ ($D{=}512 \Rightarrow 192\text{KB} > 99\text{KB}$ per-CTA limit on sm_8x/sm_120). Split-D chunks the $D$ axis, keeping SRAM fixed at $B_r \times 16$ (with $B_r=B_c$) for Q, K and V, yielding constant SRAM complexity $O(B_r \times 16) \approx O(1)$.

TiledMMA: The M4N2 layout breaks the register bottleneck. The $QK^\top$ has $N{=}B_c$ (fixed, independent of $D$), so its acc is $O(1)$; the $PV$ GEMM instead has $N{=}D$, so the $O$ acc costs $D/(2{\cdot}N_w)$ regs/thread. M8N1 (FA-2 style, $N_w{=}1$) $\Rightarrow O(D/2)$: at $D{=}512$ this already reaches 256 regs/thread, over the 255 architectural limit and spilling. Splitting $N$ to M4N2 (FA-1 style, $N_w{=}2$) halves it to $O(D/4)$, keeping $D{=}1024$ just feasible (256 regs/thread).

Dispatch: M8N1 for $D \le 512$, M4N2 for $D > 512$. On RTX 5090, M4N2 delivers 1.55ร— the throughput of M8N1 at $D{=}1024$ (154T vs 100T, where M8N1 collapses from register spilling).

Benchmark

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.


BF16 Attention for Large Headdim: FFPA vs SDPA (FWD/BWD) across NVIDIA H200 and B200, 6x-15xโ†‘.


FP8 Attention for Large/Small Headdim: FFPA vs SDPA (FWD) on NVIDIA RTX 5090, 3x-6xโ†‘.


FP4 Attention for Large/Small Headdim: FFPA vs SDPA (FWD) on NVIDIA RTX 5090, 4x-7xโ†‘.

Backends

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())

Persistent Autotune

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 --overwrite

End-to-End Training

NVIDIA-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.

End-to-End Inference

FP8 Attention for D=128: FFPA vs SageAttention-2 on NVIDIA RTX PRO 5000/6000/5090.




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).

FP4 Attention for D=128: FFPA vs SageAttention-3 on NVIDIA RTX PRO 5000/6000/5090.




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 1024

FLUX.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.

License

Apache License 2.0

Citations

@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}
}

References

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distributions

No source distribution files available for this release.See tutorial on generating distribution archives.

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

ffpa_attn-0.2.3-py3-none-any.whl (519.4 kB view details)

Uploaded Python 3

File details

Details for the file ffpa_attn-0.2.3-py3-none-any.whl.

File metadata

  • Download URL: ffpa_attn-0.2.3-py3-none-any.whl
  • Upload date:
  • Size: 519.4 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.10.12

File hashes

Hashes for ffpa_attn-0.2.3-py3-none-any.whl
Algorithm Hash digest
SHA256 d083f772433e9e6fb0f925b03b129d9d2830a310a86c7b5a0e28ecba277a89aa
MD5 7f826898ed2cc33bb232bd1ac48291ae
BLAKE2b-256 81616b56635c40fd009b450eeb66c4b5021aaf5e61c055993078fec039845f05

See more details on using hashes here.

Release history Release notifications | RSS feed

0.2.4

1 file

This release

0.2.3 This release

1 file

0.2.2

1 file

0.2.1

1 file

0.2.0

6 files

0.1.23

6 files

0.1.22

6 files

0.1.21

6 files

0.1.20

6 files

0.1.19

6 files

0.1.18

5 files

0.1.17

5 files

0.1.16

5 files

0.1.15

5 files

0.1.14

5 files

0.1.13

5 files

0.1.12

5 files

0.1.11

5 files

0.1.10

1 file

0.1.9

1 file

0.1.8

1 file

0.1.7

5 files

0.1.6

5 files

0.1.4

5 files

0.1.3

5 files

0.1.2

1 file

0.1.0

1 file

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page