Building a Faster GQA Decode Kernel for Blackwell SM100
LLM inference begins with prefill, which processes the prompt and builds a key and value cache. Decode then generates new tokens using that cached history. Grouped query attention, used in models such as Llama 3.1 and Qwen3, reduces the cache by letting several query heads share a KV head. Reading that cache is still a recurring cost at every attention layer and generation step.
The Colfax post on S/P ping pong describes a scheduling improvement for FA4 decode on Blackwell. In the original schedule, for the next KV block waits for softmax on the current block, despite having no mathematical dependency on it. Alternating between two score and probability slots in tensor memory allows those operations to overlap. The post reports gains of up to 16% on supported configurations, with the implementation in FA4 PR #2817.
The S/P ping pong schedule from Colfax. The next QK overlaps softmax for the current block.
FA4 with S/P ping pong enabled. The highlighted QK issue ranges overlap the softmax path.
Here, is issued while the softmax path for block is still active, as intended by S/P ping pong. FA4 also uses split P arrival to signal that an initial portion of is ready, allowing the PV MMA to start before the remaining probabilities are written. This early start still depends on correction releasing the output accumulator.
This implementation addresses a different limitation. PackGQA groups query heads so they can reuse the same KV tile, but the packed query dimension can still be much smaller than FA4's chosen tile. With 64 query heads and eight KV heads, a group supplies only eight useful rows to a 128 row tile. The remaining rows are padding in the computation, not extra tokens in the cache.
The CUTLASS GQA decode design changes the operand order to compute and . This puts KV positions on the large matrix dimension and the small query head group on the other. For an eight head group and a block of 128 KV tokens, the score tile can be instead of . Softmax still reduces over KV positions.
This smaller score tile also reduces the data read from TMEM. The softmax threads keep their score fragments in registers and release the TMEM slot before computing the probabilities. FA4 already retains scores in registers too; the advantage here is the smaller score footprint, not a new caching mechanism. The simple GQA pipeline publishes the complete smaller tile to shared memory in one handoff rather than using split P arrival. It reduces padded score work and movement without that additional signaling scheme.
I implemented this kernel using the CUTLASS design and removed the workspace memset from the kernel reduction path. My kernel reaches up to 2.26× the FA4 performance in my B200 BF16 single token benchmarks. The PR for this kernel is here. Those results include the later changes I will explain here, not just the simple pipeline below.
Starting with the simple pipeline
I will start with CUTLASS gqa_decode_simple and explain the changes from there. Its two matmuls are shown as KQ and VP below. There is one softmax warpgroup and one output accumulator; correction must rescale that accumulator before its next update.
A shared timeline for the simple kernel. The aligned lanes show which loads and computations can overlap. Grey blocks show VP waiting for its inputs and corrected output. Durations are illustrative, not measured.
The load and MMA roles below run concurrently for a nonempty split. These are the simple kernel's actual loops and calls, with tensor view construction, pipeline initialization and the final statistics writeback omitted. The comments describe the warm-up and the shared ring handoffs.
if warp_idx == tma_qo_warp_id: # warp 11
for k in cutlass.range_constexpr(KQ_num_k_tiles):
gQ_k = tBgQ[None, None, None, k]
smem_Q_q_idx = smem_Q[None, None, None, k]
cute_ext.tma_load(gQ_k, smem_Q_q_idx, q_load_mbar_ptr.value, update_expect_tx=False)
O_final_nbar.arrive_and_wait() # correction has staged the final split output
for dm in cutlass.range_constexpr(VP_num_m_tiles):
# Select gmem_O_part_local for this dm tile.
cute_ext.tma_store(smem_O_tma[None, None, dm, 0], gmem_O_part_local)
cute.arch.cp_async_bulk_commit_group()
cute.arch.cp_async_bulk_wait_group(0, read=True)
elif warp_idx == tma_kv_warp_id: # warp 10
# prefetch_iters = 2, prefetch_tiles = 2 * kv_splits.
# Split-local order: K0 | K1 | K2 V0 | K3 V1 | ... | V(n-2) | V(n-1).
# The first two iterations load only K. The extra two iterations drain V.
kv_token = cutlass.Boolean(True)
for s in cutlass.range(cta_s, prefetch_tiles + s_blks, kv_splits):
if s < s_blks:
# Form tAgK for block s.
for k in cutlass.range_constexpr(KQ_num_k_tiles):
gK_k = tAgK[None, None, None, k]
k_stage_token, k_idx = KV_pipe.producer_acquire_and_get_stage(token=kv_token)
k_mbar = cute_ext.get_mbarrier(k_stage_token)
smem_K_kidx = smem_K[None, None, None, k_idx]
cute_ext.tma_load(gK_k, smem_K_kidx, k_mbar)
KV_pipe.producer_commit_and_advance()
kv_token = KV_pipe.producer_try_acquire()
if s >= prefetch_tiles:
v_block_coord = (0, s - prefetch_tiles, cta_l)
# Form the V block and its per-sk tAgV views.
for sk in cutlass.range_constexpr(VP_num_k_tiles):
for dm in cutlass.range_constexpr(VP_num_m_tiles):
gV_k = tAgV[None, None, None, dm]
v_stage_token, v_idx = KV_pipe.producer_acquire_and_get_stage(token=kv_token)
v_mbar = cute_ext.get_mbarrier(v_stage_token)
smem_V_vidx = smem_V[None, None, None, v_idx]
cute_ext.tma_load(gV_k, smem_V_vidx, v_mbar)
KV_pipe.producer_commit_and_advance()
kv_token = KV_pipe.producer_try_acquire()
# This warp subsequently waits for final M/L and writes split statistics.
elif warp_idx == mma_kq_warp_id: # warp 8
vp_tiles_per_iter = VP_num_m_tiles * VP_num_k_tiles
mma_atom = cute.make_mma_atom(tiled_mma_kq.op)
cute.arch.mbarrier_wait(q_load_mbar_ptr, phase=0)
for s in cutlass.range(cta_s, s_blks, kv_splits):
s_token = S_pipe.producer_try_acquire()
mma_atom.set(tcgen05.Field.ACCUMULATE, False)
# From K3 onward, skip V entries belonging to VP's consumer cursor.
if s >= cta_s + (prefetch_iters + 1) * kv_splits:
for _ in cutlass.range_constexpr(vp_tiles_per_iter):
KV_pipe.consumer_state = KV_pipe.increment_state(KV_pipe.consumer_state)
mma_order_vp_nbar.arrive_and_wait()
k_token = KV_pipe.consumer_try_wait()
_s_stage_token, kq_idx = S_pipe.producer_acquire_and_get_stage(token=s_token)
tmem_S_sliced = tmem_S[None, None, None, kq_idx]
for k_tidx in cutlass.range_constexpr(KQ_num_k_tiles):
_k_stage_token, K_sidx = KV_pipe.consumer_wait_and_get_stage(token=k_token)
for instr_idx in cutlass.range_constexpr(KQ_num_instr_k):
K_slice = smem_K[None, None, instr_idx, K_sidx]
Q_slice = smem_Q[None, None, instr_idx, k_tidx]
cute_ext.dot(mma_atom, K_slice, Q_slice, tmem_S_sliced)
mma_atom.set(tcgen05.Field.ACCUMULATE, True)
if k_tidx == KQ_num_k_tiles - 1:
mma_order_kq_nbar.arrive() # let VP advance through the shared ring
KV_pipe.consumer_release_and_advance()
if k_tidx != KQ_num_k_tiles - 1:
k_token = KV_pipe.consumer_try_wait()
S_pipe.producer_commit_and_advance() # softmax waits for this score tile
# No K remains, but VP still needs two ordering handoffs for its tail.
for _ in cutlass.range_constexpr(prefetch_iters):
mma_order_vp_nbar.arrive_and_wait()
mma_order_kq_nbar.arrive()
elif warp_idx == mma_vp_warp_id: # warp 9
kq_tiles_per_iter = KQ_num_k_tiles
mma_atom = cute.make_mma_atom(tiled_mma_vp.op)
mma_order_vp_nbar.arrive() # seed the first KQ handoff
# Skip K0 and K1 without releasing them; KQ is their actual consumer.
for prefetch_iter in cutlass.range_constexpr(prefetch_iters):
if cta_s + prefetch_iter * kv_splits < s_blks:
for _ in cutlass.range_constexpr(kq_tiles_per_iter):
KV_pipe.consumer_state = KV_pipe.increment_state(KV_pipe.consumer_state)
mma_order_kq_nbar.arrive_and_wait()
mma_order_vp_nbar.arrive()
for s in cutlass.range(cta_s, s_blks, kv_splits):
# First skip K2 and wait for its handoff, then consume V0.
# Later iterations skip K(j+2); the V-only tail has no K to skip.
if s + prefetch_iters * kv_splits < s_blks:
for _ in cutlass.range_constexpr(kq_tiles_per_iter):
KV_pipe.consumer_state = KV_pipe.increment_state(KV_pipe.consumer_state)
mma_order_kq_nbar.arrive_and_wait()
p_token = P_pipe.consumer_try_wait()
v_token = KV_pipe.consumer_try_wait()
o_token = O_pipe.producer_try_acquire()
_p_stage_token, p_idx = P_pipe.consumer_wait_and_get_stage(token=p_token)
_o_stage_token, _idx = O_pipe.producer_acquire_and_get_stage(token=o_token)
for k_tidx in cutlass.range_constexpr(VP_num_k_tiles):
for dm in cutlass.range_constexpr(VP_num_m_tiles):
_v_stage_token, v_sidx = KV_pipe.consumer_wait_and_get_stage(token=v_token)
tmem_O_sliced = tmem_O[None, None, None, dm, 0]
mma_atom.set(tcgen05.Field.ACCUMULATE, True)
for instr_idx in cutlass.range_constexpr(VP_num_instr_k):
V_slice = smem_V[None, None, instr_idx, v_sidx]
P_slice = smem_P_nk[None, None, instr_idx, k_tidx, p_idx]
cute_ext.dot(mma_atom, V_slice, P_slice, tmem_O_sliced)
mma_atom.set(cute.nvgpu.tcgen05.Field.ACCUMULATE, True)
if dm == VP_num_m_tiles - 1 and k_tidx == VP_num_k_tiles - 1:
mma_order_vp_nbar.arrive() # KQ may advance to the next K block
KV_pipe.consumer_release_and_advance()
if dm != VP_num_m_tiles - 1 or k_tidx != VP_num_k_tiles - 1:
v_token = KV_pipe.consumer_try_wait()
P_pipe.consumer_release_and_advance()
O_pipe.producer_commit_and_advance() # correction waits for VP completion
O_pipe.producer_tail()
Softmax supplies P to the VP warp. Correction releases O only after rescaling it, so VP needs both inputs before proceeding. I moved denominator accumulation into softmax after publishing P. Although that regresses some shapes on its own, it removes the full FP32 P store and load handoff through TMEM. I expect the saving to matter more with two softmax warpgroups and multiple token rows per thread.
Comparison with FA4
This compares the current simple kernel above, not the later PR kernel, with FA4 flash_attn_func using automatic splitting. I measured all 12 shapes on B200 in BF16 with one query token, CUDA graphs and at least 512 MiB of rotating KV inputs per path. Both use identical values in their native contiguous layouts. Layout conversion is outside timing, and both timings include the final combine.
Each kernel uses its own automatic split heuristic. The FA4 build includes S/P ping pong; its dispatch enables it for the two batch 32 rows below and leaves it off for the other rows. Speedup is FA4 time divided by the current simple kernel time. The table uses two independent runs, with 14 alternating order samples per run.
| B | Hq/Hkv | d | KV | FA4 µs | Current simple µs | Speedup |
|---|---|---|---|---|---|---|
| 1 | 64/8 | 128 | 1024 | 9.886 | 6.521 | 1.516× |
| 1 | 64/8 | 128 | 8192 | 17.318 | 10.751 | 1.611× |
| 1 | 64/8 | 128 | 32768 | 32.663 | 26.155 | 1.249× |
| 1 | 64/8 | 128 | 131072 | 92.583 | 82.659 | 1.120× |
| 1 | 16/1 | 128 | 8192 | 14.841 | 8.313 | 1.785× |
| 1 | 16/1 | 128 | 32768 | 23.337 | 12.964 | 1.800× |
| 1 | 32/1 | 128 | 8192 | 15.125 | 9.713 | 1.557× |
| 1 | 32/1 | 128 | 32768 | 23.921 | 14.536 | 1.646× |
| 1 | 64/8 | 64 | 8192 | 12.967 | 8.753 | 1.481× |
| 1 | 64/8 | 64 | 32768 | 25.146 | 16.785 | 1.498× |
| 32 | 64/8 | 128 | 1024 | 24.371 | 28.263 | 0.862× |
| 32 | 64/8 | 128 | 8192 | 169.545 | 157.549 | 1.076× |
The batch 32 case with 1024 KV tokens is still slower than FA4.
Overlapping successive KV blocks
I captured the current simple kernel with IKET using one sequence, 64 query heads, eight KV heads, head dimension 128 and 32768 KV tokens. The trace samples one CTA with eight splits, so it contains several iterations of the pipeline.

The softmax tiles run consecutively on one warpgroup, and VP repeatedly waits for P. CUTLASS opt assigns even and odd tiles to two softmax warpgroups so their work can overlap. Updates to the shared running maximum still pass through an ordered critical section.
There is also only one O accumulator in the simple kernel. Correction must wait for the previous VP, rescale O and release it before the next VP can update it. O acquisition is short in this capture, so the trace does not establish it as the dominant stall. Two rolling O accumulators nevertheless remove that single buffer restriction, allowing correction on one slot to overlap VP on the other. A slot still cannot be reused until correction releases it.
I checked the opt kernel reduction path against an FP32 reference in BF16 and FP16. The single excerpt below shows its softmax and correction handoffs. Layout setup, ordinary copies, startup edge cases and the final merge are omitted. Unlike my current simple copy, opt reduces local probability rows in softmax but leaves the running denominator merge in correction.
# Two softmax warpgroups take alternate tiles; shared M stays ordered.
if warpgroup_idx in softmax_warpgroup_ids:
phase = warpgroup_idx - softmax_warpgroup_ids[0]
phase_M_acquire_nbar = with_phase(sM_mutex_nbar, phase)
phase_M_release_nbar = with_phase(sM_mutex_nbar, phase ^ 1)
phase_L_consumer_nbar = with_phase(L_consumer_nbar, phase)
phase_L_producer_nbar = with_phase(L_producer_nbar, phase)
tmem_L_phase = tmem_L[(None, None), 0, 0, phase]
if phase == 1:
S_pipe.consumer_state = S_pipe.increment_state(S_pipe.consumer_state)
P_pipe.producer_state = P_pipe.increment_state(P_pipe.producer_state)
phase_M_release_nbar.arrive() # seed the even group's first turn
for iter_idx in cutlass.range(phase, iters_s, SOFTMAX_WARPGROUPS):
# Wait for S and copy it to registers (copy code omitted).
cute.arch.fence_view_async_tmem_load()
S_pipe.consumer_release_and_advance()
# Mask the tail and reduce scores to lane_max .
phase_M_acquire_nbar.arrive_and_wait()
M_consumer_nbar.arrive_and_wait()
if lane_store_max:
smem_M_ptr = smem_M.iterator + smem_M.layout(lane_idx)
smem_fmax(smem_M_ptr, lane_max)
M_producer_nbar.arrive_and_wait()
cute.autovec_copy(smem_M, rmem_M)
phase_M_release_nbar.arrive() # the other group may update M
_, p_idx = P_pipe.producer_acquire_and_get_stage(token=p_token)
rmem_colmax = rmem_M_valid.load().reshape((g_tile, 1))
probs = exp2(scale_s_log2_e * scores - rmem_colmax)
# Convert P and store it to its SMEM stage .
cute.arch.fence_view_async_shared()
P_pipe.producer_commit_and_advance()
# Send row partial sums, not the full FP32 probability tile.
rmem_L_valid.store(probs.reduce(
cute.ReductionOp.ADD, Float32(0.0), (None, 0)).reshape((g_tile,)))
phase_L_consumer_nbar.arrive_and_wait()
cute_ext.partition_and_copy(thr_r2t_copy, rmem_L, tmem_L_phase)
cute.arch.fence_view_async_tmem_store()
phase_L_producer_nbar.arrive()
# Skip the other group's stage after the normal pipeline advance.
S_pipe.consumer_state = S_pipe.increment_state(S_pipe.consumer_state)
P_pipe.producer_state = P_pipe.increment_state(P_pipe.producer_state)
elif warpgroup_idx == correction_warpgroup_id:
# Startup zeroes O0/O1, seeds L barriers and reads the first two maxima.
phase = 0
for s in cutlass.range(iters_s - self.O_stages, unroll=self.O_stages):
with_phase(L_producer_nbar, phase).arrive_and_wait()
# Read this phase's L into rmem_L, then release its mailbox.
cute.arch.fence_view_async_tmem_load()
with_phase(L_consumer_nbar, phase).arrive()
M_producer_nbar.arrive_and_wait()
# Read smem_M into M_cur; final-block signaling is omitted.
M_consumer_nbar.arrive()
o_token = O_pipe.consumer_try_wait()
_, _o_idx = O_pipe.consumer_wait_and_get_stage(token=o_token)
rmem_corr_lane = exp2(M_prev2 - M_cur)
for gi in cutlass.range_constexpr(g_tile):
rmem_corr[gi] = cute.arch.shuffle_sync(rmem_corr_lane, gi)
# Correct the older slot while VP can work on the other slot.
for dm in cutlass.range_constexpr(tiles_dm):
tmem_O_sliced = tmem_O[(None, None), 0, 0, dm, phase]
# Load this O tile into rmem_O_acc (copy code omitted).
cute.arch.fence_view_async_tmem_load()
rmem_O_acc.store(rmem_O_acc.load() * rmem_corr.load())
cute_ext.partition_and_copy(thr_r2t_copy, rmem_O_acc, tmem_O_sliced)
cute.arch.fence_view_async_tmem_store()
O_pipe.consumer_release_and_advance() # VP may reuse this slot
# Rescale and accumulate this phase's L history .
M_prev2, M_prev = M_prev, M_cur
phase ^= 1
The two O slots carry different maximum histories. The final merge rescales the older slot into the newer slot's maximum before adding them.
I captured gqa_decode_opt with kernel reduction and PDL disabled on both launches. The image below uses 32 query heads sharing one KV head, head dimension 64, 32768 KV tokens and eight splits. This wider group makes both overlaps visible; it is not a shape matched timing comparison with the earlier simple trace.

SM_math(payload=8) on Warp04 overlaps SM_local_max(payload=9) on Warp08. The even group is computing and publishing probabilities while the odd group reduces the next tile's local maximum. Their shared maximum updates remain ordered, but the surrounding work no longer has to run on one warpgroup.
Below that, O_rescale(payload=0) on Warp12 overlaps VP_use_O(payload=1) on Warp01, including a VP_issue(payload=1) interval. Correction rescales one output slot while VP issues work into the other. This avoids making those operations take turns on the same O storage.
The combine kernel still runs after decode. PDL has not been introduced here. These are instrumented warp activity ranges, not tensor core utilization or an isolated speedup measurement.
Combining the KV splits
Splitting the KV sequence gives a small decode batch enough CTAs to use the GPU. Each split produces an unnormalized FP32 output, a running maximum and a denominator. Those outputs cannot simply be added because each split used its own maximum when computing softmax.
For one query head, let split return . The merge brings every partial into the same scale before normalizing. The maxima here use the kernel's base two score scaling.
With kernel reduction, decode writes these partials to global memory and a separate combine kernel merges them. The combine uses FP32 arithmetic and a fixed split traversal rather than concurrent output additions. Without PDL, the stream waits for decode to finish before starting combine. For a short decode, that second launch is a noticeable part of the operation.
Starting combine early with PDL
Programmatic dependent launch lets a dependent kernel start before its producer finishes. In this kernel, the split statistics warp calls griddepcontrol_launch_dependents() after writing its statistics. The output store and other producer work may still be draining.
Once the producer CTAs have reached their launch trigger, combine becomes eligible to start its independent prologue. It then executes griddepcontrol_wait() before loading any split maxima, sums or output partials. That wait protects the dependency until the producer has completed and its writes are visible. PDL overlaps startup with the producer tail, not the weighted sum with unfinished inputs. Actual overlap still depends on scheduling and available resources.
Both launches opt into PDL. Decode also waits before reading its inputs so it can safely follow a PDL enabled producer. The reduction formula has not changed.
Removing the max reset
The original CUTLASS kernel reduction maintains an additional global M_final buffer. Every split atomically updates it with its local maximum, so the buffer must be reset to negative infinity before every invocation. Replaying a graph without that reset could leave a maximum from a previous input.
I removed that buffer and its global atomic update. Each split already writes its own maximum, and combine already loads those maxima to rescale the partials. I take their maximum in combine instead. This trades a small reduction over staged values for an extra reset launch and inter CTA atomic updates.
The output and partial workspaces can now come from torch.empty. Every value combine consumes is written by the current invocation. Empty splits are skipped, and an entirely empty sequence produces zero output. No memset here means no global workspace or output reset in kernel mode; the register accumulators still start at zero.
This is the combine path after that change. Tensor view construction, register accumulator initialization and the optional LSE store are omitted.
if cutlass.const_expr(self.use_pdl):
cute.arch.fence_acq_rel_cta()
cute.arch.griddepcontrol_wait() # Do not read unfinished decode outputs.
# Stage each split's denominator and maximum after the dependency wait.
if tidx < num_valid_splits:
cute_ext.simt_auto_vec_copy(gSum_partial[None, tidx], sSum_partial[None, tidx], async_op=True)
cute_ext.simt_auto_vec_copy(gMax_partial[None, tidx], sMax_partial[None, tidx], async_op=True)
for split_idx in cutlass.range(threads_per_cta + tidx, num_valid_splits, threads_per_cta):
cute_ext.simt_auto_vec_copy(gSum_partial[None, split_idx], sSum_partial[None, split_idx], async_op=True)
cute_ext.simt_auto_vec_copy(gMax_partial[None, split_idx], sMax_partial[None, split_idx], async_op=True)
cute.arch.cp_async_commit_group()
cute.arch.cp_async_wait_group(0)
cute.arch.sync_threads()
# Replace the reset global M_final buffer with a reduction over staged maxima.
row_max = -Float32.inf
for split_idx in cutlass.range(num_valid_splits, unroll=8):
row_max = cute.arch.fmax(row_max, sMax_partial[0, split_idx])
row_sum = Float32(0)
if row_max > -Float32.inf and hdim_in_bounds:
for split_idx in cutlass.range(num_valid_splits, unroll=8):
row_max_split = sMax_partial[0, split_idx]
if row_max_split > -Float32.inf: # Ignore empty splits.
acc_scale_split = exp2_fast(row_max_split - row_max)
row_sum += acc_scale_split * sSum_partial[0, split_idx]
tOrO_final += acc_scale_split * tOgO_partial[None, split_idx].load()
tOrO_final *= cute.arch.rcp_approx(row_sum)
if cutlass.const_expr(self.use_pdl):
cute.arch.griddepcontrol_launch_dependents()
if hdim_in_bounds:
tOgO.store(tOrO_final.to(mO.element_type)) # Overwrite O, rather than add to it.
Reducing inside a thread block cluster
Atomic reduction takes a different route. The KV splits for one batch item and head group form a thread block cluster. They exchange their small maximum and denominator arrays through distributed shared memory. A butterfly first finds the common maximum, then sums the denominators after rescaling them into that maximum's scale.
Each split now knows its final normalization factor. It scales its own output partial and uses TMA reduce add to accumulate it directly into the final output. The full output partials do not need a separate global workspace and combine kernel.
This is not a free replacement for kernel reduction. The output must be zeroed on every call, including every graph replay. The BF16 or FP16 output additions are order dependent and can round differently from the FP32 combine. This implementation also requires a power of two split count, capped at 16, and the cluster's CTAs must be scheduled together.
I keep both modes. Kernel reduction supports more splits and deterministic FP32 combination without a reset. Atomic reduction can avoid the combine launch, particularly when fewer splits are needed, but pays for output zeroing and cluster synchronization. Automatic selection chooses atomic reduction for at most four splits, or eight splits with at most 64 CTAs, and kernel reduction otherwise. PDL defaults to on for kernel reduction and off for atomic reduction; standalone atomic mode has no separate combine kernel whose startup it can hide. Atomic mode still accepts an explicit PDL override for a larger dependent pipeline.
The final split reduction trace
This capture uses kernel reduction with PDL on and no memset. It has one query token, batch 1, 32 query heads sharing one KV head, 32768 KV tokens and eight splits. I used head dimension 256 to make the reduction easier to inspect; the benchmark matrix below uses dimensions 64 and 128. All decode and combine CTAs are recorded.

producer_trigger occurs while the output partial tail is still active. The combine prologue starts before the producer's recorded warp lifetime ends, then griddepcontrol_wait blocks. Only after it returns do the split statistic loads begin, followed by split_max_fold and weighted_output_merge.
This shows early launch and dependency waiting, not simultaneous merge arithmetic and unfinished output production. In this capture, the prologue does not overlap the measured output tail or store issue scopes. The store issue range also does not measure asynchronous completion.
split_max_fold is the small reduction that replaces the global maximum buffer and its reset. A separate capture of the complete wrapper contains only decode and combine, with no memset or reset kernel. The IKET ranges themselves are not benchmark timings.
Performance against FA4
I reran the PR's 60 shape matrix against upstream FA4 at 47e91f1, using the final implementation from PR #2973. These are B200 BF16 single token decode measurements, not full model serving results. FA4 uses its automatic split and S/P ping pong choices; my modes use their default split heuristics rather than a per shape search.
The timer uses CUDA graphs with complete input pool cycles covering at least 512 MiB of KV data. Each latency is the fastest of three interleaved round medians, with seven graph replays per round. The timed operation includes decode and combine where applicable, and output zeroing for atomic mode. Both paths use the same input tensors and timer.
The two plots use the same layout and mode curves as the PR, refreshed with this run. Horizontal positions are equally spaced samples, not a linear scale. Bandwidth is effective problem bandwidth, counting Q, K, V and O once and dividing by the complete operation's latency. It is not a hardware counter measurement of HBM traffic. The axes use decimal TB/s and the tables use decimal GB/s.


Kernel reduction with PDL and no memset has a median per shape speedup of 1.46×, ranging from 0.88× to 2.26×. Automatic selection reaches 1.47× at the median. The regressions remain in the plots and tables; neither reduction mode wins everywhere.
In the paired PDL comparisons, kernel mode saves 0.82 µs at the median, while enabling PDL in standalone atomic mode adds 0.53 µs. Both kernel variants already have the no memset change, so this comparison does not measure the benefit of removing the reset.
The main tables compare FA4 with kernel reduction, PDL on, no memset. Speedup is FA4 latency divided by kernel latency, calculated before display rounding.
MQA
Batch 1 and Hq/Hkv 16/1.
| d | KV tokens | FA4 µs | Kernel µs | FA4 GB/s | Kernel GB/s | Speedup |
|---|---|---|---|---|---|---|
| 64 | 1024 | 7.78 | 6.27 | 34.2 | 42.5 | 1.24× |
| 64 | 4096 | 10.12 | 6.88 | 104.0 | 153.0 | 1.47× |
| 64 | 16384 | 20.21 | 9.22 | 207.7 | 455.1 | 2.19× |
| 64 | 32768 | 21.82 | 9.78 | 384.5 | 858.3 | 2.23× |
| 64 | 65536 | 24.43 | 12.63 | 687.0 | 1328.4 | 1.93× |
| 64 | 131072 | 28.78 | 17.37 | 1166.0 | 1931.6 | 1.66× |
| 128 | 1024 | 9.97 | 6.36 | 53.4 | 83.7 | 1.57× |
| 128 | 4096 | 12.57 | 7.27 | 167.4 | 289.7 | 1.73× |
| 128 | 16384 | 21.63 | 9.57 | 388.3 | 877.6 | 2.26× |
| 128 | 32768 | 23.84 | 10.79 | 704.1 | 1555.9 | 2.21× |
| 128 | 65536 | 27.03 | 14.71 | 1241.7 | 2281.4 | 1.84× |
| 128 | 131072 | 32.40 | 21.57 | 2071.6 | 3111.8 | 1.50× |
GQA
Batch 1 and Hq/Hkv 16/2.
| d | KV tokens | FA4 µs | Kernel µs | FA4 GB/s | Kernel GB/s | Speedup |
|---|---|---|---|---|---|---|
| 64 | 1024 | 8.07 | 6.21 | 65.4 | 85.1 | 1.30× |
| 64 | 4096 | 10.42 | 6.77 | 201.7 | 310.3 | 1.54× |
| 64 | 16384 | 16.25 | 9.32 | 516.4 | 900.4 | 1.74× |
| 64 | 32768 | 18.35 | 10.35 | 914.7 | 1621.1 | 1.77× |
| 64 | 65536 | 23.08 | 13.14 | 1453.7 | 2554.6 | 1.76× |
| 64 | 131072 | 30.25 | 18.25 | 2218.8 | 3677.5 | 1.66× |
| 128 | 1024 | 10.07 | 6.50 | 104.9 | 162.6 | 1.55× |
| 128 | 4096 | 12.73 | 6.95 | 330.2 | 604.8 | 1.83× |
| 128 | 16384 | 18.94 | 10.12 | 886.0 | 1658.9 | 1.87× |
| 128 | 32768 | 21.96 | 12.19 | 1528.6 | 2754.2 | 1.80× |
| 128 | 65536 | 28.20 | 17.34 | 2380.4 | 3870.8 | 1.63× |
| 128 | 131072 | 37.47 | 26.69 | 3582.6 | 5029.8 | 1.40× |
MHA
Batch 1 and Hq/Hkv 16/16.
| d | KV tokens | FA4 µs | Kernel µs | FA4 GB/s | Kernel GB/s | Speedup |
|---|---|---|---|---|---|---|
| 64 | 1024 | 8.46 | 6.07 | 496.4 | 691.4 | 1.39× |
| 64 | 4096 | 12.31 | 7.86 | 1363.4 | 2136.0 | 1.57× |
| 64 | 16384 | 23.72 | 14.64 | 2829.1 | 4584.7 | 1.62× |
| 64 | 32768 | 34.76 | 23.94 | 3861.4 | 5606.1 | 1.45× |
| 64 | 65536 | 58.93 | 42.17 | 4555.3 | 6366.3 | 1.40× |
| 64 | 131072 | 114.07 | 78.67 | 4706.7 | 6824.7 | 1.45× |
| 128 | 1024 | 10.11 | 6.45 | 830.6 | 1302.2 | 1.57× |
| 128 | 4096 | 15.13 | 10.55 | 2218.5 | 3182.4 | 1.43× |
| 128 | 16384 | 29.01 | 24.25 | 4627.0 | 5534.3 | 1.20× |
| 128 | 32768 | 48.66 | 42.44 | 5516.8 | 6326.0 | 1.15× |
| 128 | 65536 | 93.82 | 78.89 | 5722.4 | 6805.7 | 1.19× |
| 128 | 131072 | 175.01 | 151.84 | 6135.4 | 7071.8 | 1.15× |
Batch scaling
Head dimension 128.
| B | Hq/Hkv | KV tokens | FA4 µs | Kernel µs | FA4 GB/s | Kernel GB/s | Speedup |
|---|---|---|---|---|---|---|---|
| 1 | 32/8 | 1024 | 10.55 | 6.16 | 399.2 | 683.6 | 1.71× |
| 1 | 32/8 | 8192 | 17.63 | 10.11 | 1904.3 | 3321.4 | 1.74× |
| 1 | 32/8 | 32768 | 32.89 | 24.00 | 4081.8 | 5593.6 | 1.37× |
| 8 | 32/8 | 1024 | 15.17 | 9.52 | 2221.1 | 3540.2 | 1.59× |
| 8 | 32/8 | 8192 | 51.36 | 42.32 | 5229.1 | 6346.0 | 1.21× |
| 8 | 32/8 | 32768 | 186.11 | 152.64 | 5770.1 | 7035.1 | 1.22× |
| 32 | 32/8 | 1024 | 24.18 | 26.36 | 5571.6 | 5111.4 | 0.92× |
| 32 | 32/8 | 8192 | 183.06 | 155.18 | 5868.3 | 6922.7 | 1.18× |
| 32 | 32/8 | 32768 | 721.06 | 596.34 | 5957.2 | 7203.1 | 1.21× |
| 64 | 32/8 | 1024 | 45.51 | 47.88 | 5921.4 | 5628.1 | 0.95× |
| 64 | 32/8 | 8192 | 360.41 | 304.27 | 5961.4 | 7061.3 | 1.18× |
| 64 | 32/8 | 32768 | 1445.19 | 1186.28 | 5944.6 | 7241.9 | 1.22× |
| 1 | 64/8 | 1024 | 10.60 | 6.28 | 398.8 | 673.0 | 1.69× |
| 1 | 64/8 | 8192 | 17.87 | 10.01 | 1879.0 | 3355.4 | 1.79× |
| 1 | 64/8 | 32768 | 33.26 | 24.23 | 4036.0 | 5541.8 | 1.37× |
| 8 | 64/8 | 1024 | 15.06 | 10.04 | 2246.1 | 3368.4 | 1.50× |
| 8 | 64/8 | 8192 | 51.40 | 42.56 | 5227.4 | 6313.5 | 1.21× |
| 8 | 64/8 | 32768 | 175.15 | 152.52 | 6131.9 | 7041.6 | 1.15× |
| 32 | 64/8 | 1024 | 24.91 | 28.21 | 5431.0 | 4795.6 | 0.88× |
| 32 | 64/8 | 8192 | 174.17 | 158.51 | 6171.0 | 6780.7 | 1.10× |
| 32 | 64/8 | 32768 | 706.14 | 602.04 | 6083.8 | 7135.7 | 1.17× |
| 64 | 64/8 | 1024 | 46.02 | 51.87 | 5878.4 | 5215.6 | 0.89× |
| 64 | 64/8 | 8192 | 353.46 | 307.37 | 6081.5 | 6993.5 | 1.15× |
| 64 | 64/8 | 32768 | 1413.95 | 1186.62 | 6076.6 | 7240.7 | 1.19× |
Conclusion
I kept CUTLASS's register redistribution for the wider head groups. For a 32 head tile, the MMA and TMA warps use setmaxregister_decrease with a limit of 64 registers per thread, the softmax warpgroups decrease to 120, and correction uses setmaxregister_increase to request 208. Correction holds output fragments and denominator state, so it gets more of the CTA's register budget instead of reserving the same amount for every role. This path is active in the 32 head trace configurations, but not in the smaller head groups of the benchmark matrix above.
One change I tried but did not keep was removing the mutex between the softmax warpgroups. I gave the even and odd groups separate SMEM maximum slots so each could maintain its own running maximum without taking turns on a shared slot. Correction still needed the per group handoffs, and the final merge had to rescale both maximum histories into a common scale before combining their outputs and denominators. The earlier A/B tests showed no clear performance gain, with results mostly flat or slightly slower.
The profiles help explain that result. In the wider head group trace above, the recorded mutex acquisition ranges have a median of 32 ns and a maximum of 96 ns. Waiting for score tiles is more visible. The mutex was not a large observed stall in these captures, so removing it offered little to recover while adding separate state and extra final rescaling. I kept the shared maximum protocol rather than adding complexity without a measured gain.
Acknowledgements
Thanks to Colfax Research for Optimization Diaries on S/P Ping Pong for FlashAttention 4 Decode, and to @hassan-abdallah and the FlashAttention contributors for the FA4 decode scheduling optimization in PR #2817. Their explanation and implementation of overlapping QK with softmax helped shape the scheduling discussion in this post.
I also thank the NVIDIA CUTLASS contributors for gqa_decode_simple and gqa_decode_opt. Both the simple kernel and the optimized kernel supplied the swap-AB design and warp specialized pipeline that my FA4 implementation is adapted from.