Featured projects

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 with TLX (Triton Low-level Extensions), which add explicit, hardware-aware control on top of Triton’s high-level, tile-based programming model. 

Attention is the single slowest kernel in GEM, and reaching peak Blackwell performance has traditionally required hand-written CuteDSL or CUDA — slow to develop and hard to extend to new variants. We show that TLX closes this gap on both fronts. On development efficiency, the TLX attention kernel is about 3.2K lines of concise Triton-level code — roughly 3× less than the ~10K-line CuteDSL kernels of the state-of-the-art FlashAttention-4 (FA4). On performance, it outperforms FA4 (May 2026 version)  on the jagged shapes that matter for GEM — by ~13% on the forward pass and ~50% on the backward pass. And by staying in Triton’s high-level model, the kernel is approachable enough for modeling engineers, not just kernel specialists, to read, extend, and fuse.

We organize the work in two parts — the structural changes the TLX rewrite makes possible, and the optimizations we layer on top — all benchmarked in bfloat16 on B200 against FA4, the current state of the art on Blackwell.

Code available at: https://github.com/facebookresearch/ads_model_kernel_library/tree/main/tlx_jfa

Introduction

Meta’s ads models, including the Generative Ads Model (GEM) [2] and the Kunlun architecture [5], run attention over jagged (variable-length or ragged) user sequences. As described in our GEM training deep dive, padding these sequences to a fixed length can waste up to 50% of compute. GEM instead packs sequences contiguously and records each sequence’s boundaries in an offsets tensor.

Jagged Flash Attention is the kernel-level mechanism that makes this padding-free representation efficient: it applies the FlashAttention algorithm directly to packed Q/K/V tensors and their offsets, without ever materializing padded tokens. The earlier GEM post described this challenge at the training-system level; here, we focus on optimizing the underlying attention kernel for Blackwell.

Hitting peak throughput on Blackwell means keeping the tensor cores continuously fed. Plain Triton leaves most of these decisions to the compiler; TLX [1] exposes them as first-class primitives (explicit SMEM/TMEM allocation, async_task warp specialization, barriers, async TMA and MMA, and CLC). The rest of the post uses that control to turn a compiler-scheduled, memory-bound kernel into a tightly pipelined, warp-specialized one.

1. The Challenge: Performant and Extensible Attention on Blackwell

Optimizing attention on Blackwell means solving two problems at once. The first is performance: attention is the single slowest kernel in GEM, and it reaches peak throughput only when the global-memory loads, softmax, and matmuls are tightly overlapped. The second is iteration speed: attention is where new modeling ideas land (sliding window, block-sparse, and other variants), so the kernel is edited often — yet the state-of-the-art path to that performance is a hand-written CuteDSL/CUDA kernel that is hard to read, extend, or fuse. Our aim with TLX is to strike the balance: competitive performance in a kernel that stays easy to iterate on.

Why a compiler-scheduled baseline falls short. Our starting point is a mature Triton-based JFA kernel — algorithmically correct, but it leaves all on-chip data movement and scheduling to the compiler: no control over shared-memory allocation or pipeline depth, no explicit barriers or warp specialization, and a single homogeneous warp group doing loads, softmax, and matmuls at once, so the tensor cores stall whenever it is not issuing an MMA. Breaking that ceiling is exactly what TLX’s explicit control enables.

The production case, broadcast-Q, further shapes the design: a single dense Q is broadcast across every sequence in the batch, so its gradient dQ must be summed across the entire batch — turning the dQ epilogue into a heavily contended cross-program reduction that several of our optimizations exist to make fast. Unless noted otherwise, the rest of the post assumes broadcast_q=True.

2. Optimizing Jagged Flash Attention with TLX

Our work falls into two categories. Structural changes (§2.1) are the ground-up reorganization of the kernel that TLX makes possible — they change how work and memory are laid out across the CTA’s warps. Optimizations (§2.2) keep that structure fixed and squeeze out additional performance, mostly by removing stalls and wasted work. We motivate each by the bottleneck it targets rather than the mechanics of its implementation.

2.1 Structural changes

Warp specialization, explicit shared-/tensor-memory management, and barrier pipelining are standard ingredients of any high-performance Blackwell attention kernel, so we keep this brief — what matters is that TLX lets us express them in high-level Triton code rather than hand-written CuteDSL or CUDA, which is what made the §2.2 optimizations fast to prototype. Concretely, we split the CTA into role-specialized async tasks — dedicated warps for TMA loads, tensor-core matmuls, softmax/correction math, and the epilogue store, plus a dedicated dQ-reduction warp in the backward — so the tensor cores stay fed with back-to-back MMAs while softmax runs concurrently, instead of one warp group stalling the matmuls as in the baseline.

The same explicit control extends to memory and scheduling: we hand-allocate every on-chip buffer and choose its pipeline depth (e.g. triple-buffered K/V so the load warp runs ahead), alias TMEM buffers with non-overlapping lifetimes (QK scores, P, and the softmax statistics share one allocation while the PV accumulator keeps its own), hand data between warps through explicit producer/consumer barriers for a deterministic pipeline, and make both passes persistent — one CTA per SM looping over tiles. None of this is expressible in plain Triton, and the persistent structure in particular gives us a place to make our own scheduling decisions — the foundation for §2.2.

2.2 Optimizations

With the warp-specialized, persistent structure in place, the remaining work is to eliminate the stalls and wasted work that keep the tensor cores from staying busy. The following optimizations are largely independent and each targets a specific bottleneck.

We found them with a consistent methodology: NVIDIA Nsight Compute (NCU) to read hardware counters such as tensor-core/TMEM pipeline utilization, SM-busy, and register spills (which show up as local-memory traffic); the ptxas dump logs to confirm spills; and TritonBench in profiler-measure mode for clean before/after A/B latency. The optimizations fall into two groups. Some we adopt from the FA4 implementation — which itself shows that TLX can express, in readable Python, the same techniques a hand-written CuteDSL kernel uses. Others we derive ourselves, from profiling and analysis of our own kernel, targeting the jagged broadcast-Q workload specifically.

At a glance, the optimizations below are:

  • Scheduling jagged tiles across SMs — software load balancing plus Cluster Launch Control keep every SM busy despite the order-of-magnitude length skew of jagged inputs (forward and backward).
  • Multi-stage dQ staging — a double-buffered SMEM staging pipeline hides the heavily contended broadcast-Q dQ reduce-add, the #1 backward bottleneck.
  • Early tensor-memory release — freeing the dQ tensor-memory buffer before its final stores lets the MMA warp start the next tile’s dQ matmul sooner.
  • Loop peeling — splitting the KV loop into a branch-free bulk pass and a tiny masked tail removes per-iteration mask overhead and the register spills it causes.
  • 2-CTA collaborative MMA — two CTAs cooperate on a single matmul to raise tensor-core utilization in the matmul-heavy backward (adopted from FA4).

Scheduling jagged tiles across SMs

Jagged inputs create severe load imbalance across SMs: with sequence lengths varying by orders of magnitude, a naive tile-to-SM mapping leaves some SMs idle while others grind through the long sequences. An SM-occupancy heatmap (Triton-MPP) makes this visual — a few rows stay hot long after the rest go idle. A truly greedy rebalancing is near-optimal but cannot be vectorized, so it blows up host-side launch overhead.

In the forward pass the outer loop is over the dense Q and the inner loop over the jagged K/V, so a tile’s cost is proportional to its batch’s key length — the tiles are highly non-uniform. Balancing this is our own idea for the jagged case: on the host we sort all tiles by KV workload (descending) and hand them to SMs in a zigzag (boustrophedon) pattern — even passes fill SMs left-to-right, odd passes right-to-left — so every SM ends up with a balanced mix of long and short tiles at negligible CPU cost. This recovered roughly 20% on the forward kernel.

The backward needs a different strategy. Its dK/dV pass loops over K/V blocks on the outside and the dense Q on the inside, so now the K/V blocks are the tiles: each one is roughly uniform in cost (a K/V block against the shared dense Q), but the number of tiles per batch varies with the jagged key length. Sorting by workload no longer helps — instead we precompute the list of valid (batch, head, K/V-block) tiles on the host and round-robin them across SMs, using a block-innermost ordering so consecutive tiles reuse the same K/V in L2, and padding each head’s block count to a multiple of the cluster size for the 2-CTA path. Distributing a variable count of equal-cost tiles this way keeps every SM busy without any sorting.

Static balancing handles the skew we can predict ahead of time; the residual variance that only shows up at runtime is handled dynamically by Cluster Launch Control (CLC) [3]. A Blackwell feature, CLC hands out the next tile index on demand, so whichever SM finishes first grabs the next piece of work and no single long sequence stalls the grid.

CTAs workload heapmat without CLC

CTAs workload heatmap with CLC

We apply CLC to both the forward and backward kernels; in TLX it is a lightweight producer/consumer protocol over a shared scheduling context.

Although software load balancing and CLC may look redundant — both spread work across SMs — they are complementary: CLC schedules dynamically yet cannot tell which tiles are empty, so on jagged (and especially sparse) inputs it still spends cycles handing out the sentinel/empty tiles that padding introduces, whereas our host-side valid-tile precompute prunes those up front. That is why software load balancing stays valuable even on a CLC-capable kernel.

Multi-stage dQ staging

Triton-MPP barrier analysis of the backward pinned the dQ epilogue as the #1 bottleneck: its reduce-add to HBM accounted for ~9–11% of lost tensor-core utilization, the single largest actionable gap. Two things make it slow. First, because dQ is summed across the whole batch, every SM issues reduce-adds into the same dQ locations — heavy contention. Second, the high-level reduce-add API serializes them — each column slice’s add must finish before the next begins, with no pipelining.

We make the staging explicit and double-buffered so that one slice’s reduce-add to HBM overlaps the next slice’s copy out of TMEM, keeping a store in flight at all times:

NCOL   = BLOCK_D // (EPILOGUE_SUBTILE * 2)   # narrow column slices, sized to fit SMEM
STAGES = 2                                   # double-buffered staging
dq_smem = tlx.local_alloc((BLOCK_M, NCOL), dq.dtype, STAGES)

for s in range(BLOCK_D // NCOL):
    # TMEM -> registers
    dq = tlx.local_load(dq_tmem[:, s * NCOL : (s + 1) * NCOL]) * LN2

    # registers -> SMEM (ping-pong)
    tlx.local_store(dq_smem[s % STAGES], dq.to(dq.dtype))
    tlx.fence_async_shared()

    tlx.async_descriptor_store(desc_dq, dq_smem[s % STAGES],
                               offsets, store_reduce="add")  # SMEM -> HBM, accumulating

    # keep at most 1 store in flight
    tlx.async_descriptor_store_wait(STAGES - 1)

The technique generalizes to any shared-memory-bound reduce-store epilogue, not just attention.

Early tensor-memory release

The same analysis showed why that epilogue stalls the math from the TMEM side: the reduction warp holds the dQ tensor-memory buffer until all of its slices have been drained to HBM, so the MMA warp blocks on the dq-empty barrier before it can start the next tile’s dQ matmul. A waterfall test that simply removed the dQ store confirmed it, recovering +8–11% of tensor-core utilization across sequence lengths.

The improvement is to pre-load the last one or two slices into registers and release the TMEM buffer before issuing their stores, so the MMA warp can begin reusing that memory while the reduction warp finishes writing to HBM:

# Drain most slices the normal way
# (TMEM -> SMEM -> async reduce-add to HBM).
for s in range(N_SLICES - EARLY_RELEASE_SUBTILES):
    reduce_add_slice(s)

# Pre-load the final 1-2 slices into registers FIRST, then release
# the dQ TMEM buffer immediately
# so the MMA warp can start the next tile's dQ matmul.
dq_tail = tlx.local_load(dq_tmem[:, last_slice])   # TMEM -> registers
tlx.barrier_arrive(dq_empties[buf])                # <-- TMEM freed early
reduce_add_from_registers(dq_tail)

How many slices to release early is autotuned (1 or 2); the optimal choice is small because releasing more raises register pressure enough to hurt. Together with multi-stage staging, this lets the MMA warp begin the next dQ matmul while the previous tile’s dQ is still being written out.

Loop peeling

NCU profiling showed the forward was MMA-issue starved — not memory- or compute-bound — with the tensor-core (TMEM) pipeline at ~57% utilization versus FA4’s ~82% on the same shape, and the lost cycles coming from softmax-warp overhead bubbling into the MMA-issue gaps. Anything that shrinks or simplifies the softmax body therefore lets the MMA warp pick up work sooner.

Two signals point at the culprit. First, the ptxas dumps and NCU show register spills (local-memory traffic) in the register-heavy softmax warp — in an earlier round, simply granting the warps more registers cut that traffic by ~40% for a ~6% forward gain, confirming the pressure is real. Second, the worst offender is a per-iteration mask branch in the KV loops that only ever matters on the partial last tile.

The reason this branch is costly is subtle and specific to how the compiler allocates registers: Triton/TLX assign registers to variables statically, so a mask branch living inside the hot loop forces the compiler to reserve registers for the variables used only in the rarely-taken masked path (the column offsets, the mask tensor, the select) across the entire loop. That permanent reservation raises register pressure and can trigger spills — especially damaging for the register-heavy softmax warp.

The fix is to peel the loop into a branch-free bulk pass plus a tiny masked tail. Because the mask flag is a compile-time constant, the bulk body contains no mask variables at all, so the compiler frees those registers and can schedule a tighter MMA-issue cadence:

# APPLY_MASK is constexpr: with APPLY_MASK=False the compiler statically removes the
# comparison, the offs_n / mask tensors, and the select -> a branch-free body that
# does not reserve registers for the (rare) masked path.
aligned = (klen // BLOCK_N) * BLOCK_N

for start_n in tl.range(lo, aligned, BLOCK_N):  # bulk: straight-line, no mask code
    softmax_iter(start_n, ..., APPLY_MASK=False)

for start_n in tl.range(aligned, klen, BLOCK_N): # tail: 0 or 1 iteration
    softmax_iter(start_n, ..., APPLY_MASK=True)  # mask only the partial last tile

We apply the same peeling to the register-heavy paths in both the forward softmax loop and the backward, where the gain comes not from skipping a cheap runtime branch but from the reduced register pressure it unlocks; in the backward it specifically reclaims the ~9% latency that the correctness masks would otherwise add.

2-CTA collaborative MMA (backward)

The backward is matmul-heavy — computing dQ, dK, and dV takes five GEMMs per K/V block — and a single CTA does not fully utilize the Blackwell tensor cores on its own. We adopt FA4’s 2-CTA (paired-CTA tcgen05) scheme: two CTAs in a cluster cooperate on one wider collaborative matmul over two adjacent K/V blocks of the same (batch, head), splitting the accumulator rows across the two SMs and exchanging the dS intermediate on-chip over distributed shared memory (DSMEM) rather than through global memory. Feeding a single MMA from both SMs raises tensor-core utilization on the backward.

What is ours here is the TLX implementation: porting the scheme to our jagged, broadcast-Q layout and running it under our persistent and CLC multi-tile schedulers.

Scoped to the production broadcast-Q, HEAD_DIM=128 case, it adds about +12% throughput (−11% latency) in the backward over the single-CTA path.

Performance

We benchmark on B200 (bf16) in two regimes: the production jagged case (Hierarchical Seed Pooling (HSP) — a broadcast dense Q against jagged K/V) and an LLM-style dense case (equal-length Q/K/V), both against FA4, the state-of-the-art open-source CuteDSL FlashAttention-4 [6] kernel (May 2026 version). On the jagged shapes — the ones that matter for Ads — TLX outperforms FA4 on both passes: the forward is faster across most shapes and sparsities (about +13% on average, trailing only on the longest sequences at high density), and the backward is faster everywhere, by about +50% on average. On the dense shapes the forward is competitive (~87% of FA4) while the backward again wins (~+17%). The jagged charts sweep sparsity from 0.1 (highly variable lengths) through 0.5 to 0.9 (nearly uniform lengths, with 1.0 meaning all sequences are the same length) — the range seen in production.

Jagged (broadcast-Q) — forward TFLOPS (bf16, B200)

Jagged (broadcast-Q) — backward TFLOPS (bf16, B200)

LLM dense — forward TFLOPS (bf16, B200, B=768, H=4, head_dim=128)

LLM dense — backward TFLOPS (bf16, B200, B=768, H=4, head_dim=128)

3. Attention Variants: the flexibility of TLX

A practical benefit of building on TLX is that the kernel is structured enough to fork for new requirements without a ground-up rewrite. Because the warp specialization, memory allocation, barriers, and scheduling are written once in readable Python-level code, retargeting the kernel to a new numerical format or a new attention pattern mostly comes down to changing the math, not the machinery. Two variants illustrate this.

Low-precision (MXFP8) attention

For FP8-tolerant training we built a microscaling-FP8 variant of both the forward and backward by swapping the BF16 matmuls for TLX’s block-scaled MMA (FP8 E4M3 data with E8M0 per-block scale factors), reusing the same warp-specialized skeleton, barriers, and TMEM layout. The softmax output P is quantized to FP8 on the fly and its scale factors are kept in tensor memory so the block-scaled MMA can consume them directly.

After tuning, the MXFP8 forward lands above FA4’s own FP8 kernel and the backward reaches parity with FA4 at dense — both comfortably beating the BF16 baseline — a large throughput win obtained largely by changing the GEMM calls rather than the kernel structure.

Block-sparse attention

For long jagged histories where full O(N²) attention dominates, we built a two-stage sparse variant [4]. A cheap scoring kernel average-pools each Q and K block and selects the top-k most relevant KV blocks per Q block (a tunable selection ratio); a TLX attention kernel — forked from the dense one with per-tile sparse block iteration added — then runs attention over only the selected blocks.

It reuses the same warp-specialized structure, CLC dispatch, and software load balancing (here used to skip unselected and empty tiles), stays PT2-friendly, and still supports broadcast-Q, GQA, and windowing. At a 0.5 selection ratio the forward is roughly 1.3–1.5× faster than dense across the sequence lengths we tested.

In both cases the bulk of the kernel carried over unchanged; only the math — the GEMM precision, or the set of KV blocks each tile visits — differed. That reuse is the practical payoff of TLX: the same kernel can be re-targeted to new precisions and sparsity patterns while staying at the Python level, instead of dropping down to hand-written CUDA.

4. Summary and Future Work

The structural changes — warp specialization, persistent execution, explicit SMEM/TMEM allocation with reuse aliasing, and barrier-mediated pipelining — set the stage, and the optimizations layered on top — software load balancing, Cluster Launch Control, multi-stage dQ staging, early TMEM release, loop peeling, and 2-CTA collaborative MMA — recovered most of the remaining headroom. Together they let the TLX kernel match or beat FA4 on both passes for the production jagged shapes, while retaining full broadcast-Q, sliding-window, and GQA support.

Beyond raw performance, the biggest win is development efficiency. Much of the low-level plumbing that CuteDSL writes by hand — async pipelines, barrier management, tcgen05 MMA setup — is instead handled by the TLX compiler, so we express intent rather than machinery, in a fraction of the code. That leverage lets the team prototype, profile, and land every optimization above quickly, and then fork the same kernel into MXFP8 and block-sparse variants with little extra code. Taken together — faster on both passes for the jagged Ads shapes, competitive on dense, and far cheaper to build and maintain — TLX JFA is a better fit than the SOTA open-source FA4 for our Ads models.

References

[1] TLX: Enabling Cluster Launch Control with Triton — https://pytorch.org/blog/enabling-cluster-launch-control-with-tlx/

[2] GEM: Meta’s Generative Ads Model — https://engineering.fb.com/2025/11/10/ml-applications/metas-generative-ads-model-gem-the-central-brain-accelerating-ads-recommendation-ai-innovation/

[3] NVIDIA Blackwell Tuning Guide (thread-block clusters) — https://docs.nvidia.com/cuda/blackwell-tuning-guide/index.html#thread-block-clusters

[4] TLX Block Attention: A Warp-Specialized Blackwell Kernel for Fixed-Block Sparse Self-Attention — https://pytorch.org/blog/tlx-block-attention-a-warp-specialized-blackwell-kernel-for-fixed-block-sparse-self-attention/

[5] Kunlun: Establishing Scaling Laws for Massive-Scale Recommendation Systems through Unified Architecture Design — https://arxiv.org/abs/2602.10016

[6] FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling — https://arxiv.org/abs/2603.05451