Skip to content

Enable 2CTA for SM100 block-sparse backward - #2661

Merged
drisspg merged 1 commit into
mainfrom
drisspg/stack/45
Jul 14, 2026
Merged

drisspg merged 1 commit into
mainfrom
drisspg/stack/45

Conversation

@drisspg

@drisspg drisspg commented Jun 16, 2026 •

Copy link
Copy Markdown
Collaborator

Enable 2CTA for SM100 block-sparse backward

Summary

Okay lots of lines of code but a decent amoutn of plumbing work for kv_subtile factor and some refactors around load utilities in blocksparese and the load loop.

The main chunk of new code is allowing for subtitling, this is basically a mirror of what we do in the fwd for q_subtile_factor where if q_block_size is a multiple of tile_m we subtile with tile_m. We now allow for the same along the k_seqlen dim with tile_n. This applies to both the fwd and the backward. This now lets us use 2cta which requires at least 2 tile_m chunks. And critical (the genesis of this PR) we can we support the DSV3 mla shapes.

Testing Testing Testing

You can see I added tests to the PR but besides that I(codex) wrote up a fuzztester;

import argparse
import math
import random
import sys
from dataclasses import dataclass
from pathlib import Path

import torch
import cutlass.cute as cute
from torch.nn.attention.flex_attention import flex_attention

REPO = Path(__file__).resolve().parents[1]
TESTS = REPO / "tests" / "cute"
if str(TESTS) not in sys.path:
    sys.path.insert(0, str(TESTS))

from score_mod_definitions import rel_bias_eager, score_mod_rel_bias  # noqa: E402
from test_mask_mod import (  # noqa: E402
    _build_block_sparse_masks_for_bwd,
    assert_bwd_matches_reference,
    assert_fwd_matches_reference,
    create_tensors,
    get_mask_pair,
    run_cute_mask_bwd,
    run_flex_reference_bwd,
)
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd  # noqa: E402
from flash_attn.cute.cache_utils import get_jit_cache  # noqa: E402
from flash_attn.cute.testing import is_fake_mode  # noqa: E402


@cute.jit
def score_mod_bwd_rel_bias(grad, score, b_idx, h_idx, q_idx, kv_idx, seqlen_info, aux_tensors):
    return grad


@dataclass(frozen=True)
class Case:
    seed: int
    seqlen_q: int
    seqlen_k: int
    headdim: int
    headdim_v: int
    sparse_tile_m: int
    sparse_tile_n: int
    mask_name: str
    score_mod_name: str = "none"
    dtype: torch.dtype = torch.bfloat16
    batch_size: int = 1
    nheads: int = 1
    tile_m: int = 128
    tile_n: int = 128


def make_cases(mode: str) -> list[Case]:
    curated = [
        Case(1001, 129, 257, 128, 128, 256, 256, "causal"),
        Case(1002, 257, 129, 128, 128, 256, 256, "causal"),
        Case(1003, 383, 769, 128, 128, 256, 512, "causal"),
        Case(1004, 384, 384, 192, 128, 256, 256, "causal"),
        Case(1005, 255, 385, 192, 128, 256, 256, "causal"),
        Case(1006, 513, 769, 128, 128, 256, 256, "sliding_window"),
        Case(1007, 513, 769, 128, 128, 256, 512, "sliding_window"),
        Case(1008, 512, 640, 128, 128, 256, 256, "causal"),
        Case(1009, 640, 512, 128, 128, 256, 512, "causal"),
        Case(1010, 768, 896, 192, 128, 256, 256, "sliding_window"),
        Case(2001, 383, 769, 192, 128, 256, 512, "causal"),
        Case(2002, 513, 769, 192, 128, 256, 512, "sliding_window"),
        Case(2003, 127, 255, 128, 128, 256, 256, "causal"),
        Case(2004, 128, 256, 192, 128, 256, 256, "causal"),
        Case(2005, 256, 512, 128, 128, 256, 512, "block_diagonal"),
        Case(2006, 257, 513, 192, 128, 256, 256, "mini_causal"),
        Case(2007, 385, 513, 128, 128, 256, 512, "mini_causal"),
        Case(2008, 511, 1025, 192, 128, 256, 512, "causal"),
        Case(2009, 1025, 511, 128, 128, 256, 256, "sliding_window"),
        Case(2010, 1024, 1024, 192, 128, 256, 512, "block_diagonal"),
    ]
    if mode == "quick":
        return curated[:4]
    if mode == "sanitizer":
        return [curated[2]]
    if mode == "racecheck":
        return [curated[4]]
    if mode == "flex200":
        return make_flex_reference_cases()
    return curated


def make_flex_reference_cases() -> list[Case]:
    seqlens = [
        (127, 255),
        (128, 256),
        (129, 257),
        (191, 383),
        (255, 385),
        (256, 512),
        (257, 513),
        (383, 769),
        (384, 640),
        (511, 1025),
        (512, 512),
        (513, 769),
        (640, 512),
        (768, 896),
        (896, 768),
        (1024, 1024),
        (1025, 511),
        (1152, 1280),
        (1536, 1024),
        (2048, 1536),
    ]
    mask_names = ["causal", "sliding_window", "block_diagonal", "mini_causal", "prefix_lm"]
    variants = [(128, 128, 256), (192, 128, 512)]
    return [
        Case(
            seed=3000 + case_idx,
            seqlen_q=seqlen_q,
            seqlen_k=seqlen_k,
            headdim=headdim,
            headdim_v=headdim_v,
            sparse_tile_m=256,
            sparse_tile_n=sparse_tile_n,
            mask_name=mask_name,
            score_mod_name="rel_bias",
        )
        for case_idx, (mask_name, (seqlen_q, seqlen_k), (headdim, headdim_v, sparse_tile_n)) in enumerate(
            (mask_name, seqlen, variant)
            for mask_name in mask_names
            for seqlen in seqlens
            for variant in variants
        )
    ]


def window_size(mask_name: str, seqlen_q: int, seqlen_k: int):
    if mask_name == "sliding_window":
        return max(64, min(seqlen_q, seqlen_k) // 3)
    return None


def case_score_mods(case: Case):
    if case.score_mod_name == "none":
        return None, None, None
    if case.score_mod_name == "rel_bias":
        return score_mod_rel_bias, score_mod_bwd_rel_bias, rel_bias_eager
    raise ValueError(f"Unknown score_mod_name: {case.score_mod_name}")


def run_score_mask_reference_bwd(q, k, v, block_mask, grad_out, eager_score_mod, dtype=None):
    if dtype is not None:
        q_ref = q.transpose(1, 2).to(dtype).requires_grad_(True)
        k_ref = k.transpose(1, 2).to(dtype).requires_grad_(True)
        v_ref = v.transpose(1, 2).to(dtype).requires_grad_(True)
        grad_out_ref = grad_out.transpose(1, 2).to(dtype)
    else:
        q_ref = q.transpose(1, 2).requires_grad_(True)
        k_ref = k.transpose(1, 2).requires_grad_(True)
        v_ref = v.transpose(1, 2).requires_grad_(True)
        grad_out_ref = grad_out.transpose(1, 2)

    out_ref = flex_attention(
        q_ref,
        k_ref,
        v_ref,
        score_mod=eager_score_mod,
        block_mask=block_mask,
        enable_gqa=True,
    )
    dq_ref, dk_ref, dv_ref = torch.autograd.grad(out_ref, (q_ref, k_ref, v_ref), grad_out_ref)
    return (
        out_ref.transpose(1, 2),
        dq_ref.transpose(1, 2),
        dk_ref.transpose(1, 2),
        dv_ref.transpose(1, 2),
    )


def run_case(case: Case) -> None:
    torch.manual_seed(case.seed)
    random.seed(case.seed)
    cute_score_mod, cute_score_mod_bwd, eager_score_mod = case_score_mods(case)
    mask_mod_cute, mask_mod_flex = get_mask_pair(
        case.mask_name,
        seqlen_q=case.seqlen_q,
        seqlen_k=case.seqlen_k,
        window_size=window_size(case.mask_name, case.seqlen_q, case.seqlen_k),
    )
    tensors = create_tensors(
        case.batch_size,
        case.seqlen_q,
        case.seqlen_k,
        case.nheads,
        case.nheads,
        case.headdim,
        case.headdim_v,
        case.dtype,
    )
    block_sparse_mask_fwd, block_sparse_mask_bwd, block_mask = _build_block_sparse_masks_for_bwd(
        mask_mod_flex=mask_mod_flex,
        batch_size=case.batch_size,
        nheads=case.nheads,
        seqlen_q=case.seqlen_q,
        seqlen_k=case.seqlen_k,
        tile_m=case.tile_m,
        tile_n=case.tile_n,
        spt=False,
        sparse_tile_m=case.sparse_tile_m,
        sparse_tile_n=case.sparse_tile_n,
        return_block_mask=True,
    )
    out_cute, lse_cute = _flash_attn_fwd(
        q=tensors["q"],
        k=tensors["k"],
        v=tensors["v"],
        out=tensors["out"],
        lse=tensors["lse"],
        cu_seqlens_q=None,
        cu_seqlens_k=None,
        seqused_q=None,
        seqused_k=None,
        page_table=None,
        softmax_scale=1.0 / math.sqrt(case.headdim),
        causal=False,
        softcap=None,
        window_size_left=None,
        window_size_right=None,
        learnable_sink=None,
        tile_mn=(case.tile_m, case.tile_n),
        pack_gqa=False,
        _arch=None,
        score_mod=cute_score_mod,
        mask_mod=mask_mod_cute,
        block_sparse_tensors=block_sparse_mask_fwd,
        return_lse=True,
    )
    grad_out = torch.randn_like(out_cute)
    if cute_score_mod is None:
        dq_cute, dk_cute, dv_cute = run_cute_mask_bwd(
            tensors["q"],
            tensors["k"],
            tensors["v"],
            out_cute,
            lse_cute,
            grad_out,
            mask_mod_cute,
            block_sparse_mask_bwd=block_sparse_mask_bwd,
            tile_m=case.tile_m,
            tile_n=case.tile_n,
        )
    else:
        dq_cute, dk_cute, dv_cute = _flash_attn_bwd(
            q=tensors["q"],
            k=tensors["k"],
            v=tensors["v"],
            out=out_cute,
            dout=grad_out,
            lse=lse_cute,
            causal=False,
            m_block_size=case.tile_m,
            n_block_size=case.tile_n,
            score_mod=cute_score_mod,
            score_mod_bwd=cute_score_mod_bwd,
            mask_mod=mask_mod_cute,
            block_sparse_tensors=block_sparse_mask_bwd,
        )
    if is_fake_mode():
        return
    if eager_score_mod is None:
        out_ref_fp32, dq_ref_fp32, dk_ref_fp32, dv_ref_fp32 = run_flex_reference_bwd(
            tensors["q"], tensors["k"], tensors["v"], block_mask, grad_out, dtype=torch.float32
        )
        out_pt, dq_pt, dk_pt, dv_pt = run_flex_reference_bwd(
            tensors["q"], tensors["k"], tensors["v"], block_mask, grad_out
        )
    else:
        out_ref_fp32, dq_ref_fp32, dk_ref_fp32, dv_ref_fp32 = run_score_mask_reference_bwd(
            tensors["q"], tensors["k"], tensors["v"], block_mask, grad_out, eager_score_mod, dtype=torch.float32
        )
        out_pt, dq_pt, dk_pt, dv_pt = run_score_mask_reference_bwd(
            tensors["q"], tensors["k"], tensors["v"], block_mask, grad_out, eager_score_mod
        )
    assert_fwd_matches_reference(out_cute, out_ref_fp32, out_pt)
    assert_bwd_matches_reference(
        dq_cute,
        dk_cute,
        dv_cute,
        dq_ref_fp32,
        dk_ref_fp32,
        dv_ref_fp32,
        dq_pt,
        dk_pt,
        dv_pt,
        case.dtype,
        min(case.seqlen_q, case.seqlen_k),
    )
    torch.cuda.synchronize()


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--mode", choices=["quick", "full", "sanitizer", "racecheck", "flex200"], default="full")
    parser.add_argument("--case", type=int, default=None)
    args = parser.parse_args()
    if torch.cuda.get_device_capability()[0] != 10:
        raise RuntimeError("SM100 fuzz requires compute capability 10.x")
    _flash_attn_fwd.compile_cache = get_jit_cache(f"agent_space.fuzz_fwd.{args.mode}")
    _flash_attn_bwd.compile_cache = get_jit_cache(f"agent_space.fuzz_bwd.{args.mode}")
    cases = make_cases(args.mode)
    if args.case is not None:
        cases = [cases[args.case]]
    for idx, case in enumerate(cases):
        print(f"CASE {idx}: {case}", flush=True)
        run_case(case)
        print(f"CASE {idx}: PASS", flush=True)


if __name__ == "__main__":
    main()

Results in;
image

Full Test suite

I used 32 threads and one of em got a lil too big for their britches and oomed but reran in isolation and we good;
image

Also test various 1 offs

Performance

Here is the chart comparing the new blocksizes with 2cta vs 1 cta:
image
At first I was like damnnn that sucks. But if you look at random, this make more sense. We have went from 256,128 -> 256, 256 so end up needing to visit way more tiles in the random case.

Below is a more fair comparison with Same blocksizes (256, 256) with 2cta vs 1cta.
bkv256_2cta_vs_bkv256_1cta_median

And now we can run deepseekshapes

drisspg added a commit that referenced this pull request Jun 16, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 72ee600 to 5b5c03e Compare June 16, 2026 19:19
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 16, 2026 19:37
drisspg added a commit that referenced this pull request Jun 16, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 5b5c03e to 7d9b4c5 Compare June 16, 2026 19:37
@drisspg
drisspg changed the base branch from main to drisspg/stack/44 June 16, 2026 19:37
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 16, 2026 21:43
drisspg added a commit that referenced this pull request Jun 16, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 7d9b4c5 to 124f501 Compare June 16, 2026 21:43
@drisspg
drisspg changed the base branch from main to drisspg/stack/44 June 16, 2026 21:43
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 16, 2026 23:23
drisspg added a commit that referenced this pull request Jun 16, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 124f501 to 3fea3c1 Compare June 16, 2026 23:23
@drisspg
drisspg changed the base branch from main to drisspg/stack/44 June 16, 2026 23:23
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 17, 2026 00:38
drisspg added a commit that referenced this pull request Jun 17, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 3fea3c1 to aebf6c0 Compare June 17, 2026 00:38
@drisspg
drisspg changed the base branch from main to drisspg/stack/44 June 17, 2026 00:38
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 17, 2026 00:42
drisspg added a commit that referenced this pull request Jun 17, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from aebf6c0 to b89e53d Compare June 17, 2026 00:42
@drisspg
drisspg changed the base branch from main to drisspg/stack/44 June 17, 2026 00:42
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 17, 2026 01:02
drisspg added a commit that referenced this pull request Jun 17, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from b89e53d to c589194 Compare June 17, 2026 01:02
@drisspg
drisspg changed the base branch from main to drisspg/stack/44 June 17, 2026 01:02
drisspg added a commit that referenced this pull request Jun 17, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from c589194 to a356b2a Compare June 17, 2026 01:08
@drisspg
drisspg changed the base branch from drisspg/stack/44 to main June 17, 2026 01:08
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 9d33fa6 to cb8f308 Compare June 17, 2026 19:34
@drisspg
drisspg marked this pull request as ready for review June 17, 2026 19:34
@drisspg
drisspg marked this pull request as draft June 17, 2026 19:35
@drisspg
drisspg marked this pull request as ready for review June 17, 2026 19:35

@chatgpt-codex-connector chatgpt-codex-connector Bot 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.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: cb8f3089f9

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread flash_attn/cute/block_sparse_utils.py Outdated
@drisspg
drisspg marked this pull request as draft June 17, 2026 19:53
drisspg added a commit that referenced this pull request Jun 17, 2026
stack-info: PR: #2661, branch: drisspg/stack/45
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from cb8f308 to c5e6ae0 Compare June 17, 2026 19:53
@drisspg
drisspg marked this pull request as ready for review June 17, 2026 19:53
@drisspg drisspg changed the title Enable 2CTA for SM100 block-sparse backward Enable 2CTA for SM100 block-sparse backward and kv subtiling for fwd Jun 17, 2026
@drisspg
drisspg marked this pull request as draft June 18, 2026 03:55
@drisspg drisspg changed the title Enable 2CTA for SM100 block-sparse backward and kv subtiling for fwd Enable 2CTA for SM100 block-sparse backward Jun 18, 2026
@drisspg
drisspg marked this pull request as ready for review June 18, 2026 03:55
@drisspg drisspg mentioned this pull request Jun 18, 2026

@chatgpt-codex-connector chatgpt-codex-connector Bot 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.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: c5e6ae04cf

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread flash_attn/cute/interface.py
@drisspg
drisspg marked this pull request as draft June 18, 2026 04:12
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from c5e6ae0 to 41a3d67 Compare June 18, 2026 04:12
@drisspg
drisspg marked this pull request as ready for review June 18, 2026 04:12
@drisspg
drisspg marked this pull request as draft June 18, 2026 04:19
@drisspg
drisspg force-pushed the drisspg/stack/45 branch from 41a3d67 to 580cb53 Compare June 18, 2026 04:20
@drisspg

drisspg commented Jul 4, 2026

Copy link
Copy Markdown
Collaborator Author

Just rebased -> still figure out my env to got an ima but I think it might just be 4.6 churn

@drisspg

drisspg commented Jul 6, 2026

Copy link
Copy Markdown
Collaborator Author

We good -> found a latent bug in some mask defs with OOB indexings so also fixed here

@chatgpt-codex-connector chatgpt-codex-connector Bot 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.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: db98127eaf

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread flash_attn/cute/block_sparse_utils.py
stack-info: PR: #2661, branch: drisspg/stack/45
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants