Skip to content

[Perf][foundry][Attention] Serve MHA backward on a warp-specialized SM90 kernel and time FA3's backward alone - #2281

Merged
lcy-seso merged 4 commits into
tile-ai:mainfrom
lcy-seso:perf/mha-bwd-ws
Sep 28, 2026
Merged

lcy-seso merged 4 commits into
tile-ai:mainfrom
lcy-seso:perf/mha-bwd-ws

Conversation

@lcy-seso

@lcy-seso lcy-seso commented Sep 27, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

  • MHA backward ran eight launches per call; it now runs three: a preprocess, the backward kernel and a dQ landing pass.
  • MHA backward with head dim 128 and a sequence that splits into 128-row key blocks now runs on a new persistent warp-specialized SM90 kernel, MHABwdWsKernel: 2.08x to 4.21x faster than main, 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.
  • The pipelined kernel serialized all its WGMMAs by reading an accumulator before its wait, and released its Q/dO stages before dK/dV had read them; both are fixed, and with the new preprocess GQA backward is 1.35x to 2.99x faster than main.
  • The FA3 and torch-sdpa baselines in bench_mha.py and bench_gqa.py timed forward plus backward; they now time the backward alone.
  • MultiHeadAttentionBwdOp is removed: its signature is GQA's with H_kv = H and it only delegated to GroupedQueryAttentionBwdOp, which now serves MHA calls directly. Its four manifest rows move to the GQA entry with H_kv = H and a -mha label.
  • Test nodes: tests/ops/test_gqa.py 41 -> 46 (+5) and tests/ops/test_mha.py 16 -> 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: Reference is the unsplit f32 backward check holds the other two to; Base is the placement before this change; Tiled is the placement this PR implements. check against torch autograd (1x256x2x128, allclose atol 5e-3 rtol 1e-2 on all three outputs) passes for all three.

#!/usr/bin/env python3
"""Causal MHA backward, BSHD, as HIR."""
import math
import os

from tilefoundry import func, module
from tilefoundry.dsl import Mesh, Tensor, Topology, tf
from tilefoundry.dsl.tf import *  # noqa: F401, F403
from tilefoundry.target import CudaTarget

B, S, H, D = (int(x) for x in os.environ.get("MHA_SHAPE", "2,2048,32,128").split(","))
DT = os.environ.get("MHA_DT", "f16")
SM_SCALE = 1.0 / math.sqrt(D)
LOG2E = 1.4426950408889634

_H200 = CudaTarget("nvidia.h200_sxm")


@module(entry="mha_bwd", target=_H200, topologies=(Topology("cta", 132),))
class Reference:
    """The whole backward in f32, unsplit: the program check holds every placement to."""

    @func
    def mha_bwd(
        q: Tensor[(B, S, H, D), DT],
        k: Tensor[(B, S, H, D), DT],
        v: Tensor[(B, S, H, D), DT],
        o: Tensor[(B, S, H, D), DT],
        do: Tensor[(B, S, H, D), DT],
        lse: Tensor[(B, H, S), "f32"],
    ):
        q32 = tf.transpose(tf.cast(q, dtype="f32"), perm=(0, 2, 1, 3))
        k32 = tf.transpose(tf.cast(k, dtype="f32"), perm=(0, 2, 1, 3))
        v32 = tf.transpose(tf.cast(v, dtype="f32"), perm=(0, 2, 1, 3))
        o32 = tf.transpose(tf.cast(o, dtype="f32"), perm=(0, 2, 1, 3))
        do32 = tf.transpose(tf.cast(do, dtype="f32"), perm=(0, 2, 1, 3))
        delta = tf.reduce(o32 * do32, axes=(-1,), keepdim=True, kind="sum")
        scores = tf.matmul(q32, k32, b_layout="NK")
        logits = scores * tf.full_like(scores, value=SM_SCALE * LOG2E)
        prob = tf.exp2(logits - tf.reshape(lse, new_shape=(B, H, S, 1)))
        rows = tf.reshape(tf.arange(type=Tensor[(S,), "i32"]), new_shape=(S, 1))
        cols = tf.reshape(tf.arange(type=Tensor[(S,), "i32"]), new_shape=(1, S))
        prob = tf.where(tf.cmp_le(cols, rows), prob, tf.full_like(prob, value=0.0))
        dv = tf.matmul(prob, do32, a_layout="KM")
        dprob = tf.matmul(do32, v32, b_layout="NK")
        dscore = prob * (dprob - delta) * tf.full_like(prob, value=SM_SCALE)
        dq = tf.matmul(dscore, k32)
        dk = tf.matmul(dscore, q32, a_layout="KM")
        return (
            tf.cast(tf.transpose(dq, perm=(0, 2, 1, 3)), dtype=DT),
            tf.cast(tf.transpose(dk, perm=(0, 2, 1, 3)), dtype=DT),
            tf.cast(tf.transpose(dv, perm=(0, 2, 1, 3)), dtype=DT),
        )


BM = int(os.environ.get("MHA_BM", "128"))  # key/value rows one CTA owns
BN = int(os.environ.get("MHA_BN", "64"))  # query rows per loop step
NB = S // BM


@module(entry="mha_bwd", target=_H200, topologies=(Topology("cta", 132),))
class Base:
    """The placement before this change: dK and dV leave the CTA as f32 adds into zeroed
    gmem buffers, and two more launches land them in the input dtype."""

    @func
    def delta_rows(o: Tensor[(B, S, H, D), DT], do: Tensor[(B, S, H, D), DT]):
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            o_r = tf.reshard(o, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            do_r = tf.reshard(do, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            delta = tf.reduce(tf.cast(o_r, dtype="f32") * tf.cast(do_r, dtype="f32"), axes=(-1,), keepdim=True, kind="sum")
            return tf.reshard(delta, (B @ cta.b, S @ cta.n, H @ cta.h, 1), "gmem")

    @func
    def dq_land(dq_acc: Tensor[(B, S, H, D), "f32"]):
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            dq_r = tf.reshard(dq_acc, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            return tf.reshard(tf.cast(dq_r, dtype=DT), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")

    @func
    def land(acc: Tensor[(B, S, H, D), "f32"]):
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            acc_r = tf.reshard(acc, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            return tf.reshard(tf.cast(acc_r, dtype=DT), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")

    @func
    def mha_bwd(
        q: Tensor[(B, S, H, D), DT],
        k: Tensor[(B, S, H, D), DT],
        v: Tensor[(B, S, H, D), DT],
        o: Tensor[(B, S, H, D), DT],
        do: Tensor[(B, S, H, D), DT],
        lse: Tensor[(B, H, S), "f32"],
    ):
        delta = tf.reshard(delta_rows(o, do), (B, S, H, 1), "gmem")  # noqa: F821
        lse_rows = tf.reshape(lse, new_shape=(B, H, 1, S))
        delta_rows_t = tf.reshape(tf.transpose(delta, perm=(0, 2, 1, 3)), new_shape=(B, H, 1, S))
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            k_s = tf.reshard(k, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
            v_s = tf.reshard(v, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
            kv_pos = tf.reshard(
                tf.reshape(tf.arange(type=Tensor[(S,), "i32"]), new_shape=(1, 1, S, 1)),
                (1, 1, S @ cta.n, 1),
                "rmem",
            )
            q_pos_all = tf.reshape(tf.arange(type=Tensor[(S,), "i32"]), new_shape=(1, 1, 1, S))
            acc0 = tf.zeros(Tensor[(B @ cta.b, H @ cta.h, S @ cta.n, D), "f32", "rmem"])
            dk = tf.full_like(acc0, value=0.0)
            dv = tf.full_like(acc0, value=0.0)
            dq0 = tf.zeros(Tensor[(B @ cta.b, S, H @ cta.h, D), "f32", "gmem"])
            dq_acc = tf.full_like(dq0, value=0.0)
            for t in tile(S, BN):
                t0 = t + 0
                # K and V are loaded once per CTA; carrying them keeps that load outside the loop.
                k_s = tf.reshard(k_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
                v_s = tf.reshard(v_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
                q_s = tf.reshard(q[:, t0 : t0 + BN, :, :], (B @ cta.b, BN, H @ cta.h, D), "smem")
                do_s = tf.reshard(do[:, t0 : t0 + BN, :, :], (B @ cta.b, BN, H @ cta.h, D), "smem")
                lse_r = tf.reshard(lse_rows[:, :, :, t0 : t0 + BN], (B @ cta.b, H @ cta.h, 1, BN), "rmem")
                delta_r = tf.reshard(delta_rows_t[:, :, :, t0 : t0 + BN], (B @ cta.b, H @ cta.h, 1, BN), "rmem")
                q_pos = tf.reshard(q_pos_all[:, :, :, t0 : t0 + BN], (1, 1, 1, BN), "rmem")
                # WGMMA reads its shared-memory operands; these reshards are those reads.
                k_r = tf.transpose(tf.reshard(k_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                v_r = tf.transpose(tf.reshard(v_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                q_r = tf.transpose(tf.reshard(q_s, (B @ cta.b, BN, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                do_r = tf.transpose(tf.reshard(do_s, (B @ cta.b, BN, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                # Transposed scores: rows are keys, columns are queries.
                st = tf.cast(tf.matmul(k_r, q_r, b_layout="NK"), dtype="f32")
                pt = tf.exp2(st * tf.full_like(st, value=SM_SCALE * LOG2E) - lse_r)
                pt = tf.where(tf.cmp_le(kv_pos, q_pos), pt, tf.full_like(pt, value=0.0))
                p16 = tf.cast(pt, dtype=DT)
                dv = dv + tf.cast(tf.matmul(p16, do_r), dtype="f32")
                dpt = tf.cast(tf.matmul(v_r, do_r, b_layout="NK"), dtype="f32")
                ds16 = tf.cast(pt * (dpt - delta_r) * tf.full_like(pt, value=SM_SCALE), dtype=DT)
                dk = dk + tf.cast(tf.matmul(ds16, q_r), dtype="f32")
                # dQ contracts over the key rows the mesh split: a partial over cta.n,
                # read back from smem as the WGMMA A operand and settled by one add in gmem.
                ds_s = tf.reshard(ds16, (B @ cta.b, H @ cta.h, S @ cta.n, BN), "smem")
                ds_r = tf.reshard(ds_s, (B @ cta.b, H @ cta.h, S @ cta.n, BN), "rmem")
                dq_step = tf.transpose(tf.cast(tf.matmul(ds_r, k_r, a_layout="KM"), dtype="f32"), perm=(0, 2, 1, 3))
                dq_tile = tf.reshard(dq_step, (B @ cta.b, BN, H @ cta.h, D), "gmem")
                dq_acc = tf.insert_slice(dq_acc, dq_tile, (0, t0, 0, 0))
            zero = tf.zeros(Tensor[(B @ cta.b, S @ cta.n, H @ cta.h, D), "f32", "gmem"])
            dk_z = tf.full_like(zero, value=0.0)
            dv_z = tf.full_like(zero, value=0.0)
            dk_g = tf.reshard(tf.transpose(dk, perm=(0, 2, 1, 3)), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")
            dv_g = tf.reshard(tf.transpose(dv, perm=(0, 2, 1, 3)), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")
            dk_f = dk_z + dk_g
            dv_f = dv_z + dv_g
        dq = dq_land(tf.reshard(dq_acc, (B, S, H, D), "gmem"))  # noqa: F821
        dk_out = land(tf.reshard(dk_f, (B, S, H, D), "gmem"))  # noqa: F821
        dv_out = land(tf.reshard(dv_f, (B, S, H, D), "gmem"))  # noqa: F821
        return dq, dk_out, dv_out


@module(entry="mha_bwd", target=_H200, topologies=(Topology("cta", 132),))
class Tiled:
    """One CTA per (batch, key block, head): K and V stay in smem, dK and dV in registers,
    and each query step settles its dQ partial in gmem."""

    @func
    def delta_rows(o: Tensor[(B, S, H, D), DT], do: Tensor[(B, S, H, D), DT]):
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            o_r = tf.reshard(o, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            do_r = tf.reshard(do, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            delta = tf.reduce(tf.cast(o_r, dtype="f32") * tf.cast(do_r, dtype="f32"), axes=(-1,), keepdim=True, kind="sum")
            return tf.reshard(delta, (B @ cta.b, S @ cta.n, H @ cta.h, 1), "gmem")

    @func
    def dq_land(dq_acc: Tensor[(B, S, H, D), "f32"]):
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            dq_r = tf.reshard(dq_acc, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem")
            dq_scaled = dq_r * tf.full_like(dq_r, value=SM_SCALE)
            return tf.reshard(tf.cast(dq_scaled, dtype=DT), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")

    @func
    def mha_bwd(
        q: Tensor[(B, S, H, D), DT],
        k: Tensor[(B, S, H, D), DT],
        v: Tensor[(B, S, H, D), DT],
        o: Tensor[(B, S, H, D), DT],
        do: Tensor[(B, S, H, D), DT],
        lse: Tensor[(B, H, S), "f32"],
    ):
        delta = tf.reshard(delta_rows(o, do), (B, S, H, 1), "gmem")  # noqa: F821
        lse_rows = tf.reshape(lse, new_shape=(B, H, 1, S))
        delta_rows_t = tf.reshape(tf.transpose(delta, perm=(0, 2, 1, 3)), new_shape=(B, H, 1, S))
        with Mesh(("cta",), layout=(B, NB, H), names=("b", "n", "h")) as cta:
            k_s = tf.reshard(k, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
            v_s = tf.reshard(v, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
            kv_pos = tf.reshard(
                tf.reshape(tf.arange(type=Tensor[(S,), "i32"]), new_shape=(1, 1, S, 1)),
                (1, 1, S @ cta.n, 1),
                "rmem",
            )
            q_pos_all = tf.reshape(tf.arange(type=Tensor[(S,), "i32"]), new_shape=(1, 1, 1, S))
            acc0 = tf.zeros(Tensor[(B @ cta.b, H @ cta.h, S @ cta.n, D), "f32", "rmem"])
            dk = tf.full_like(acc0, value=0.0)
            dv = tf.full_like(acc0, value=0.0)
            dq0 = tf.zeros(Tensor[(B @ cta.b, S, H @ cta.h, D), "f32", "gmem"])
            dq_acc = tf.full_like(dq0, value=0.0)
            for t in tile(S, BN):
                t0 = t + 0
                # K and V are loaded once per CTA; carrying them keeps that load outside the loop.
                k_s = tf.reshard(k_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
                v_s = tf.reshard(v_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "smem")
                q_s = tf.reshard(q[:, t0 : t0 + BN, :, :], (B @ cta.b, BN, H @ cta.h, D), "smem")
                do_s = tf.reshard(do[:, t0 : t0 + BN, :, :], (B @ cta.b, BN, H @ cta.h, D), "smem")
                lse_r = tf.reshard(lse_rows[:, :, :, t0 : t0 + BN], (B @ cta.b, H @ cta.h, 1, BN), "rmem")
                delta_r = tf.reshard(delta_rows_t[:, :, :, t0 : t0 + BN], (B @ cta.b, H @ cta.h, 1, BN), "rmem")
                q_pos = tf.reshard(q_pos_all[:, :, :, t0 : t0 + BN], (1, 1, 1, BN), "rmem")
                # WGMMA reads its shared-memory operands; these reshards are those reads.
                k_r = tf.transpose(tf.reshard(k_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                v_r = tf.transpose(tf.reshard(v_s, (B @ cta.b, S @ cta.n, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                q_r = tf.transpose(tf.reshard(q_s, (B @ cta.b, BN, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                do_r = tf.transpose(tf.reshard(do_s, (B @ cta.b, BN, H @ cta.h, D), "rmem"), perm=(0, 2, 1, 3))
                # Transposed scores: rows are keys, columns are queries.
                st = tf.cast(tf.matmul(k_r, q_r, b_layout="NK"), dtype="f32")
                pt = tf.exp2(st * tf.full_like(st, value=SM_SCALE * LOG2E) - lse_r)
                pt = tf.where(tf.cmp_le(kv_pos, q_pos), pt, tf.full_like(pt, value=0.0))
                p16 = tf.cast(pt, dtype=DT)
                dv = dv + tf.cast(tf.matmul(p16, do_r), dtype="f32")
                dpt = tf.cast(tf.matmul(v_r, do_r, b_layout="NK"), dtype="f32")
                # The softmax scale is applied to dK and dQ once, at the end.
                ds16 = tf.cast(pt * (dpt - delta_r), dtype=DT)
                dk = dk + tf.cast(tf.matmul(ds16, q_r), dtype="f32")
                # dQ contracts over the key rows the mesh split: a partial over cta.n,
                # read back from smem as the WGMMA A operand and settled by one add in gmem.
                ds_s = tf.reshard(ds16, (B @ cta.b, H @ cta.h, S @ cta.n, BN), "smem")
                ds_r = tf.reshard(ds_s, (B @ cta.b, H @ cta.h, S @ cta.n, BN), "rmem")
                dq_step = tf.transpose(tf.cast(tf.matmul(ds_r, k_r, a_layout="KM"), dtype="f32"), perm=(0, 2, 1, 3))
                dq_tile = tf.reshard(dq_step, (B @ cta.b, BN, H @ cta.h, D), "gmem")
                dq_acc = tf.insert_slice(dq_acc, dq_tile, (0, t0, 0, 0))
            dk_scaled = dk * tf.full_like(dk, value=SM_SCALE)
            dk_out = tf.reshard(tf.transpose(tf.cast(dk_scaled, dtype=DT), perm=(0, 2, 1, 3)), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")
            dv_out = tf.reshard(tf.transpose(tf.cast(dv, dtype=DT), perm=(0, 2, 1, 3)), (B @ cta.b, S @ cta.n, H @ cta.h, D), "gmem")
        dq = dq_land(tf.reshard(dq_acc, (B, S, H, D), "gmem"))  # noqa: F821
        return dq, dk_out, dv_out

The shipped op is held to the Tiled program through a runtime twin. delta_rows runs the preprocess kernel and mha_bwd runs GroupedQueryAttentionBwdOp; dq_land is 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:Fast passes against the evaluator on Tiled.mha_bwd and against torch autograd (1x256x2x128, same predicates, max_violation 0 on all three outputs).

#!/usr/bin/env python3
"""The runtime twin of the Tiled placement: TileOPs' MHA backward op and its preprocess."""
from mha_bwd import SM_SCALE, Tiled
from tilefoundry.runtime import runtime_func, runtime_module

from tileops.kernels.attention import FlashAttnBwdPreprocessKernel
from tileops.ops import GroupedQueryAttentionBwdOp


@runtime_module(Tiled)
class Fast:
    @runtime_func
    def delta_rows(self, o, do):
        batch, seq_len, heads, dim = o.shape
        delta, _ = FlashAttnBwdPreprocessKernel(batch, heads, seq_len, dim, o.dtype)(o, do)
        return delta.transpose(1, 2).unsqueeze(-1)

    @runtime_func
    def dq_land(self, dq_acc):
        # The shipped landing pass reads the accumulator in WGMMA register order, which
        # this function's argument does not have; the scale and cast are what it does.
        return (dq_acc * SM_SCALE).half()

    @runtime_func
    def mha_bwd(self, q, k, v, o, do, lse):
        return GroupedQueryAttentionBwdOp(is_causal=True)(q, k, v, o, do, lse)

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:

Module Shape gmem read gmem write f32 flops ideal-ns bound
Base 2x2048x32x128 897.8 MB 801.0 MB 4.40 G 413444 compute
Tiled 2x2048x32x128 385.8 MB 353.0 MB 4.13 G 409437 compute
Base 4x512x32x128 448.8 MB 400.5 MB 0.59 G 185526 memory
Tiled 4x512x32x128 192.8 MB 176.5 MB 0.55 G 80668 memory

What the analysis decided before any kernel existed:

  • Landing dK and dV in the input dtype from registers removes 512 MB of gmem reads and 448 MB of writes on the long shape (the zero-fill, the f32 add and the two land launches); the long shape stays compute-bound, the short one memory-bound, in both placements.
  • Folding the softmax scale out of dS removes 0.27 G f32 flops.
  • A BM/BN sweep of Tiled (64/64, 128/64, 128/128, 64/128) left gmem traffic unchanged and moved ideal-ns by 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.
  • The loop is authored over every query block with the causal mask applied, because a loop start that depends on the mesh coordinate is not priced; the causal kernel runs 272 of the 512 steps on the long shape, so its compute bound is about 0.53 of the ideal-ns above.
  • The short-shape rows price every step's dQ partial as a gmem round trip; the 33.5 MB accumulator fits in L2, so the kernel pays less than that.

Performance

Operator: GroupedQueryAttentionBwdOp, the -mha rows (H_kv = H) and the GQA rows (H_kv = 8)

Environment Value
image ghcr.io/tile-ai/tileops-runner:cu132-torch2.13-tl-afcebed1-tilefoundry-latest
digest sha256:75c76b1502d7730d9cc63eba6e418ce0cd5ae00e3e282a063fa73e425dcaf028
gpu NVIDIA H200
driver 595.71.05
cuda 13.2
torch 2.13.0+cu132
tilelang 0.1.11+cu132.gitafcebed1
tilefoundry 0.1.dev163+g3571ac3f7
flash_attn_3 3.0.0
timer native CUPTI device-busy time

Method: pytest benchmarks/ops/bench_mha.py benchmarks/ops/bench_gqa.py -k bwd in the image above, on one idle GPU, first on this branch (5b65f98) and then on its base main (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 through MultiHeadAttentionBwdOp before its removal; it delegated to GroupedQueryAttentionBwdOp, so the calls and kernels are the ones the -mha rows 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.

Workload B H S D Dtype This PR (us) main (us)
/ candidate
FA3 bwd (us)
/ candidate
llama-8b-short-float16 4 32 512 128 float16 101.4 220.1
🟢 2.171x
102.8
🟢 1.014x
llama-8b-short-bfloat16 4 32 512 128 bfloat16 100.7 423.8
🟢 4.209x
102.5
🟢 1.018x
llama-8b-long-float16 2 32 2048 128 float16 376.7 784.0
🟢 2.081x
370.9
🔴 0.985x
llama-8b-long-bfloat16 2 32 2048 128 bfloat16 372.5 1176.7
🟢 3.159x
366.0
🔴 0.983x
llama-70b-short-float16 2 64 512 128 float16 101.5 224.5
🟢 2.212x
103.0
🟢 1.015x
llama-70b-short-bfloat16 2 64 512 128 bfloat16 101.9 426.8
🟢 4.188x
102.7
🟢 1.008x
llama-70b-long-float16 1 64 2048 128 float16 376.6 782.7
🟢 2.078x
371.0
🔴 0.985x
llama-70b-long-bfloat16 1 64 2048 128 bfloat16 372.8 985.6
🟢 2.644x
366.7
🔴 0.984x
geometric mean 194.9 531.4
🟢 2.727x
194.6
🔴 0.999x

GQA rows (H_kv = 8), same run, served by the pipelined kernel:

B H H_kv S D Dtype This PR (us) main (us)
/ candidate
FA3 bwd (us)
/ candidate
4 32 8 512 128 float16 133.7 180.9
🟢 1.353x
125.2
🔴 0.936x
4 32 8 512 128 bfloat16 133.9 382.3
🟢 2.855x
124.9
🔴 0.933x
2 32 8 2048 128 float16 494.5 711.2
🟢 1.438x
410.9
🔴 0.831x
2 32 8 2048 128 bfloat16 490.4 1106.2
🟢 2.256x
403.6
🔴 0.823x
2 64 8 512 128 float16 125.0 171.9
🟢 1.375x
118.3
🔴 0.946x
2 64 8 512 128 bfloat16 125.3 374.8
🟢 2.991x
117.8
🔴 0.940x
1 64 8 2048 128 float16 489.8 688.2
🟢 1.405x
398.0
🔴 0.813x
1 64 8 2048 128 bfloat16 487.1 891.9
🟢 1.831x
393.0
🔴 0.807x
geometric mean 251.9 464.1
🟢 1.842x
220.8
🔴 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

  • Against main every 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.
  • Per launch at 2x2048x32 fp16 (ncu, fixed clocks): preprocess 32.7 us against FA3's 32.4, main kernel 315 us against 311, landing pass 23.7 us against 23.9. The remaining gap is the main kernel. Its DRAM, L2 and shared-memory traffic match FA3's to within 1%, so the difference is scheduling, not work.
  • What remains in the main kernel: each tile's epilogue stages dK and dV through the Q and dO stages and costs about 2 us per tile, during which the next tile's first Q and dO cannot load; and the loop body exists twice in SASS, once per consumer warpgroup (1074 instructions against FA3's 480 shared by both), because TileLang infers a fragment's layout from the threads of its enclosing block and cannot give two warpgroups one code path.
  • Measured and not kept: staging dK and dV in the dQ stages through 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).
  • GQA keeps the pipelined kernel; its rows are 0.81x to 0.95x of FA3's backward and are not addressed here.
  • Hopper only, head dim 128, sequence a multiple of 128: other calls keep the pipelined kernel. One kernel instance keeps its tile counter between launches and serves one stream at a time.

@lcy-seso
lcy-seso requested review from a team and a lite review from Copilot September 27, 2026 19:23
@lcy-seso lcy-seso added perf Performance improvements foundry Kernel generated by the TileFoundry tool labels Sep 27, 2026

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

…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.
…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
lcy-seso merged commit 2d60099 into tile-ai:main Sep 28, 2026
16 checks passed
@lcy-seso
lcy-seso deleted the perf/mha-bwd-ws branch September 28, 2026 07:45
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.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

foundry Kernel generated by the TileFoundry tool perf Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants