Resources
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
Maximum Direct-P forward throughput relative to BF16 FlashAttention-4.
Fast Direct-P attention speedup relative to the BF16 baseline.
Complete single-GPU Llama 3.1-style 8B update speedup at local batch 4.
NVFP4-projection plus FP8-P/V throughput ratio versus BF16 across the 100B-token schedule.
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 paperMore in AI Hardware
Browse all 34 papers →AI as a Compiler: Compiling Triton kernels without the Triton compiler
François Costa, Charly Castes, Thomas Bourgeat, Azalia Mirhoseini
An LLM learns to replace parts of the GPU compiler stack by translating Triton code directly into fast, verified PTX kernels.
Purlin: Separating Orchestration from the Datapath of Collectives
Osayamen Jonathan Aimuyo, Swapnil Gandhi, Christos Kozyrakis
Purlin makes GPU collective communication more modular and faster, improving large-scale LLM and diffusion inference across modern hardware.
RESOLVE: Language-Agnostic Validation of GPU Kernels Through Testing, Reduction, and Proof
Ashkan Vedadi Gargary, Guido Martínez, Sebastian Burckhardt, Gabriel Ebner, Abhinav Jangda, Madan Musuvathi, Tyler Sorensen
RESOLVE makes AI-written GPU kernels safer by combining race-finding tests with formal proofs that optimized code still computes the right result.