AivexaNewsSearch
AI news for builders and product teamsChecked every hour
PyTorchFirst partyIndustry

Optimizing Jagged Flash Attention with TLX: The Road Toward SOTA FA4 on Blackwell

Collected Oct 1, 2026

Meta has published work on Jagged Flash Attention (JFA), the attention kernel behind its Generative Ads Model (GEM), running on NVIDIA Blackwell (B200) and built with TLX (Triton Low-level Extensions). TLX adds explicit, hardware-aware control on top of Triton's high-level, tile-based programming model. The kernel is benchmarked in bfloat16 against FlashAttention-4 (FA4, May 2026 version), described as the current state of the art on Blackwell. Code is available in Meta's ads_model_kernel_library repository under tlx_jfa.

Two headline numbers frame the result. The TLX attention kernel is about 3.2K lines of Triton-level code, roughly 3x less than the ~10K lines of CuteDSL in FA4. On the jagged shapes that matter for GEM, it outperforms FA4 by about 13% on the forward pass and about 50% on the backward pass. On dense, LLM-style shapes the forward is competitive at roughly 87% of FA4, while the backward wins by about 17%. The jagged charts sweep sparsity from 0.1 through 0.5 to 0.9.

GEM runs attention over jagged, variable-length user sequences, packing them contiguously and recording boundaries in an offsets tensor; padding to fixed length can waste up to 50% of compute. JFA applies the FlashAttention algorithm directly to packed Q/K/V tensors and offsets without materializing padded tokens. Attention is described as the single slowest kernel in GEM. The production case, broadcast-Q, broadcasts one dense Q across the batch, so dQ must be summed across the entire batch, making the dQ epilogue a heavily contended cross-program reduction.

The structural rewrite splits the CTA into role-specialized async tasks: dedicated warps for TMA loads, tensor-core matmuls, softmax/correction, epilogue store, plus a dQ-reduction warp in the backward. On-chip buffers are hand-allocated with chosen pipeline depth, such as triple-buffered K/V, and both passes are persistent with one CTA per SM looping over tiles.

Optimizations include software load balancing for jagged tiles, recovering roughly 20% on the forward kernel; Cluster Launch Control, a Blackwell feature that hands out tile indices on demand; double-buffered staging of the dQ reduce-add; early tensor-memory release, autotuned to one or two slices; loop peeling of the KV loop into a branch-free bulk pass plus a masked tail; and a 2-CTA collaborative MMA scheme adopted from FA4, adding about 12% backward throughput in the broadcast-Q, HEAD_DIM=128 case.

Why it matters: developers working on attention kernels get a path to near-state-of-the-art Blackwell performance without hand-written CuteDSL or CUDA, in code Meta says modeling engineers can read, extend and fuse. The staging technique is noted as generalizing to any shared-memory-bound reduce-store epilogue.

Read at PyTorch

Based on reporting from the original publisher. Visit the source for full context and later updates.

Publisher excerpt

TL;DR In this blog post, we present our work on Jagged Flash Attention (JFA) — the attention kernel behind Meta’s Generative Ads Model (GEM) — on NVIDIA Blackwell (B200), built... The post Optimizing Jagged Flash Attention with TLX: The Road Toward SOTA FA4 on Blackwell appeared first on PyTorch .