Research

Meta Optimizes Jagged Flash Attention on Blackwell

Meta researchers optimized Jagged Flash Attention for NVIDIA Blackwell GPUs using Triton Low-level Extensions, beating FlashAttention-4 performance with far simpler code.

PyTorch Blog8 hrs agoResearch
Image: PyTorch Blog

Meta has developed a highly optimized Jagged Flash Attention (JFA) kernel for NVIDIA Blackwell B200 GPUs using Triton Low-level Extensions (TLX). JFA is the core attention mechanism powering Meta's Generative Ads Model (GEM) and Kunlun architecture, which process variable-length user sequences without wasteful padding. Traditionally, hitting peak performance on Blackwell required writing complex CUDA or CuteDSL code. By using TLX, Meta created a kernel with just 3,200 lines of code—about three times shorter than the 10,000 lines required for the state-of-the-art FlashAttention-4 (FA4) library.

In benchmarks on B200 hardware using bfloat16, the new TLX kernel outperformed the May 2026 version of FA4 on the jagged workloads critical for ads models. It achieved a 13 percent speedup on the forward pass and a 50 percent improvement on the backward pass. On standard dense shapes, the forward pass reached 87 percent of FA4's speed, while the backward pass won by 17 percent. The team achieved these gains through structural changes like warp specialization and several targeted optimizations. For instance, a host-side zigzag sorting pattern for load balancing recovered 20 percent on the forward pass, while early tensor-memory release reclaimed 8 to 11 percent of tensor-core utilization.

Other key optimizations included loop peeling to reduce register pressure, which reclaimed 9 percent latency in the backward pass, and a two-CTA collaborative matrix multiplication that added 12 percent throughput. Because TLX keeps the code in readable Python, the researchers easily adapted the kernel for other use cases. They built a microscaling FP8 variant that beats FA4's FP8 forward kernel and matches its backward performance. They also created a two-stage sparse variant that runs 1.3 to 1.5 times faster than dense attention at a 0.5 selection ratio.

For machine learning practitioners, this development bridges the gap between high-level programmability and bare-metal hardware performance. Instead of relying on a small pool of specialized kernel engineers to write inscrutable CUDA code, standard modeling engineers can now read, extend, and fuse attention kernels directly in Python. This dramatically accelerates the iteration cycle for testing new attention variants like sliding window or block-sparse configurations on next-generation Blackwell hardware.

This is our own summary of reporting by PyTorch Blog

More in Research