NTH

Hardware-Aware FP4 FlashAttention-4

AuthorsRobert Hu

September 8, 2026 2 min read
Watch on YouTube
The one-line take

This work redesigns FlashAttention for Blackwell’s FP4 hardware, achieving faster inference and training while showing that aggressively quantized distributed training can become unstable.

Key results

2.13
GB200 forward speedup

Maximum Direct-P forward throughput relative to BF16 FlashAttention-4.

2.09
Wan2.1-14B speedup

Fast Direct-P attention speedup relative to the BF16 baseline.

1.14
8B update speedup

Complete single-GPU Llama 3.1-style 8B update speedup at local batch 4.

1.112
Distributed throughput ratio

NVFP4-projection plus FP8-P/V throughput ratio versus BF16 across the 100B-token schedule.

512
TMEM allocation

Blackwell D128 tensor-memory columns consumed by two score and two output banks.

What the paper found

Hardware-Aware FP4 FlashAttention-4 targets the bottleneck that remains after NVIDIA Blackwell tensor cores accelerate the QK and PV matrix products: softmax conversion, probability scaling, synchronization, and tensor-memory ownership. Its Direct-P method keeps NVFP4 for Q and K, maps log-softmax scores directly into MXFP4 E2M1 probability codes using packed FMA and native conversion, and computes the denominator from the same rounded probabilities consumed by PV. On favorable NVIDIA GB200 shapes, this reaches up to 2.13× BF16 forward throughput, while Wan2.1 inference reaches 2.09× on the 14B model, with greater latent drift than FP8 controls. For causal training, forward saves quantized Q/K payloads, scales, and LSE values for backward, which reconstructs probabilities and uses FP8 gradient operands, including E5M2 output gradients. In a complete Llama 3.1-style 8B update, speedup reaches 1.14× at local batch 4. Distributed training over 64 GPUs completes a 100B-token schedule at 1.112× the BF16 throughput, but FP8 P/V is selected for stability because every tested MXFP4 P/V trajectory diverged. The hardware analysis identifies the 512-column TMEM allocation—two score banks and two output banks—as the main limit on overlap, implying that future gains require another allocatable score destination or finer-grained K32 PV issue semantics, not merely faster arithmetic.

Original abstract

Blackwell's 4-bit floating-point (FP4) tensor cores do not automatically make attention faster because softmax conversion and on-chip dependencies dominate once its matrix products shrink. We address this with \emph{Direct-P} for noncausal inference and a causal path that passes the forward quantization directly into backward. Direct-P maps scores directly to FP4 probabilities and reaches up to 2.13$\times$ the bfloat16 (BF16) forward throughput on an NVIDIA GB200. The causal path reconstructs probabilities from saved quantized queries and keys and uses 8-bit floating-point (FP8) gradient operands, accelerating a complete single-GPU 8-billion-parameter update by up to 1.14$\times$. Matched distributed training retains FP8 probabilities and values; every tested MXFP4 probability/value training trajectory diverges.

Read the original paper

More in AI Hardware

Browse all 34 papers →