Meta's TLX Beats FlashAttention-4 on Blackwell GPUs by 50%
PyTorch's new Jagged Flash Attention kernel on Blackwell beats FlashAttention-4 by 13% forward and 50% backward using Triton Low-level Extensions.
- PyTorch and Meta released a Jagged Flash Attention kernel built with TLX for Blackwell B200.
- Beats FlashAttention-4 by ~13% forward and ~50% backward on jagged ads shapes.
- 3.2K lines of Triton-level code vs FA4's ~10K lines of CuteDSL, roughly 3x smaller.
- Powers Meta's Generative Ads Model (GEM) attention over variable-length user sequences.
- Key tricks: warp specialization, zigzag tile scheduling, double-buffered dQ staging, early TMEM release, loop peeling.
- Code available on GitHub, with MXFP8 and block-sparse variants forked from the same base.
Meta’s TLX attention kernel outpaces FlashAttention-4 on B200
Engineers from Meta and PyTorch rebuilt the attention kernel behind Meta’s Generative Ads Model (GEM) for NVIDIA’s Blackwell B200 GPU. Their Jagged Flash Attention implementation, written with TLX extensions, outperformed the May 2026 version of FlashAttention-4 (FA4) on GEM’s production jagged workloads while using about one-third as much code.
Attention is GEM’s slowest kernel, according to the team’s engineering report. High throughput requires the kernel to overlap memory transfers, softmax calculations, and tensor-core matrix multiplications. Developers have typically needed CuteDSL or CUDA to control that scheduling. TLX exposes comparable hardware controls through Python-level Triton code.
TLX exposes Blackwell’s machinery
Standard Triton uses a tile-based programming model and leaves scheduling, shared-memory allocation, and pipelining largely to the compiler. TLX adds explicit control over shared memory (SMEM), Blackwell tensor memory (TMEM), warp specialization, synchronization barriers, and asynchronous data movement and matrix operations.
- Warp specialization: The
async_taskprimitive assigns separate warps to loading, matrix multiplication, softmax, correction math, and output storage. - TMA: NVIDIA’s Tensor Memory Accelerator moves multidimensional tiles between high-bandwidth memory (HBM) and on-chip memory asynchronously.
- MMA: Matrix multiply-accumulate instructions execute the attention kernel’s tensor-core operations.
- Cluster Launch Control: Blackwell hardware can reassign queued work to whichever streaming multiprocessor (SM) becomes available first.
The resulting kernel contains about 3,200 lines of Triton-level code, compared with roughly 10,000 lines for FA4’s CuteDSL implementation. On the tested jagged workloads, the team reported average gains of about 13% for the forward pass and 50% for the backward pass.
GEM’s sequences do not fit a fixed grid
GEM processes user histories whose lengths vary substantially. Padding every sequence to a common length can spend as much as half of the available compute on empty positions. The model instead packs sequences into contiguous tensors and stores their boundaries in an offsets tensor. Jagged Flash Attention operates directly on those packed values without materializing padding.
GEM also broadcasts one dense query across every sequence in a batch. During backpropagation, each sequence contributes to the same query gradient, dQ, so the kernel must reduce partial results across the batch. Concurrent programs contend for that shared output, making the reduction epilogue a major backward-pass bottleneck.
Specialized warps keep tensor cores busy
The engineers divided each cooperative thread array (CTA), CUDA’s term for a thread block, into role-specific asynchronous tasks. Dedicated warps handle TMA loads, tensor-core operations, softmax and correction calculations, output storage, and the backward pass’s dQ reduction. This organization allows matrix operations to continue while other warps process softmax or move data.
Explicit on-chip memory allocation supports a three-stage K/V pipeline, allowing the load warp to fetch future inputs while tensor cores process the current tile. Buffers with non-overlapping lifetimes share TMEM capacity. Both passes use persistent execution, with one CTA per SM repeatedly claiming tiles instead of launching a separate CTA for every tile.
Five changes unlocked the gains
NVIDIA Nsight Compute profiling identified load imbalance, reduction traffic, tensor-memory lifetimes, and register pressure as the main sources of lost utilization. The team addressed each source separately:
- Balance jagged tiles before launch. The CPU sorts tiles by key-value workload, then assigns them to SMs in a back-and-forth, or boustrophedon, order. Each SM receives a mixture of long and short tiles. This scheduling change improved forward performance by about 20%.
- Assign remaining work dynamically. Cluster Launch Control gives the next queued tile to an SM as soon as it finishes its current work. The mechanism complements the initial static ordering and reduces idle time when tile costs remain uneven.
-
Pipeline the query-gradient reduction. Writing and adding
dQvalues in HBM accounted for roughly 9% to 11% of lost tensor-core utilization. A double-buffered SMEM pipeline keeps one reduction store in flight while the kernel prepares the next. -
Release tensor memory earlier. The kernel moves the final one or two
dQslices into registers before releasing their TMEM allocation. The matrix warp can then begin the next tile sooner, improving reported tensor-core utilization by 8% to 11%. - Peel the masked loop tail. The main K/V loop runs without masking, while a separate final iteration handles incomplete tiles. Removing the branch from the common path prevents the compiler from reserving registers for a rarely used masked path on every iteration.
A separate two-CTA collaborative MMA path pairs two thread blocks within a cluster to process one wider matrix multiplication. For GEM’s production configuration with a broadcast query and a head dimension of 128, this path raised backward throughput by about 12% and reduced latency by about 11% compared with the single-CTA implementation.
Benchmark gains concentrate on jagged workloads
The team benchmarked bfloat16 kernels on an NVIDIA B200 against the May 2026 version of FA4. Reported averages varied by workload:
| Workload | Pass | TLX result versus FA4 |
|---|---|---|
| Jagged | Forward | About 13% faster on average; slower on the longest, densest cases |
| Jagged | Backward | About 50% faster on average across all tested shapes |
| Dense | Forward | About 87% of FA4 throughput |
| Dense | Backward | About 17% faster on average |
These results apply to the tested B200 shapes, data types, sparsities, and FA4 snapshot. They do not establish the same advantage for other GPUs or attention configurations, and the dense forward result shows that the largest gains come from GEM’s jagged workload and broadcast-query backward pass.
The same skeleton supports lower precision and sparsity
The warp-specialized structure also served as the base for two variants, allowing the developers to retain the scheduling and memory pipeline while replacing the attention computation:
- MXFP8 attention: Block-scaled matrix operations use E4M3 FP8 values with E8M0 scale factors for each block. The reported forward pass exceeded FA4’s FP8 implementation, while the backward pass reached similar performance on dense shapes.
- Block-sparse attention: A scoring kernel selects the top k key-value blocks for each query block, after which attention runs only on the selected blocks. Processing half of the blocks produced a forward speedup of roughly 1.3 to 1.5 times over dense attention.
A practical Blackwell playbook
TLX gave the team enough control to schedule Blackwell’s memory transfers, tensor-core operations, and persistent work queues without moving the implementation into raw CUDA. The smaller codebase also made it possible to reuse the same kernel structure for FP8 and block-sparse variants.
Several techniques have broader use on Blackwell: dynamic tile assignment can help workloads with uneven task sizes, staged epilogues can hide contended reduction stores, early TMEM release can shorten dependencies between tiles, and loop peeling can reduce register pressure. The zigzag schedule depends on predictable jagged workloads, while the two-CTA path and dQ pipeline target GEM’s broadcast-query configuration.
Meta has released the kernel source, including the TLX implementation and its workload-specific paths, for developers who want to inspect or adapt the design.