[Perf][foundry][Attention] Serve MHA backward on a warp-specialized SM90 kernel and time FA3's backward alone - #2281
Merged
Conversation
lcy-seso
force-pushed
the
perf/mha-bwd-ws
branch
from
September 27, 2026 23:12
1925d99 to
0a8319c
Compare
…M90 kernel MHA backward with head dim 128 and 128-row key blocks now runs on a hand-written warp-specialized kernel: one producer warp streams Q, dO, LSE and delta through a two-stage TMA ring, two consumer warpgroups keep dK and dV in registers, and a second producer warp adds each dQ tile into an f32 accumulator with a TMA reduction. The accumulator is laid out in WGMMA register order so the consumers write it with 128-bit stores; a second launch lands it in q's layout and dtype. The general pipelined kernel read S before waiting on its WGMMA, which made ptxas serialize every WGMMA in the kernel, and released q_frag's stage while dK could still read it. Both are fixed, and with one query head per KV head it writes dK/dV in the input dtype instead of zeroing, adding and casting f32 buffers. The FA3 and torch-sdpa backward baselines in bench_mha and bench_gqa ran the forward inside the timed call; they now time the backward alone.
…il dK and dV finish The pipelined backward kernel's pipeline hands a stage back to the producer after its buffer's last access, and the last accesses of q_frag and do_shared are the dK and dV WGMMA issues, which do not wait for the reads. The producer could overwrite a stage those WGMMAs were still reading, which made gqa-bwd's compiled and eager results disagree in CI. Reading both buffers after the step's final wait moves the release past the WGMMAs and replaces the wait that only narrowed the window.
… its per-step path The warp-specialized kernel now runs one CTA per SM that claims key-block tiles in launch order through an atomic counter, so the next tile's K and V load while the current tile's dK and dV are written out. The counter lives in the kernel instance and the last CTA of every launch zeroes it. Inside a step: each consumer warpgroup's half of a dQ tile is stored conflict-free and added into global memory by its own writer warp with one 16 KB bulk reduction; the dS hand-off between the warpgroups is split into an arrive and a later wait; Q and dO move on separate barriers so dO's stage is released after dV, which removes the slowdown at 64 heads; LSE and delta are read into registers before the scores are needed; half the exponentials run while dP is in flight; and dS lives in one 3D buffer, which drops the per-slot copies of the loop body and the instruction-cache misses they caused. The dQ post pass reads the accumulator in storage order, and the preprocess zeroes it in head-major order so each block clears one contiguous span.
lcy-seso
force-pushed
the
perf/mha-bwd-ws
branch
from
September 28, 2026 06:01
86770fb to
5b65f98
Compare
…ckward through GroupedQueryAttentionBwdOp MHA backward is GQA backward with one KV head per query head: the MHA signature is the GQA signature with H_kv = H, and MultiHeadAttentionBwdOp only delegated to GroupedQueryAttentionBwdOp, whose dispatch already picks the warp-specialized kernel when heads_kv == heads. The op class, its workload and benchmark adapters and its tests duplicated the GQA ones line for line. The four MHA manifest rows move to GroupedQueryAttentionBwdOp with H_kv = H and a -mha label, so the benchmark still runs every shape; the five MHA backward test cases move into the GQA backward fixture unchanged, and the compile-boundary case builds the GQA op.
lcy-seso
added a commit
to tile-ai/TileOPs.github.io
that referenced
this pull request
Sep 28, 2026
- `docs/api/attention.md` drops the `MultiHeadAttentionBwdOp` entry; `GroupedQueryAttentionBwdOp`, already on the page, now serves MHA backward. - tile-ai/TileOPs#2281 removed the op, so the daily render job's `check_api_pages.py` fails and the deploy stops until the page drops it. - Against TileOPs `main` at `2d600995c`, `check_api_pages.py` reports 176 ops, all on a page and exported.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
MHABwdWsKernel: 2.08x to 4.21x faster thanmain, 1.01x to 1.02x of FA3's backward on the short rows and 0.98x on the long rows. Every other call keeps the pipelined kernel.main.bench_mha.pyandbench_gqa.pytimed forward plus backward; they now time the backward alone.MultiHeadAttentionBwdOpis removed: its signature is GQA's withH_kv = Hand it only delegated toGroupedQueryAttentionBwdOp, which now serves MHA calls directly. Its four manifest rows move to the GQA entry withH_kv = Hand a-mhalabel.tests/ops/test_gqa.py41 -> 46 (+5) andtests/ops/test_mha.py16 -> 12 (-4): the MHA backward cases move into the GQA backward fixture, plus one causal smoke case for the new kernel.TileFoundry Description
The program has three modules over one causal 2x2048x32x128 fp16 call:
Referenceis the unsplit f32 backwardcheckholds the other two to;Baseis the placement before this change;Tiledis the placement this PR implements.checkagainst torch autograd (1x256x2x128,allcloseatol 5e-3 rtol 1e-2 on all three outputs) passes for all three.The shipped op is held to the
Tiledprogram through a runtime twin.delta_rowsruns the preprocess kernel andmha_bwdrunsGroupedQueryAttentionBwdOp;dq_landis the scale and cast in torch, because the shipped landing pass reads the accumulator in WGMMA register order, which that function's argument does not have.tilefoundry check twin.py:Fastpasses against the evaluator onTiled.mha_bwdand against torch autograd (1x256x2x128, same predicates,max_violation0 on all three outputs).TileFoundry decided the placement below and gives the floor the kernel is read against. The kernel-level work that followed (the WGMMA waits, the warp-specialized schedule, the dQ store order, the split Q and dO rings, the persistent tile loop) sits below what HIR expresses, which has no warpgroup or pipeline stage, and was found in TileLang with ncu stall sampling, ablations and globaltimer timestamps.
tilefoundry analyze --compute-cost --memory --roofline, H200 target:What the analysis decided before any kernel existed:
BM/BNsweep ofTiled(64/64, 128/64, 128/128, 64/128) left gmem traffic unchanged and movedideal-nsby at most 4%; 128/64 is the largest tile whose dK, dV, S and dP accumulators fit a consumer warpgroup's registers, which is the tile FA3 uses for head dim 128.ideal-nsabove.Performance
Operator:
GroupedQueryAttentionBwdOp, the-mharows (H_kv = H) and the GQA rows (H_kv = 8)Method:
pytest benchmarks/ops/bench_mha.py benchmarks/ops/bench_gqa.py -k bwdin the image above, on one idle GPU, first on this branch (5b65f98) and then on its basemain(94bfedb) with this branch's two benchmark files, back to back. The harness flushes L2 before every iteration, times each tag forward then reversed, and reports the median device-busy time, the union of a call's kernel intervals; the p10-p90 spread of every row on this branch is within 2.5%. The candidate is the public op constructed with manifest arguments, with dispatch choosing the kernel. The MHA table was measured throughMultiHeadAttentionBwdOpbefore its removal; it delegated toGroupedQueryAttentionBwdOp, so the calls and kernels are the ones the-mharows run now. The FA3 column is FA3's backward alone on both runs' inputs; the nightly's FA3 row, which included FA3's forward, is not used.Ratio in comparator columns: implementation / candidate. 🟢 > 1 means the candidate is faster; 🔴 <= 1 means it is not.
/ candidate
/ candidate
🟢 2.171x
🟢 1.014x
🟢 4.209x
🟢 1.018x
🟢 2.081x
🔴 0.985x
🟢 3.159x
🔴 0.983x
🟢 2.212x
🟢 1.015x
🟢 4.188x
🟢 1.008x
🟢 2.078x
🔴 0.985x
🟢 2.644x
🔴 0.984x
🟢 2.727x
🔴 0.999x
GQA rows (
H_kv = 8), same run, served by the pipelined kernel:/ candidate
/ candidate
🟢 1.353x
🔴 0.936x
🟢 2.855x
🔴 0.933x
🟢 1.438x
🔴 0.831x
🟢 2.256x
🔴 0.823x
🟢 1.375x
🔴 0.946x
🟢 2.991x
🔴 0.940x
🟢 1.405x
🔴 0.813x
🟢 1.831x
🔴 0.807x
🟢 1.842x
🔴 0.877x
Result And Limitations
not SOTA: faster than FA3's backward on the four short MHA rows (1.008x to 1.018x), 1.5% to 1.7% slower on the four long rows
mainevery row is faster, 2.08x to 4.21x for MHA and 1.35x to 2.99x for GQA; bf16 gains most because main's preprocess took 449 us in bf16 against 57 us in fp16.T.view(TileLang applies the fp16 view's swizzle to the whole buffer and corrupts dQ), dQ added from registers with vector atomics (4% slower on the long rows), counting landed dQ reductions to fuse the landing pass (a release fence per step costs about 5 us), a polynomial exp2 on the FMA pipe, dS transposed into shared memory, and issuing the next step's S before this step's dK drains (spills at the 240-register consumer budget).