Skip to content

Feat([FA4][CUTE DSL]) Add head_dim=256 support (forward + backward) - #2412

Merged
Johnsonms merged 2 commits into
Dao-AILab:mainfrom
wangsiyu:head_dim_256
Apr 23, 2026
Merged

Johnsonms merged 2 commits into
Dao-AILab:mainfrom
wangsiyu:head_dim_256

Conversation

@wangsiyu

@wangsiyu wangsiyu commented Mar 30, 2026 •

Copy link
Copy Markdown
Contributor

Summary

This PR adds head_dim=256 support to the FA4 FlashAttention implementation built with the CUTLASS CUTE DSL.

What’s included

  • A dedicated kernel path specialized for head_dim=256
  • Forward and backward support
  • No API changes for existing head dimensions

Motivation

  • head_dim=256 is common in newer model variants, but FA4 CUTE-DSL coverage is currently limited.
  • Adding this path avoids falling back to a different implementation and improves end-to-end coverage.

Implementation

  • Forward: uses a 2-CTA design and introduces a new pipeline to better hide memory latency; includes a TMEM-based design for intermediate storage.
  • Backward: uses a 2-kernel approach and a 2-CTA design for the backward path.

Performance

Performance numbers will be added in a follow-up update once the benchmark suite and configurations are finalized.

Testing

  • Checked correctness against a reference implementation on representative shapes

Author

This feature is authored by @wangsiyu @dishengbin @cherichy @Johnsonms

@tridao

tridao commented Mar 30, 2026

Copy link
Copy Markdown
Member

Thanks for the contribution! Can you say more about the new pipeline for fwd?
Ideally we'd have a unified impl for fwd in 1 file.
For bwd i see you're computing dQ and then dKdV separately. That makes sense.

@wangsiyu

Copy link
Copy Markdown
Contributor Author

Thanks for the contribution! Can you say more about the new pipeline for fwd? Ideally we'd have a unified impl for fwd in 1 file. For bwd i see you're computing dQ and then dKdV separately. That makes sense.

Thank you for the feedback! We will try our best to clean up and reorganize the code to match this design.

@cherichy

cherichy commented Mar 31, 2026 •

Copy link
Copy Markdown

Thanks for the contribution! Can you say more about the new pipeline for fwd?

The main differences in the pipeline are:

  • Two CTAs cooperatively process a single 256×128 tile. This halves per-CTA shared memory usage, allowing us to increase the pipeline depth  (kv_stage=4 and qk_acc_stage=2 for QK).
image
  • The second QK GEMM is pulled ahead in the pipeline. Since head_dim=256 exceeds the MMA K-dimension limit (128), each QK and PV computation is split into two sub-iterations (iterations_qk=2). The pipeline interleaves the K-block loads so that the second K-block of the next iteration can begin loading while the current PV is still in flight.
image
  • Only one softmax warp group instead of two. Because head_dim=256 doubles the GEMM time per KV block, GEMM rather than softmax becomes the throughput bottleneck. A single warp group is sufficient to keep up, so we use 12 warps total (4 softmax + 4 correction + 1 MMA + 1 load + 2 scheduler/empty) instead of 16.

Regarding unifying into a single file: we agree that this is the ideal end state. However, the key blocker is that head_dim=256 requires a different tile shape and warp count (12 vs. 16), which in turn affects the number of pipeline stages, TMEM layouts, and shared memory allocation. We could unify them into one file by parameterizing these, but it would add significant complexity, so we prefer to keep them separate for now and work toward unification as a follow-up.

@tridao

tridao commented Mar 31, 2026

Copy link
Copy Markdown
Member

i see there's a separate test file for hdim 256. Do we still need that or does test_flash_attn.py cover what we need?

@wangsiyu

wangsiyu commented Mar 31, 2026 •

Copy link
Copy Markdown
Contributor Author

i see there's a separate test file for hdim 256. Do we still need that or does test_flash_attn.py cover what we need?

Yes, merging is possible. Currently in the process of organizing.

@wangsiyu

wangsiyu commented Mar 31, 2026 •

Copy link
Copy Markdown
Contributor Author

Forward kernel performance is improved via STG trick. The benchmark is ready to refresh.

@Johnsonms

Copy link
Copy Markdown
Collaborator

@wangsiyu Here are the performance benchmark results on our side for commit c784c2b. Please check whether they align with what your team observed. I will also benchmark the latest commits later today. Cc: @tzadouri

PR #2412 — Summary (hdim=256, B200)
• Forward:
Slower at short seq (≤1K), but wins from ≥2K, scaling to ~3.2× at 32K
Peak ~1307 TFLOPS (~58% MFU)
• Backward:
Faster at all lengths, up to ~3.9× at 32K
Gains driven by no full attention materialization (HBM savings)

Bottom line
Short seq: SDPA competitive
Long seq: FA4 clearly better (fwd + bwd)
Why: better memory efficiency + kernel design (2CTA)

image

@Johnsonms

Copy link
Copy Markdown
Collaborator

@wangsiyu @cherichy
Here are the latest benchmark:
image

Compared with last benchmark:

image

Key Takeaways:

  1. Short sequences (≤ 4096) improved most dramatically — kernel 2–3x faster, MFU jumped from as low as 8.5% to 27–62%. Previously losing to SDPA at short seqlen; now consistently winning.
  2. Long sequences (≥ 8192) improved moderately — 7–31% faster, MFU ceiling raised from ~58% to ~62–63%, suggesting the kernel is now compute-bound throughout.
  3. MFU is now consistently ~62% across all configs in c784c2b, vs highly variable (8.5%–57.8%) in 6010daa — the STG improvement eliminated the short-seqlen inefficiency.

@dishengbin

dishengbin commented Apr 1, 2026 •

Copy link
Copy Markdown

@Johnsonms, thanks for updating the perf numbers. The results are in line with our expectations. Our previous numbers are even slightly better, and we'll update the perf numbers from our side once they are ready.

@wangsiyu

wangsiyu commented Apr 1, 2026

Copy link
Copy Markdown
Contributor Author

@wangsiyu @cherichy Here are the latest benchmark: image

Compared with last benchmark:

image Key Takeaways:
  1. Short sequences (≤ 4096) improved most dramatically — kernel 2–3x faster, MFU jumped from as low as 8.5% to 27–62%. Previously losing to SDPA at short seqlen; now consistently winning.
  2. Long sequences (≥ 8192) improved moderately — 7–31% faster, MFU ceiling raised from ~58% to ~62–63%, suggesting the kernel is now compute-bound throughout.
  3. MFU is now consistently ~62% across all configs in c784c2b, vs highly variable (8.5%–57.8%) in 6010daa — the STG improvement eliminated the short-seqlen inefficiency.

That makes sense with these shapes. We will provide more benchmark with longer sequences

@cherichy

cherichy commented Apr 2, 2026

Copy link
Copy Markdown

@Johnsonms Thanks for the updated benchmark!
Here are our latest numbers comparing FA3 on H200 vs FA4 hd256 on B200 (head_dim=256, batch=1, seqlen up to 128k).

image

Key Takeaways:

Forward:

  • FA4 outperforms FA3 across all sequence lengths, with 2.0–2.3x speedup at seqlen ≥ 8k
  • Peak throughput reaches ~1839 TFLOPS (GQA, non-causal, seqlen=16k)
  • At short seqlen (4k), FA4 still wins by ~1.15x
  • At 128k, FA3 runs OOM while FA4 sustains ~1300 TFLOPS

Backward:

  • FA4 is consistently faster, with 1.4x at 4k scaling up to 2.6x at 64k, reaching ~950 TFLOPS at long seqlen

We will update CLC enabled data later once the performance is optimized.

@Johnsonms

Copy link
Copy Markdown
Collaborator

@Johnsonms Thanks for the updated benchmark! Here are our latest numbers comparing FA3 on H200 vs FA4 hd256 on B200 (head_dim=256, batch=1, seqlen up to 128k).

image Key Takeaways:

Forward:

  • FA4 outperforms FA3 across all sequence lengths, with 2.0–2.3x speedup at seqlen ≥ 8k
  • Peak throughput reaches ~1839 TFLOPS (GQA, non-causal, seqlen=16k)
  • At short seqlen (4k), FA4 still wins by ~1.15x
  • At 128k, FA3 runs OOM while FA4 sustains ~1300 TFLOPS

Backward:

  • FA4 is consistently faster, with 1.4x at 4k scaling up to 2.6x at 64k, reaching ~950 TFLOPS at long seqlen

We will update CLC enabled data later once the performance is optimized.

Will benchmark soon, let's check the alignment

@Johnsonms

Johnsonms commented Apr 2, 2026 •

Copy link
Copy Markdown
Collaborator

@Johnsonms Thanks for the updated benchmark! Here are our latest numbers comparing FA3 on H200 vs FA4 hd256 on B200 (head_dim=256, batch=1, seqlen up to 128k).

image Key Takeaways:

Forward:

  • FA4 outperforms FA3 across all sequence lengths, with 2.0–2.3x speedup at seqlen ≥ 8k
  • Peak throughput reaches ~1839 TFLOPS (GQA, non-causal, seqlen=16k)
  • At short seqlen (4k), FA4 still wins by ~1.15x
  • At 128k, FA3 runs OOM while FA4 sustains ~1300 TFLOPS

Backward:

  • FA4 is consistently faster, with 1.4x at 4k scaling up to 2.6x at 64k, reaching ~950 TFLOPS at long seqlen

We will update CLC enabled data later once the performance is optimized.

My benchmark in fa3 on H100 SMX and fa4 on B200

Cc: @tzadouri
image

Key observations:

  • Large seqlens (≥32768): FA4/B200 is consistently ~2x faster than FA3/H100 in both fwd and bwd, tracking closely
    with the raw hardware improvement (B200 ~2.25x BF16 peak over H100).
  • Small seqlen + causal (4096): Only 1.06–1.21x speedup — likely tile scheduling overhead at small problem sizes,
    not yet hitting peak occupancy.
  • Backward is catching up well: At large seqlens the bwd speedup (2.0–2.14x) matches or slightly exceeds fwd,
    suggesting the FA4 bwd kernel is well-tuned for B200.
  • GQA is slightly better than MHA at capturing the hardware ratio, especially at large seqlens.

Comparison: the contributor vs. ours Benchmarks

The two benchmarks are highly consistent, with only minor numerical differences.

Key Takeaways

  • FA4 (B200) > FA3 (H200):
    ~1.5×–2.1× speedup across forward/backward, causal/non-causal, and all sequence lengths.

  • Speedup grows with sequence length:
    ~1.1×–1.6× (4K) → ~1.9×–2.1× (64K–128K), driven by higher arithmetic intensity.

  • GQA > MHA (slightly):
    More stable and higher speedups, especially at long seq and backward, due to lower KV bandwidth.

  • Non-causal > causal:
    Higher TFLOPS and speedup; causal is limited by masking and reduced parallelism.

  • Backward ≥ Forward:
    Backward shows similar or slightly higher gains (~1.8×–2.1×); forward has more variance at small seq.


Conclusion: Both benchmarks align closely and confirm strong, scalable gains of FA4 on B200.

One case (16K fwd, non-causal) needs further investigation, with the contributor reaching higher peak throughput (~1800 TFLOPS vs. ~1300–1500). Cc: @tridao
image

@dishengbin

dishengbin commented Apr 3, 2026 •

Copy link
Copy Markdown

One case (16K fwd, non-causal) needs further investigation, with the contributor reaching higher peak throughput (~1800 TFLOPS vs. ~1300–1500).

@Johnsonms It’s probably because different numbers of iterations were used during benchmarking, which led to the discrepancy in the results. I reran this 16K fwd kernel separately: when the iteration count is 10, the throughput is about 1780 TFLOPS, and when the iteration count is 50, it’s about 1610 TFLOPS(see the following figures). Running kernels back-to-back continuously can cause the frequency to drop, which in turn degrades performance. I noticed that your test data all use 50 repetitions, right? In that case, we can adopt the same setting (GPU,CUDA/Driver version and CuTe DSL version as well) on our side and update the numbers accordingly.
image
image

@wangsiyu

wangsiyu commented Apr 3, 2026

Copy link
Copy Markdown
Contributor Author

@Johnsonms I think we should align on our scripts. We’ve noticed some discrepancies in dimensions with the script we’re currently using. Could you please share your script so we can test it in our environment?

@tridao tridao mentioned this pull request Apr 3, 2026
@wangsiyu

wangsiyu commented Apr 4, 2026

Copy link
Copy Markdown
Contributor Author

i see there's a separate test file for hdim 256. Do we still need that or does test_flash_attn.py cover what we need?

Unit tests have been merged into test_flash_attn.py and test_flash_atten_varlen.py。Unsupported cases for 256 dim will be temporally skipped.

interface.py's refine is on going.

@Johnsonms

Copy link
Copy Markdown
Collaborator

i see there's a separate test file for hdim 256. Do we still need that or does test_flash_attn.py cover what we need?

Unit tests have been merged into test_flash_attn.py and test_flash_atten_varlen.py。Unsupported cases for 256 dim will be temporally skipped.

interface.py's refine is on going.

Confirmed the test passed:
image

@Johnsonms

Copy link
Copy Markdown
Collaborator

@Johnsonms I think we should align on our scripts. We’ve noticed some discrepancies in dimensions with the script we’re currently using. Could you please share your script so we can test it in our environment?

Hi @wangsiyu Here is the script I used, Cc: @tzadouri

python bench_sm100_hd256.py --compare-baseline --nheads 16 --nheads-kv 16 --rep 50 --warmup 10

bench_sm100_hd256.py

#!/usr/bin/env python
"""SM100 Blackwell head_dim=256 benchmark (forward + backward).

Benchmarks the new 2CTA kernels added in PR #2412.  Uses the unified
`_flash_attn_fwd` / `_flash_attn_bwd` API which auto-routes to the
hd256 2CTA path on SM100/SM110 when head_dim=256.

Usage:
    # Default: fwd + bwd, seqlens 1k–16k, causal + non-causal
    python benchmarks/bench_sm100_hd256.py

    # Forward only
    python benchmarks/bench_sm100_hd256.py --direction fwd

    # Backward only
    python benchmarks/bench_sm100_hd256.py --direction bwd

    # Custom seqlens / batch
    python benchmarks/bench_sm100_hd256.py --seqlen 2048,4096,8192 --batch 2

    # Causal only
    python benchmarks/bench_sm100_hd256.py --causal-only

    # Compare FA hd256 vs PyTorch SDPA baseline
    python benchmarks/bench_sm100_hd256.py --compare-sdpa

    # Compile kernels without running (for two-pass workflow)
    python benchmarks/bench_sm100_hd256.py --compile-only

Two-pass workflow (compile in parallel, then run):
    FLASH_ATTENTION_FAKE_TENSOR=1 FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED=1 \\
        python benchmarks/bench_sm100_hd256.py --compile-only

    FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED=1 \\
        python benchmarks/bench_sm100_hd256.py
"""
import argparse
import contextlib
import io
import math
import os
import sys

import torch
import torch.nn.functional as F

from flash_attn.cute.interface import _flash_attn_fwd, _flash_attn_bwd
from benchmark_attn import get_peak_flops


@contextlib.contextmanager
def _suppress_stdout_stderr():
    """Suppress both Python-level (sys.stdout/err) and C-level (fd 1/2) output.

    Kernel debug prints (e.g. 'H>> shared_storage.size_in_bytes()') are emitted
    via Python print() during JIT compilation.  Redirecting only the fd is not
    enough because Python buffers writes through sys.stdout; we must also swap
    the Python stream objects.
    """
    devnull_fd = os.open(os.devnull, os.O_WRONLY)
    old_stdout_fd = os.dup(1)
    old_stderr_fd = os.dup(2)
    old_py_stdout = sys.stdout
    old_py_stderr = sys.stderr
    try:
        sys.stdout = io.StringIO()
        sys.stderr = io.StringIO()
        os.dup2(devnull_fd, 1)
        os.dup2(devnull_fd, 2)
        yield
    finally:
        # Restore Python streams first so subsequent prints go to real stdout
        sys.stdout = old_py_stdout
        sys.stderr = old_py_stderr
        os.dup2(old_stdout_fd, 1)
        os.dup2(old_stderr_fd, 2)
        os.close(devnull_fd)
        os.close(old_stdout_fd)
        os.close(old_stderr_fd)


# ── Helpers ────────────────────────────────────────────────────────────────

HEAD_DIM = 256
# Typical number of heads for head_dim=256 (matches newer model variants)
NHEADS = 8
NHEADS_KV = 8  # MHA; adjust to test GQA


def csv_ints(s):
    return [int(x.strip()) for x in s.split(",")]


def auto_batch(seqlen, batch_arg, total_tokens=32768):
    return batch_arg if batch_arg > 0 else max(1, total_tokens // seqlen)


def fwd_flops(batch, nheads, seqlen, hdim, causal=False):
    avg_seqlen = seqlen / 2 if causal else seqlen
    return batch * nheads * 2 * seqlen * avg_seqlen * (hdim + hdim)


def bwd_flops(batch, nheads, seqlen, hdim, causal=False):
    return 2.5 * fwd_flops(batch, nheads, seqlen, hdim, causal=causal)


def check_sm100():
    if not torch.cuda.is_available():
        print("ERROR: No CUDA device found.", file=sys.stderr)
        sys.exit(1)
    cap = torch.cuda.get_device_capability()
    name = torch.cuda.get_device_name()
    peak = get_peak_flops(0, dtype=torch.bfloat16)
    peak_str = f"  peak_bf16={peak/1e12:.0f} TFLOPS" if peak else ""
    if cap[0] not in (10, 11):
        print(
            f"WARNING: This benchmark targets SM100/SM110 (Blackwell). "
            f"Current GPU: {name} (SM{cap[0]}{cap[1]}). "
            f"The hd256 2CTA kernel may not be selected.",
            file=sys.stderr,
        )
    else:
        print(f"GPU: {name}  (SM{cap[0]}{cap[1]}){peak_str}")
    return peak


# ── Core bench functions ────────────────────────────────────────────────────

def bench_fwd(batch, seqlen, nheads, nheads_kv, causal,
              check_correctness=True, warmup=5, rep=30):
    """Benchmark hd256 forward pass. Returns (ms, tflops, max_diff_or_error)."""
    q = torch.randn(batch, seqlen, nheads,    HEAD_DIM, dtype=torch.bfloat16, device="cuda")
    k = torch.randn(batch, seqlen, nheads_kv, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
    v = torch.randn(batch, seqlen, nheads_kv, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
    scale = HEAD_DIM ** -0.5

    try:
        out, _lse = _flash_attn_fwd(q, k, v, softmax_scale=scale, causal=causal)
    except Exception as e:
        return None, None, str(e)[:120]

    max_diff = None
    if check_correctness:
        # Expand KV heads if GQA
        gqa = nheads // nheads_kv
        q_ref = q.transpose(1, 2).float()
        k_ref = k.transpose(1, 2).float().repeat_interleave(gqa, dim=1)
        v_ref = v.transpose(1, 2).float().repeat_interleave(gqa, dim=1)
        out_ref = F.scaled_dot_product_attention(
            q_ref, k_ref, v_ref, is_causal=causal, scale=scale
        )
        out_ref = out_ref.transpose(1, 2).to(torch.bfloat16)
        max_diff = (out.float() - out_ref.float()).abs().max().item()

    for _ in range(warmup):
        _flash_attn_fwd(q, k, v, softmax_scale=scale, causal=causal)

    torch.cuda.synchronize()
    start = torch.cuda.Event(enable_timing=True)
    end   = torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(rep):
        _flash_attn_fwd(q, k, v, softmax_scale=scale, causal=causal)
    end.record()
    torch.cuda.synchronize()

    ms = start.elapsed_time(end) / rep
    tflops = fwd_flops(batch, nheads, seqlen, HEAD_DIM, causal=causal) / ms / 1e9
    return ms, tflops, max_diff


def bench_bwd(batch, seqlen, nheads, nheads_kv, causal,
              check_correctness=True, warmup=5, rep=30):
    """Benchmark hd256 backward pass. Returns (ms, tflops, (dq_err, dk_err, dv_err) or error str)."""
    q = torch.randn(batch, seqlen, nheads,    HEAD_DIM, device="cuda", dtype=torch.bfloat16)
    k = torch.randn(batch, seqlen, nheads_kv, HEAD_DIM, device="cuda", dtype=torch.bfloat16)
    v = torch.randn(batch, seqlen, nheads_kv, HEAD_DIM, device="cuda", dtype=torch.bfloat16)
    scale = HEAD_DIM ** -0.5

    try:
        out, lse = _flash_attn_fwd(q, k, v, softmax_scale=scale, causal=causal, return_lse=True)
    except Exception as e:
        return None, None, str(e)[:120]

    dout = torch.randn_like(out)

    def fn():
        return _flash_attn_bwd(q, k, v, out, dout, lse, softmax_scale=scale, causal=causal)

    try:
        with _suppress_stdout_stderr():
            dq, dk, dv = fn()  # compile / warm JIT — suppresses kernel debug prints
    except Exception as e:
        return None, None, str(e)[:120]

    # Gradient correctness vs PyTorch reference
    grad_errs = None
    if check_correctness:
        gqa = nheads // nheads_kv
        q_ref = q.float().detach().requires_grad_(True)
        k_ref = k.float().detach().requires_grad_(True)
        v_ref = v.float().detach().requires_grad_(True)
        k_exp = k_ref.transpose(1, 2).repeat_interleave(gqa, dim=1)
        v_exp = v_ref.transpose(1, 2).repeat_interleave(gqa, dim=1)
        out_ref = F.scaled_dot_product_attention(
            q_ref.transpose(1, 2), k_exp, v_exp, is_causal=causal, scale=scale
        ).transpose(1, 2)
        out_ref.backward(dout.float())
        dq_err = (dq.float() - q_ref.grad).abs().max().item()
        # dK/dV reference grads are summed over GQA groups; take mean per KV head
        dk_ref = k_ref.grad
        dv_ref = v_ref.grad
        dk_err = (dk.float() - dk_ref).abs().max().item()
        dv_err = (dv.float() - dv_ref).abs().max().item()
        grad_errs = (dq_err, dk_err, dv_err)

    with _suppress_stdout_stderr():
        for _ in range(warmup):
            fn()

    torch.cuda.synchronize()
    start = torch.cuda.Event(enable_timing=True)
    end   = torch.cuda.Event(enable_timing=True)
    start.record()
    with _suppress_stdout_stderr():
        for _ in range(rep):
            fn()
    end.record()
    torch.cuda.synchronize()

    ms = start.elapsed_time(end) / rep
    tflops = bwd_flops(batch, nheads, seqlen, HEAD_DIM, causal=causal) / ms / 1e9
    return ms, tflops, grad_errs


def bench_sdpa_fwd(batch, seqlen, nheads, nheads_kv, causal, warmup=5, rep=30):
    """PyTorch SDPA baseline for forward."""
    scale = HEAD_DIM ** -0.5
    gqa = nheads // nheads_kv
    q = torch.randn(batch, nheads,    seqlen, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
    k = torch.randn(batch, nheads_kv, seqlen, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
    v = torch.randn(batch, nheads_kv, seqlen, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
    k = k.repeat_interleave(gqa, dim=1)
    v = v.repeat_interleave(gqa, dim=1)

    for _ in range(warmup):
        F.scaled_dot_product_attention(q, k, v, is_causal=causal, scale=scale)

    torch.cuda.synchronize()
    start = torch.cuda.Event(enable_timing=True)
    end   = torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(rep):
        F.scaled_dot_product_attention(q, k, v, is_causal=causal, scale=scale)
    end.record()
    torch.cuda.synchronize()

    ms = start.elapsed_time(end) / rep
    tflops = fwd_flops(batch, nheads, seqlen, HEAD_DIM, causal=causal) / ms / 1e9
    return ms, tflops


def bench_sdpa_bwd(batch, seqlen, nheads, nheads_kv, causal, warmup=5, rep=30):
    """PyTorch SDPA baseline for backward (fwd+bwd via autograd)."""
    scale = HEAD_DIM ** -0.5
    gqa = nheads // nheads_kv

    def make_inputs():
        q = torch.randn(batch, nheads,    seqlen, HEAD_DIM, dtype=torch.bfloat16,
                        device="cuda", requires_grad=True)
        k = torch.randn(batch, nheads_kv, seqlen, HEAD_DIM, dtype=torch.bfloat16,
                        device="cuda").repeat_interleave(gqa, dim=1).requires_grad_(True)
        v = torch.randn(batch, nheads_kv, seqlen, HEAD_DIM, dtype=torch.bfloat16,
                        device="cuda").repeat_interleave(gqa, dim=1).requires_grad_(True)
        return q, k, v

    q, k, v = make_inputs()
    out = F.scaled_dot_product_attention(q, k, v, is_causal=causal, scale=scale)
    dout = torch.randn_like(out)

    def fn():
        q_, k_, v_ = make_inputs()
        o = F.scaled_dot_product_attention(q_, k_, v_, is_causal=causal, scale=scale)
        o.backward(dout)

    for _ in range(warmup):
        fn()

    torch.cuda.synchronize()
    start = torch.cuda.Event(enable_timing=True)
    end   = torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(rep):
        fn()
    end.record()
    torch.cuda.synchronize()

    ms = start.elapsed_time(end) / rep
    tflops = bwd_flops(batch, nheads, seqlen, HEAD_DIM, causal=causal) / ms / 1e9
    return ms, tflops


# ── Formatting helpers ─────────────────────────────────────────────────────

def fmt_tflops_mfu(tflops, peak_flops, width=18):
    """Format TFLOPS with optional MFU% as a single string, e.g. '1373.3(61.0%)'."""
    if peak_flops is not None:
        mfu = tflops * 1e12 / peak_flops * 100
        cell = f"{tflops:.1f}({mfu:.1f}%)"
    else:
        cell = f"{tflops:.1f}"
    return f"{cell:>{width}}"


# ── Run modes ──────────────────────────────────────────────────────────────

def run_default(args, peak_flops=None):
    directions = ["fwd", "bwd"] if args.direction == "both" else [args.direction]
    causals = [True] if args.causal_only else ([False] if args.non_causal_only else [False, True])
    has_mfu = peak_flops is not None

    for direction in directions:
        dir_label = "Forward" if direction == "fwd" else "Backward"

        tflops_col = "FA4 TFLOPS(MFU%)" if has_mfu else "Throughput (TFLOPS)"
        tflops_w = max(len(tflops_col), 18)

        if direction == "fwd":
            hdr = (f"{'Config (attn-mask / seqlen)':<30} {'Batch':>6} "
                   f"{'Latency (ms)':>14} {tflops_col:>{tflops_w}} {'Max Abs Err (bf16)':>19}")
        else:
            hdr = (f"{'Config (attn-mask / seqlen)':<30} {'Batch':>6} "
                   f"{'Latency (ms)':>14} {tflops_col:>{tflops_w}} "
                   f"{'dQ Err':>10} {'dK Err':>10} {'dV Err':>10}")

        width = len(hdr)
        print(f"\n{'=' * width}")
        print(f"  SM100 Blackwell  head_dim=256  2CTA  {dir_label}  "
              f"nheads={args.nheads}  nheads_kv={args.nheads_kv}  (rep={args.rep})")
        print(f"{'=' * width}")
        print(hdr)
        print("-" * width)

        for seqlen in args.seqlen:
            batch = auto_batch(seqlen, args.batch)
            for causal in causals:
                mask_label = "causal" if causal else "non-causal"
                row_name = f"{mask_label} / seqlen={seqlen}"

                if direction == "fwd":
                    ms, tflops, diff = bench_fwd(
                        batch, seqlen, args.nheads, args.nheads_kv, causal,
                        warmup=args.warmup, rep=args.rep,
                    )
                    if ms is not None:
                        line = (f"{row_name:<30} {batch:>6} {ms:>14.3f} "
                                f"{fmt_tflops_mfu(tflops, peak_flops, tflops_w)}")
                        if diff is not None:
                            line += f" {diff:>19.6f}"
                        print(line)
                    else:
                        print(f"{row_name:<30} {batch:>6} {'FAIL':>14}  {diff}")
                else:
                    ms, tflops, grad_errs = bench_bwd(
                        batch, seqlen, args.nheads, args.nheads_kv, causal,
                        warmup=args.warmup, rep=args.rep,
                    )
                    if ms is not None:
                        line = (f"{row_name:<30} {batch:>6} {ms:>14.3f} "
                                f"{fmt_tflops_mfu(tflops, peak_flops, tflops_w)}")
                        if grad_errs is not None:
                            dq_e, dk_e, dv_e = grad_errs
                            line += f" {dq_e:>10.6f} {dk_e:>10.6f} {dv_e:>10.6f}"
                        print(line)
                    else:
                        print(f"{row_name:<30} {batch:>6} {'FAIL':>14}  {grad_errs}")


def run_compare_sdpa(args, peak_flops=None):
    """Compare FA hd256 forward vs PyTorch SDPA."""
    causals = [True] if args.causal_only else ([False] if args.non_causal_only else [False, True])
    has_mfu = peak_flops is not None

    tflops_col = "FA4 TFLOPS(MFU%)" if has_mfu else "FA TFLOPS"
    tflops_w = max(len(tflops_col), 18)

    hdr = (f"{'Config (attn-mask / seqlen)':<30} {'Batch':>6} "
           f"{'FA Latency (ms)':>16} {tflops_col:>{tflops_w}} "
           f"{'SDPA Latency (ms)':>18} {'SDPA TFLOPS':>12} {'Speedup':>8}")
    width = len(hdr)

    print(f"\n{'=' * width}")
    print(f"  SM100 Blackwell  head_dim=256  Forward:  FA 2CTA  vs  PyTorch SDPA  "
          f"nheads={args.nheads}  nheads_kv={args.nheads_kv}  (rep={args.rep})")
    print(f"{'=' * width}")
    print(hdr)
    print("-" * width)

    for seqlen in args.seqlen:
        batch = auto_batch(seqlen, args.batch)
        for causal in causals:
            mask_label = "causal" if causal else "non-causal"
            row_name = f"{mask_label} / seqlen={seqlen}"

            fa_ms, fa_tflops, diff = bench_fwd(
                batch, seqlen, args.nheads, args.nheads_kv, causal,
                check_correctness=False, warmup=args.warmup, rep=args.rep,
            )
            sdpa_ms, sdpa_tflops = bench_sdpa_fwd(
                batch, seqlen, args.nheads, args.nheads_kv, causal,
                warmup=args.warmup, rep=args.rep,
            )
            if fa_ms is not None:
                speedup = sdpa_ms / fa_ms
                line = (f"{row_name:<30} {batch:>6} {fa_ms:>16.3f} "
                        f"{fmt_tflops_mfu(fa_tflops, peak_flops, tflops_w)} "
                        f"{sdpa_ms:>18.3f} {sdpa_tflops:>12.1f} {speedup:>7.2f}x")
                print(line)
            else:
                print(f"{row_name:<30} {batch:>6} {'FAIL':>16}  {fa_tflops}")


def run_compare_baseline(args, peak_flops=None):
    """Full fwd+bwd comparison: FA hd256 2CTA vs PyTorch SDPA (the only viable baseline,
    since FA4 main does not support head_dim=256 on SM100).
    """
    causals = [True] if args.causal_only else ([False] if args.non_causal_only else [False, True])
    has_mfu = peak_flops is not None

    for direction, flops_fn, fa_bench_fn, sdpa_bench_fn in [
        ("Forward",  fwd_flops, bench_fwd,
         lambda b, s, nh, nhkv, c, **kw: bench_sdpa_fwd(b, s, nh, nhkv, c, **kw)),
        ("Backward", bwd_flops, bench_bwd,
         lambda b, s, nh, nhkv, c, **kw: bench_sdpa_bwd(b, s, nh, nhkv, c, **kw)),
    ]:
        tflops_col = "FA4 TFLOPS(MFU%)" if has_mfu else "FA TFLOPS"
        tflops_w = max(len(tflops_col), 18)

        hdr = (f"{'Config (attn-mask / seqlen)':<30} {'Batch':>6} "
               f"{'FA ms':>8} {tflops_col:>{tflops_w}} "
               f"{'SDPA ms':>9} {'SDPA TFLOPS':>12} {'Speedup':>8}")
        width = len(hdr)

        print(f"\n{'=' * width}")
        print(f"  FA hd256 2CTA  vs  PyTorch SDPA  [{direction}]  "
              f"nheads={args.nheads}  nheads_kv={args.nheads_kv}  (rep={args.rep})")
        print(f"  NOTE: FA4 main does not support head_dim=256; SDPA is the only baseline.")
        print(f"{'=' * width}")
        print(hdr)
        print("-" * width)

        for seqlen in args.seqlen:
            batch = auto_batch(seqlen, args.batch)
            for causal in causals:
                mask_label = "causal" if causal else "non-causal"
                row_name = f"{mask_label} / seqlen={seqlen}"

                fa_ms, fa_tflops, _ = fa_bench_fn(
                    batch, seqlen, args.nheads, args.nheads_kv, causal,
                    check_correctness=False, warmup=args.warmup, rep=args.rep,
                )
                sdpa_ms, sdpa_tflops = sdpa_bench_fn(
                    batch, seqlen, args.nheads, args.nheads_kv, causal,
                    warmup=args.warmup, rep=args.rep,
                )

                if fa_ms is not None:
                    speedup = sdpa_ms / fa_ms
                    line = (f"{row_name:<30} {batch:>6} {fa_ms:>8.3f} "
                            f"{fmt_tflops_mfu(fa_tflops, peak_flops, tflops_w)} "
                            f"{sdpa_ms:>9.3f} {sdpa_tflops:>12.1f} {speedup:>7.2f}x")
                    print(line)
                else:
                    print(f"{row_name:<30} {batch:>6} {'FAIL':>8}  {fa_tflops}")


def run_compile_only(args):
    """Trigger JIT compilation without timing — useful for two-pass workflow."""
    causals = [False, True]
    print("Compiling hd256 2CTA kernels (fwd + bwd) ...")
    for seqlen in args.seqlen:
        batch = auto_batch(seqlen, args.batch)
        for causal in causals:
            # fwd
            q = torch.randn(batch, seqlen, args.nheads,    HEAD_DIM, dtype=torch.bfloat16, device="cuda")
            k = torch.randn(batch, seqlen, args.nheads_kv, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
            v = torch.randn(batch, seqlen, args.nheads_kv, HEAD_DIM, dtype=torch.bfloat16, device="cuda")
            scale = HEAD_DIM ** -0.5
            try:
                out, lse = _flash_attn_fwd(q, k, v, softmax_scale=scale, causal=causal, return_lse=True)
                dout = torch.randn_like(out)
                _flash_attn_bwd(q, k, v, out, dout, lse, softmax_scale=scale, causal=causal)
                print(f"  compiled  causal={causal}  seqlen={seqlen}  batch={batch}")
            except Exception as e:
                print(f"  FAILED    causal={causal}  seqlen={seqlen}  batch={batch}  {e}")
    print("Done.")


# ── Main ──────────────────────────────────────────────────────────────────

def main():
    parser = argparse.ArgumentParser(
        description="SM100 head_dim=256 2CTA attention benchmark",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=__doc__,
    )
    parser.add_argument("--direction", choices=["fwd", "bwd", "both"], default="both")
    parser.add_argument(
        "--seqlen", type=csv_ints, default=[1024, 2048, 4096, 8192, 16384, 32768],
        help="Comma-separated sequence lengths (default: 1024,2048,4096,8192,16384,32768)",
    )
    parser.add_argument("--batch", type=int, default=0,
                        help="Batch size (0 = auto ~64k tokens)")
    parser.add_argument("--nheads",    type=int, default=NHEADS)
    parser.add_argument("--nheads-kv", type=int, default=NHEADS_KV, dest="nheads_kv")
    parser.add_argument("--warmup", type=int, default=5)
    parser.add_argument("--rep",    type=int, default=30)
    parser.add_argument("--causal-only",     action="store_true")
    parser.add_argument("--non-causal-only", action="store_true")
    parser.add_argument("--compare-sdpa",     action="store_true",
                        help="Compare FA hd256 fwd vs PyTorch SDPA")
    parser.add_argument("--compare-baseline", action="store_true",
                        help="Full fwd+bwd comparison vs PyTorch SDPA (the only viable "
                             "baseline since FA4 main does not support head_dim=256)")
    parser.add_argument("--compile-only",     action="store_true",
                        help="Compile kernels without benchmarking (two-pass step 1)")

    args = parser.parse_args()
    torch.manual_seed(0)
    peak_flops = check_sm100()

    if args.compile_only:
        run_compile_only(args)
    elif args.compare_baseline:
        run_compare_baseline(args, peak_flops=peak_flops)
    elif args.compare_sdpa:
        run_compare_sdpa(args, peak_flops=peak_flops)
    else:
        run_default(args, peak_flops=peak_flops)


if __name__ == "__main__":
    main()

@wangsiyu

wangsiyu commented Apr 7, 2026

Copy link
Copy Markdown
Contributor Author

i see there's a separate test file for hdim 256. Do we still need that or does test_flash_attn.py cover what we need?

Unit tests have been merged into test_flash_attn.py and test_flash_atten_varlen.py。Unsupported cases for 256 dim will be temporally skipped.
interface.py's refine is on going.

Confirmed the test passed: image

I am refining interface and will triggered these tests.

@dishengbin

Copy link
Copy Markdown

Hi @Johnsonms , thanks for providing the benchmarking scripts. I checked the script and found that the performance differences come from how the tensors are created.
In our previous scripts, we actually used randint and then converted the tensors to BF16, whereas your script uses randn. I believe these different initialization methods can affect the GPU frequency and thus the performance.
image

@wangsiyu

wangsiyu commented Apr 9, 2026 •

Copy link
Copy Markdown
Contributor Author

@Johnsonms All conflicts have been resolved and All relateive unit tests passed

@wangsiyu

wangsiyu commented Apr 9, 2026 •

Copy link
Copy Markdown
Contributor Author

Thanks for @Johnsonms ’s great help! Due to our current bandwidth constraints, we would greatly appreciate it if you could directly contribute code to this PR as well.

@Johnsonms

Copy link
Copy Markdown
Collaborator

Thanks for @Johnsonms ’s great help! Due to our current bandwidth constraints, we would greatly appreciate it if you could directly contribute code to this PR as well.

My pleasure. Thanks @wangsiyu

@umiswing

Copy link
Copy Markdown

Hi! Is this PR usable as-is now? And when will this PR be merged? Thanks!

@Johnsonms

Copy link
Copy Markdown
Collaborator

Hi! Is this PR usable as-is now? And when will this PR be merged? Thanks!

Thanks @umiswing for checking. Yes, the PR is usable as-is for now, and we are actively review and benchmark that and it should be merged very soon.

Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of e122e67 from Johnsonms/exp2-emu-hd256 on top of
merged main (hd256 PR Dao-AILab#2412, 27b4eb9). Original branch was based on a
pre-merge snapshot; the other five commits in that branch were absorbed
into the squash-merge.

Replace a fraction of hardware exp2 (SFU) instructions with a
polynomial FMA emulation (ex2_emulation_2) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 exp2 calls (ex2_emu_freq=4, ex2_emu_res=3) with packed FMA
polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write slot acquisition is moved earlier to
overlap any pipeline stall with the exp2 compute.

Original author benchmark (B200, bf16, hdim=256, 8 Q-heads, batch
~32k tokens, avg 9 runs, locked clocks @ 1965 MHz):

FWD Non-Causal (TFLOPS):
  seqlen  :   1k    2k    4k    8k   16k   32k   64k   96k  128k
  base    :  585  1258  1438  1525  1575  1602  1419  1448  1372
  exp2    :  728  1265  1488  1608  1680  1726  1569  1557  1560
  delta   : +25%   +0%   +3%   +5%   +7%   +8%  +11%   +8%  +14%

FWD Causal (TFLOPS):
  seqlen  :   1k    2k    4k    8k   16k   32k   64k   96k  128k
  base    :  347   702  1175  1356  1476  1552  1612  1453  1399
  exp2    :  343   709  1190  1384  1505  1586  1628  1434  1444
  delta   :  -1%   +1%   +1%   +2%   +2%   +2%   +1%   -1%   +3%

BWD: negligible impact (< 0.5% across all seqlens), no regression.

Post-rebase validation (B200, bf16, hdim=256, MHA 32:32 and GQA 32:2,
3-run means, locked clocks @ 1755 MHz, seqlens 4k..128k):
  - FWD: 19/24 cells positive; peak +7.4% at MHA 128k non-causal;
    avg MHA +2.3%, avg GQA +1.5%. Long-seqlen non-causal dominates
    (+5-7% at 64k/128k), matching the SFU-bottleneck theory.
  - Smaller magnitudes than the 1965 MHz numbers above: my bench uses
    32 Q-heads (vs 8) and lower sustained clock, both of which
    reduce the relative SFU bottleneck.
  - Correctness smoke: tests/cute/test_flash_attn.py::test_flash_attn_output
    -k "256-False-0-0.0-False-False" -> 78 passed, 78 skipped, 0 failed
    (same pass/skip as origin/main).
Johnsonms added a commit that referenced this pull request Apr 23, 2026
Follow-up polish on the freshly-merged hd256 feature (#2412), sourced
from Copilot AI review comments on the original PR.

interface.py: drop duplicate `from cutlass import Int32` (already imported
at line 17) and unused `from flash_attn.cute.mask import Sm100MaskEnum as
MaskEnum`, which is never referenced.

mask.py: remove two dead `tidx, tidy, tidx = cute.arch.thread_idx()` lines
in Sm100FusedMask.apply_mask and apply_mask_via_causal_local. Neither
`tidx` nor `tidy` is ever read in the function bodies; these calls are
leftover debug scaffolding (consistent with the commented-out
`cute.printf("tidx = ...")` lines nearby at 490/525/665).

test_flash_attn.py: drop the stray "/SM110" from two TODO comments. The
skip guard is `IS_SM100` only (capability major == 10), and the hd256
2CTA kernel path is only taken when `arch // 10 == 10` (interface.py:573,
1310), never on SM110 (major == 11).
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `27b4eb9`). Original branch was based on
a pre-merge snapshot; the other five commits in that branch were
absorbed into the squash-merge.

## Change

Replace a fraction of hardware `exp2` (SFU) instructions with a
polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed
FMA polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write-slot acquisition is moved earlier to overlap
any pipeline stall with the `exp2` compute.

Kernel-only change; no API change.

## Original author benchmark

B200, bf16, hdim=256, 8 Q-heads, batch ~32k tokens, avg 9 runs, locked
clocks @ 1965 MHz.

### FWD Non-Causal (TFLOPS)

| seqlen |   1k |   2k |   4k |   8k |  16k |  32k |  64k |  96k | 128k |
|-------:|-----:|-----:|-----:|-----:|-----:|-----:|-----:|-----:|-----:|
| base   |  585 | 1258 | 1438 | 1525 | 1575 | 1602 | 1419 | 1448 | 1372 |
| exp2   |  728 | 1265 | 1488 | 1608 | 1680 | 1726 | 1569 | 1557 | 1560 |
| delta  | +25% |  +0% |  +3% |  +5% |  +7% |  +8% | +11% |  +8% | +14% |

### FWD Causal (TFLOPS)

| seqlen |   1k |   2k |   4k |   8k |  16k |  32k |  64k |  96k | 128k |
|-------:|-----:|-----:|-----:|-----:|-----:|-----:|-----:|-----:|-----:|
| base   |  347 |  702 | 1175 | 1356 | 1476 | 1552 | 1612 | 1453 | 1399 |
| exp2   |  343 |  709 | 1190 | 1384 | 1505 | 1586 | 1628 | 1434 | 1444 |
| delta  |  -1% |  +1% |  +1% |  +2% |  +2% |  +2% |  +1% |  -1% |  +3% |

BWD: negligible impact (< 0.5% across all seqlens), no regression.

## Post-rebase validation

B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked
clocks @ 1755 MHz, seqlens 4k..128k.

### FWD delta vs origin/main (TFLOPS mean, 3 runs each)

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  -0.2%    |  +0.7%   |
| 8k     |   F    |  +0.4%    |  +0.3%   |
| 16k    |   F    |  +0.7%    |  -1.0%   |
| 32k    |   F    |  +2.3%    |  -0.3%   |
| 64k    |   F    |  **+5.4%** |  **+5.2%** |
| 128k   |   F    |  **+7.4%** |  **+7.3%** |
| 4k     |   T    |  +0.3%    |  +0.9%   |
| 8k     |   T    |  +0.8%    |  +0.9%   |
| 16k    |   T    |  +0.8%    |  +1.1%   |
| 32k    |   T    |  **+5.2%** |  -0.4%   |
| 64k    |   T    |  +1.1%    |  +1.1%   |
| 128k   |   T    |  **+3.4%** |  **+2.6%** |

- 19 of 24 cells positive; 4 slightly negative (all within the
  batch-quantization noise band we measured on origin/main).
- Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax
  SFU pressure is worst.
- Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**.
- Smaller magnitudes than the 1965 MHz numbers above are expected: my
  bench uses 32 Q-heads (vs 8) and lower sustained clock, both of which
  reduce the relative SFU bottleneck. Direction and long-seqlen
  dominance match the commit's theory.

### Correctness smoke

`pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k
"256-False-0-0.0-False-False"` on B200:
**78 passed, 78 skipped, 0 failed** — same pass/skip count as
`origin/main`.

## Caveat

Exp2 FMA emulation introduces small numerical differences vs hardware
`exp2`. The existing test tolerances accept the delta.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `27b4eb9`). The original branch was
based on a pre-merge snapshot; the other five commits in that branch
were absorbed into the squash-merge, leaving this one novel change.

## Change

Replace a fraction of hardware `exp2` (SFU) instructions with a
polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed
FMA polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write-slot acquisition is moved earlier to overlap
any pipeline stall with the `exp2` compute.

Kernel-only change; no API change. Backward is untouched.

## Validation (this PR vs current `origin/main`)

B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked
clocks @ 1755 MHz, seqlens 4k..128k.

### FWD delta vs `origin/main` (TFLOPS mean, 3 runs each)

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  -0.2%    |  +0.7%   |
| 8k     |   F    |  +0.4%    |  +0.3%   |
| 16k    |   F    |  +0.7%    |  -1.0%   |
| 32k    |   F    |  +2.3%    |  -0.3%   |
| 64k    |   F    |  **+5.4%** |  **+5.2%** |
| 128k   |   F    |  **+7.4%** |  **+7.3%** |
| 4k     |   T    |  +0.3%    |  +0.9%   |
| 8k     |   T    |  +0.8%    |  +0.9%   |
| 16k    |   T    |  +0.8%    |  +1.1%   |
| 32k    |   T    |  **+5.2%** |  -0.4%   |
| 64k    |   T    |  +1.1%    |  +1.1%   |
| 128k   |   T    |  **+3.4%** |  **+2.6%** |

- 19 of 24 cells positive; 4 slightly negative, all within the
  batch-quantization noise band that `origin/main` itself already
  showed in our 3-run regression sweep.
- Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax
  SFU pressure is worst, consistent with the theory above.
- Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**.

### Correctness smoke

`pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k
"256-False-0-0.0-False-False"` on B200:
**78 passed, 78 skipped, 0 failed** — identical pass/skip count to
`origin/main`.

## Caveat

Exp2 FMA emulation introduces small numerical differences vs hardware
`exp2`. The existing test tolerances accept the delta.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `27b4eb9`). The original branch was
based on a pre-merge snapshot; the other five commits in that branch
were absorbed into the squash-merge, leaving this one novel change.

## Change

Replace a fraction of hardware `exp2` (SFU) instructions with a
polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed
FMA polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write-slot acquisition is moved earlier to overlap
any pipeline stall with the `exp2` compute.

Kernel-only change; no API change. Backward is untouched.

## Validation (this PR vs current `origin/main`)

B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked
clocks @ 1755 MHz, seqlens 4k..128k.

### FWD delta vs `origin/main` (TFLOPS mean, 3 runs each)

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  -0.2%    |  +0.7%   |
| 8k     |   F    |  +0.4%    |  +0.3%   |
| 16k    |   F    |  +0.7%    |  -1.0%   |
| 32k    |   F    |  +2.3%    |  -0.3%   |
| 64k    |   F    |  **+5.4%** |  **+5.2%** |
| 128k   |   F    |  **+7.4%** |  **+7.3%** |
| 4k     |   T    |  +0.3%    |  +0.9%   |
| 8k     |   T    |  +0.8%    |  +0.9%   |
| 16k    |   T    |  +0.8%    |  +1.1%   |
| 32k    |   T    |  **+5.2%** |  -0.4%   |
| 64k    |   T    |  +1.1%    |  +1.1%   |
| 128k   |   T    |  **+3.4%** |  **+2.6%** |

- 19 of 24 cells positive; 4 slightly negative, all within the
  batch-quantization noise band that `origin/main` itself already
  showed in our 3-run regression sweep.
- Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax
  SFU pressure is worst, consistent with the theory above.
- Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**.

### Correctness smoke

`pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k
"256-False-0-0.0-False-False"` on B200:
**78 passed, 78 skipped, 0 failed** — identical pass/skip count to
`origin/main`.

## Caveat

Exp2 FMA emulation introduces small numerical differences vs hardware
`exp2`. The existing test tolerances accept the delta.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `27b4eb9`). The original branch was
based on a pre-merge snapshot; the other five commits in that branch
were absorbed into the squash-merge, leaving this one novel change.

## Change

Replace a fraction of hardware `exp2` (SFU) instructions with a
polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed
FMA polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write-slot acquisition is moved earlier to overlap
any pipeline stall with the `exp2` compute.

Kernel-only change; no API change. Backward is untouched.

## Validation (this PR vs `origin/main` @ `b21e204` — includes Dao-AILab#2412 hd256 base and Dao-AILab#2487 post-merge cleanup)

B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked
clocks @ 1755 MHz, seqlens 4k..128k.

### FWD delta vs `origin/main` (TFLOPS mean, 3 runs each)

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  -0.2%    |  +0.7%   |
| 8k     |   F    |  +0.4%    |  +0.3%   |
| 16k    |   F    |  +0.7%    |  -1.0%   |
| 32k    |   F    |  +2.3%    |  -0.3%   |
| 64k    |   F    |  **+5.4%** |  **+5.2%** |
| 128k   |   F    |  **+7.4%** |  **+7.3%** |
| 4k     |   T    |  +0.3%    |  +0.9%   |
| 8k     |   T    |  +0.8%    |  +0.9%   |
| 16k    |   T    |  +0.8%    |  +1.1%   |
| 32k    |   T    |  **+5.2%** |  -0.4%   |
| 64k    |   T    |  +1.1%    |  +1.1%   |
| 128k   |   T    |  **+3.4%** |  **+2.6%** |

- 19 of 24 cells positive; 4 slightly negative, all within the
  batch-quantization noise band that `origin/main` itself already
  showed in our 3-run regression sweep.
- Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax
  SFU pressure is worst, consistent with the theory above.
- Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**.

### Correctness smoke

`pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k
"256-False-0-0.0-False-False"` on B200:
**78 passed, 78 skipped, 0 failed** — identical pass/skip count to
`origin/main`.

## Caveat

Exp2 FMA emulation introduces small numerical differences vs hardware
`exp2`. The existing test tolerances accept the delta.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `49fe257` from `Johnsonms/paged-kv-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `27b4eb9` + post-merge cleanup Dao-AILab#2487,
`b21e204`). Original branch was based on a pre-merge snapshot; its
base commits were absorbed into the squash-merge.

## Change

Adds paged KV support to the SM100 hd256 2CTA forward kernel. The paged
path reuses the dense TMA load path — logical KV blocks are remapped to
physical page indices through the page table at load time, so each page
maps to exactly one TMA tile.

**Constraint:** `page_size` must equal `tile_n = 128`.

### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`

- Conditional K/V tensor layout in `__call__`: dense
  `(s_k, d, ((h_r, h_k), b))` vs paged
  `(page_size, d, h_k, num_pages)` for K (and transposed for V).
- Conditional K/V TMA setup in the load warp: dense uses
  `domain_offset` + batch indexing; paged uses `head_kv` slicing and
  keeps `num_pages` as the outer mode for per-load `page_idx` lookup.
- Conditional per-load `page_idx`: K uses mode-2 subtile + mode-3 page;
  V uses mode-1 page.
- Plumb `mPageTable` + `max_seqlen_k` through the kernel signature.
  `seqlen_k` in each of the 4 warp sections now uses `max_seqlen_k`
  for the paged path.
- Store `qhead_per_kvhead` on `self` and derive `head_kv_coord` via
  integer divide (matches the `flash_fwd_sm100` convention for
  contiguous GQA grouping).
- Relax `mPageTable` / `paged_kv_non_tma` assertions.

### `flash_attn/cute/paged_kv.py`

- Extract `_flatten_smem_sm100` / `_copy_row_async` helpers from
  `load_KV` — pure refactor, no behavior change for existing callers.

### `tests/cute/test_flash_attn.py`

- `test_flash_attn_paged_hd256_sm100_tma`: bit-exact vs dense varlen
  reference + determinism check, parametrized over `seqlen_q`.
- `test_flash_attn_paged_hd256_sm100_tma_gqa`: same check for GQA with
  `nheads_kv in {2, 4, 8}` — exercises `qhead_per_kvhead > 1`, which
  a modulo-aliasing bug would fail.

## Validation (this PR vs `origin/main` @ `b21e204`)

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 new paged tests with the existing d=256 dense subset:

```
-k "paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **84 passed, 78 skipped, 0 failed** in 2 min — 78 from the
dense d=256 subset (identical pass/skip count to `origin/main`) and
**6 from the new `paged_hd256_sm100_tma[_gqa]` tests**.

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  +0.2%    |  +0.3%   |
| 8k     |   F    |   0.0%    |  -0.1%   |
| 16k    |   F    |  -0.2%    |   0.0%   |
| 32k    |   F    |  -1.0%    |  **+2.3%** |
| 64k    |   F    |  **+2.1%** |  +0.8%  |
| 128k   |   F    |  -0.8%    |  -1.7%   |
| 4k     |   T    |  +0.2%    |  +0.2%   |
| 8k     |   T    |  +0.1%    |  +0.1%   |
| 16k    |   T    |  +0.3%    |  +0.1%   |
| 32k    |   T    |  **+3.0%** |  -1.2%  |
| 64k    |   T    |  -0.3%    |  -0.1%   |
| 128k   |   T    |  +0.1%    |  -0.7%   |

- **22 of 24 cells within ±2%.**
- Two `> 2%` outliers are both **positive** and in the batch-
  quantization noise zone that `origin/main` itself showed run-spread
  in during our 3-run baseline sweep — not regressions.
- **Aggregated means: MHA +0.31%, GQA +0.00%.**
- Paged-KV path isn't exercised by `benchmark_attn.py` (which uses
  contiguous KV); dense-path perf parity is the regression-critical
  property and is preserved.

## Caveat

- **page_size == tile_n == 128 is a hard constraint.** Callers that
  want a different page size will need a separate path.
- The paged-KV path itself is correctness-tested by the two new
  `paged_hd256_sm100_tma` tests (bit-exact vs dense reference, with
  and without GQA). Perf of the paged path was not benchmarked.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `94e63db` from `Johnsonms/seqused-k-hd256` on
top of `Johnsonms/paged-kv-hd256-v2`. Original branch was based on a
pre-merge snapshot of the hd256 tree; base commits were absorbed into
the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487
(`b21e204`).

**Stacked PR** — base branch should be `Johnsonms/paged-kv-hd256-v2`,
not `main`. Depends on that PR for the paged-KV kernel plumbing.

## Change

Enables variable per-batch KV sequence lengths via a `seqused_k`
tensor — needed for MLA-style decode (DeepSeek-V2 / V3 / R1), where
different batches have different KV cache occupancies. Works with
both dense and paged K/V.

### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`

- Drop the `mSeqUsedK` half of the `__call__` assertion; only
  `mSeqUsedQ` stays blocked for now.
- Build the `seqused_k` cute tensor in `__call__` alongside the
  `page_table` / paged-layout construction.
- Add `mSeqUsedK` to `kernel()` signature and pass `seqused_k` from
  `__call__`.
- Replace `seqlen_k` derivation in all 4 warp sections with a ternary:
  `mSeqUsedK[batch_coord] if set, else <dense/paged expression>`.
- Move `batch_coord` above `seqlen_k` in the MMA warp (second warp
  section) — it was declared later but now needs to be in scope for
  `seqused_k` indexing.
- **Zero-KV batch handling (`seqlen_k == 0`):** extend `continue_cond`
  in all four warp sections with `continue_cond or seqlen_k <= 0`, so
  load / MMA / correction / softmax warps skip in sync instead of
  deadlocking on `K0 / Vend / QK0 / PVend / first-stats` tiles.

### `flash_attn/cute/interface.py`

- Relax the `seqused_q is None and seqused_k is None` assertion to
  `seqused_q is None` (the kernel now handles `seqused_k`).
- Prefill zero-KV batches on the host with zero output and `-inf`
  LSE, so their output is defined even though the kernel skips them.

### `tests/cute/test_flash_attn.py`

- `test_flash_attn_seqused_k_hd256_sm100`: dense + `seqused_k`
  (padded K/V with per-batch valid lengths) bit-exact vs a
  `cu_seqlens_k` packed reference, parametrized over asymmetric
  per-batch lengths.
- `test_flash_attn_paged_seqused_k_hd256_sm100`: paged + `seqused_k`
  combined (MLA decode pattern), bit-exact vs packed reference.
- `test_flash_attn_seqused_k_zero_hd256_sm100`: `seqused_k = 0` for
  one batch, parametrized dense/paged. Verifies no deadlock, zero
  output, `-inf` LSE on the empty row, and finite output / LSE on
  the other row.

## Validation

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 new `seqused_k` tests, the 6 paged tests from the parent PR, and
the existing d=256 dense subset:

```
-k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or
    seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or
    (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **90 passed, 78 skipped, 0 failed** in 2m14s — 78 from the
dense d=256 subset (identical pass/skip count to `origin/main`),
6 from the parent PR's `paged_hd256_sm100_tma[_gqa]` tests, and
**6 from the 3 new `seqused_k_*_hd256_sm100` tests** introduced here.

### FWD perf

Not re-measured on this branch. Expected to track the parent PR's
numbers (within ±0.3% mean of `origin/main`) on the dense path, since
the `seqused_k` plumbing only adds a scalar-tensor indirection per
batch and a `continue_cond` check per warp section — no change to
inner-loop throughput.

## Caveat

- Host-side prefill in `interface.py` walks zero-KV batch indices in
  Python. For batch sizes in the thousands with many zero-KV rows
  this could show up in CPU profiles; the current FA4 callers aren't
  in that regime.
- `seqused_q` is still asserted `None` — the kernel would need a
  similar ternary in the Q-side iteration, out of scope for this PR.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `94e63db` from `Johnsonms/seqused-k-hd256` on
top of `Johnsonms/paged-kv-hd256-v2`. Original branch was based on a
pre-merge snapshot of the hd256 tree; base commits were absorbed into
the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487
(`b21e204`).

**Stacked PR** — base branch should be `Johnsonms/paged-kv-hd256-v2`,
not `main`. Depends on that PR for the paged-KV kernel plumbing.

## Change

Enables variable per-batch KV sequence lengths via a `seqused_k`
tensor — needed for MLA-style decode (DeepSeek-V2 / V3 / R1), where
different batches have different KV cache occupancies. Works with
both dense and paged K/V.

### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`

- Drop the `mSeqUsedK` half of the `__call__` assertion; only
  `mSeqUsedQ` stays blocked for now.
- Build the `seqused_k` cute tensor in `__call__` alongside the
  `page_table` / paged-layout construction.
- Add `mSeqUsedK` to `kernel()` signature and pass `seqused_k` from
  `__call__`.
- Replace `seqlen_k` derivation in all 4 warp sections with a ternary:
  `mSeqUsedK[batch_coord] if set, else <dense/paged expression>`.
- Move `batch_coord` above `seqlen_k` in the MMA warp (second warp
  section) — it was declared later but now needs to be in scope for
  `seqused_k` indexing.
- **Zero-KV batch handling (`seqlen_k == 0`):** extend `continue_cond`
  in all four warp sections with `continue_cond or seqlen_k <= 0`, so
  load / MMA / correction / softmax warps skip in sync instead of
  deadlocking on `K0 / Vend / QK0 / PVend / first-stats` tiles.

### `flash_attn/cute/interface.py`

- Relax the `seqused_q is None and seqused_k is None` assertion to
  `seqused_q is None` (the kernel now handles `seqused_k`).
- Prefill zero-KV batches on the host with zero output and `-inf`
  LSE, so their output is defined even though the kernel skips them.

### `tests/cute/test_flash_attn.py`

- `test_flash_attn_seqused_k_hd256_sm100`: dense + `seqused_k`
  (padded K/V with per-batch valid lengths) bit-exact vs a
  `cu_seqlens_k` packed reference, parametrized over asymmetric
  per-batch lengths.
- `test_flash_attn_paged_seqused_k_hd256_sm100`: paged + `seqused_k`
  combined (MLA decode pattern), bit-exact vs packed reference.
- `test_flash_attn_seqused_k_zero_hd256_sm100`: `seqused_k = 0` for
  one batch, parametrized dense/paged. Verifies no deadlock, zero
  output, `-inf` LSE on the empty row, and finite output / LSE on
  the other row.

## Validation

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 new `seqused_k` tests, the 6 paged tests from the parent PR, and
the existing d=256 dense subset:

```
-k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or
    seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or
    (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **90 passed, 78 skipped, 0 failed** — 78 from the dense d=256
subset (identical pass/skip count to `origin/main`), 6 from the parent
PR's `paged_hd256_sm100_tma[_gqa]` tests, and **6 from the 3 new
`seqused_k_*_hd256_sm100` tests** introduced here.

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

#### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1474 | 1485 | +0.7% |
| 8k     | F | 1582 | 1595 | +0.8% |
| 16k    | F | 1641 | 1650 | +0.5% |
| 32k    | F | 1450 | 1488 | **+2.6%** |
| 64k    | F | 1417 | 1484 | **+4.7%** |
| 128k   | F | 1398 | 1492 | **+6.7%** |
| 4k     | T | 1215 | 1217 | +0.2% |
| 8k     | T | 1411 | 1415 | +0.3% |
| 16k    | T | 1540 | 1545 | +0.3% |
| 32k    | T | 1552 | 1615 | **+4.1%** |
| 64k    | T | 1486 | 1496 | +0.7% |
| 128k   | T | 1363 | 1370 | +0.5% |

#### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1496 | 1511 | +1.0% |
| 8k     | F | 1601 | 1612 | +0.7% |
| 16k    | F | 1620 | 1653 | **+2.0%** |
| 32k    | F | 1482 | 1464 | −1.2% |
| 64k    | F | 1436 | 1472 | **+2.5%** |
| 128k   | F | 1389 | 1469 | **+5.8%** |
| 4k     | T | 1242 | 1244 | +0.2% |
| 8k     | T | 1434 | 1439 | +0.3% |
| 16k    | T | 1556 | 1562 | +0.4% |
| 32k    | T | 1624 | 1564 | −3.7% |
| 64k    | T | 1493 | 1507 | +0.9% |
| 128k   | T | 1373 | 1368 | −0.4% |

Aggregated means: **MHA fwd +1.8%, GQA fwd +0.7%**.

## Caveat — unexpected long-seqlen non-causal speedup

The intent of this change is correctness only (adds a ternary, a
`continue_cond` extension, moves a variable declaration). None of
these should improve inner-loop throughput.

Yet we observe a **reproducible** +4–7% at long-seqlen non-causal
(MHA 128k F +6.7%, GQA 128k F +5.8%, MHA 64k F +4.7%, MHA 32k T
+4.1%). 3-run variance per cell is tight (typically <1%), so this is
not run-to-run noise. Bracketed against the parent `paged-kv-v2` PR,
which measured within ±0.3% of main on the same cells, the gain is
introduced specifically by this commit's source-level reorderings
(likely register-allocation or instruction-scheduling artifacts from
ptxas).

### Why this is a concern, not just free perf

1. **Unintended change** — we can't explain it from the diff, which
   means the compiler's decision hinges on something fragile (variable
   declaration order, kernel signature shape). A future unrelated edit
   could flip this back to `±0%`, or worse, regress it.
2. **Could mask a different regression.** If the reordering also
   subtly changed some other code path we don't benchmark (e.g. paged
   path with `seqused_k = None`), we wouldn't notice until production.
3. **Not portable guidance.** We can't tell future contributors
   "move `batch_coord` earlier to get +6%" because the mechanism isn't
   a deliberate optimization.

### TODO before opening / merging

- [ ] Compare SASS between this branch and `paged-kv-v2` for the
      long-seqlen non-causal dense path; identify which instructions
      changed and whether the gain is attributable to a specific
      scheduling/allocation difference.
- [ ] Confirm the paged path with `seqused_k = None` isn't regressed
      (benchmark_attn.py doesn't exercise paged; add a quick
      paged-bench harness or run the paged tests under timing).
- [ ] Decide whether to keep the reorderings (if the mechanism is
      understood) or revert the non-essential ones (declaration move)
      to isolate the correctness change from the accidental perf gain.

## Other caveat

- Host-side prefill in `interface.py` walks zero-KV batch indices in
  Python. For batch sizes in the thousands with many zero-KV rows
  this could show up in CPU profiles; current FA4 callers aren't in
  that regime.
- `seqused_q` is still asserted `None` — the kernel would need a
  similar ternary in the Q-side iteration, out of scope for this PR.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `dbb6c98` from `Johnsonms/persistent-cluster-hd256`
on top of `Johnsonms/seqused-k-hd256-v2`. Original branch was based on a
pre-merge snapshot of the hd256 tree; base commits were absorbed into
the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487
(`b21e204`).

**Stacked PR** — base branch should be `Johnsonms/seqused-k-hd256-v2`,
not `main`. Depends on that PR for the `seqused_k` + paged-KV plumbing.

## Change

Persistent scheduling amortizes CTA launch overhead by issuing a
grid-stride loop over tiles. It was hardcoded off in hd256 since the
kernel's inception because the static tile scheduler was cluster-
unaware and split 2CTA clusters across independent work tiles,
corrupting output.

### Cluster-aware fix

- `Sm100FmhaStaticTileSchedulerParams` gains a `cluster_shape_m`
  constexpr (default 1, so 1CTA kernels are unchanged).
- Grid is sized in cluster units:
  `max_ctas = (sm_count // cluster_shape_m) * cluster_shape_m`;
  problem size is multiplied by `cluster_shape_m` so `dsl_min`
  compares apples to apples.
- `num_persistent_clusters` replaces `num_persistent_sm` as the
  grid-stride step, so `advance_to_next_work` advances by one cluster
  per iteration instead of one CTA.
- In `get_current_work`, the CTA rank within the cluster is
  reconstructed from the launch `block_idx`
  (`cta_rank = blk_coord[0] % cluster_shape_m`) and spliced back into
  `mid = m_block * cluster_shape_m + cta_rank`, so both CTAs in a 2CTA
  cluster land on their half of the same tile.

### Enablement gate (work-per-tile heuristic)

Persistent's per-tile cost (cluster-rank reconstruction + grid-stride
state) only pays off when work-per-tile is small, i.e. short KV. On
B200, persistent wins at short seqlen but regresses dense prefill at
long seqlen because launch overhead is already amortized by the
hardware at high tile counts.

- **Gate:** `... and seqlen_k <= 2048`. Keeps the decode win and
  holds long-context prefill within noise of the pre-persistent
  baseline.
- Gate is folded into `interface.py`'s compile_key so configs
  straddling the threshold compile separate kernels with the correct
  persistent flag baked in.

### Files

- `flash_attn/cute/tile_scheduler.py`: add `cluster_shape_m` on
  `Sm100FmhaStaticTileSchedulerParams` and the matching updates to
  `Sm100FmhaStaticTileScheduler` and `compute_sm100_fmha_grid`.
- `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`: drop
  `self.is_persistent = False` hardcode (now honors constructor arg),
  pass `cluster_shape_mnk` to `compute_grid`, divide `blk_idx[0]` by
  `cluster_shape_mnk[0]` when constructing `FmhaStaticTileScheduler`.
- `flash_attn/cute/sm100_hd256_2cta_fmha_backward_dqkernel.py`:
  same fix applied to the bwd dQ kernel for consistency (its
  scheduler is the same `Sm100` static scheduler).
- `flash_attn/cute/interface.py`: compute `hd256_is_persistent` above
  the compile_key (gated on `seqlen_k <= 2048` in addition to
  causal/cu_seqlens_q/...); add it to the compile_key; use it at
  `fa_fwd` construction.

## Validation

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 `seqused_k` tests (inherited from parent PR), the 6 paged tests
(grandparent PR), and the existing d=256 dense subset:

```
-k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or
    seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or
    (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **90 passed, 78 skipped, 0 failed** — identical pass/skip
count to `origin/main`. No new tests added by this commit (scheduler
refactor tests are covered via the existing d=256 dense parametrize
space, which hits both gate-ON and gate-OFF configurations).

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

#### Short seqlen — **gate-ON path** (persistent scheduling active)

##### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 1k     | F | 985  |  987 | +0.2% |
| 2k     | F | 1307 | 1301 | −0.5% |
| 1k     | T | 592  |  592 |  0.0% |
| 2k     | T | 910  |  911 | +0.1% |

##### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 1k     | F | 1012 | 1015 | +0.3% |
| 2k     | F | 1329 | 1329 |  0.0% |
| 1k     | T |  600 |  600 |  0.0% |
| 2k     | T |  919 |  919 |  0.0% |

**Observation:** the persistent path is active but delivers essentially
**no speedup** in this bench config (32 Q-heads MHA, 32:2 GQA). The
commit message cites +10–25% wins at `seqlen_k=1024` on a different
config (8 Q-heads); with 32 Q-heads the tile count per batch is 4×
larger and launch overhead is already amortized by the hardware, so the
persistent loop's benefit is saturated out. Reproducibility is very
tight (variance <1 TFLOPS across 3 runs), so this is not run noise —
just "no regression" rather than "a win" at this config.

#### Long seqlen — **gate-OFF path** (should match main)

##### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1474 | 1479 | +0.3% |
| 8k     | F | 1582 | 1588 | +0.4% |
| 16k    | F | 1641 | 1636 | −0.3% |
| 32k    | F | 1450 | 1442 | −0.6% |
| 64k    | F | 1417 | 1447 | **+2.1%** |
| 128k   | F | 1398 | 1423 | +1.8% |
| 4k     | T | 1215 | 1218 | +0.2% |
| 8k     | T | 1411 | 1416 | +0.4% |
| 16k    | T | 1540 | 1540 |  0.0% |
| 32k    | T | 1552 | 1619 | **+4.3%** |
| 64k    | T | 1486 | 1469 | −1.1% |
| 128k   | T | 1363 | 1383 | +1.5% |

##### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1496 | 1504 | +0.5% |
| 8k     | F | 1601 | 1598 | −0.2% |
| 16k    | F | 1620 | 1640 | +1.2% |
| 32k    | F | 1482 | 1539 | **+3.8%** |
| 64k    | F | 1436 | 1460 | +1.7% |
| 128k   | F | 1389 | 1363 | −1.9% |
| 4k     | T | 1242 | 1244 | +0.2% |
| 8k     | T | 1434 | 1438 | +0.3% |
| 16k    | T | 1556 | 1557 | +0.1% |
| 32k    | T | 1624 | 1574 | **−3.1%** |
| 64k    | T | 1493 | 1493 |  0.0% |
| 128k   | T | 1373 | 1369 | −0.3% |

**Observation:** long-seqlen path is within ±2% on most cells. Same
long-seqlen noise pattern as the parent PRs at 32k/64k (batch-
quantization zone); deltas swing both directions, no systematic
regression.

## Caveats

- **Short-seqlen benefit is config-dependent.** The persistent path
  trades extra per-tile state for launch-overhead amortization; the
  trade only pays off when work-per-tile is small. At 8 Q-heads the
  commit's original bench showed +10–25% at 1k; at 32 Q-heads in my
  bench the benefit collapses to 0% because tile count is already
  high enough to amortize launches. If the use case is decode-style
  small-batch / few-head, the gate-ON path is expected to deliver the
  advertised win.
- Backward dQ kernel gets the same cluster-aware scheduler change
  for consistency but is not exercised by `benchmark_attn.py` in this
  PR's validation. Forward dense path is the regression-critical one
  and is covered above.
- `cluster_shape_m` currently defaults to 1 so 1CTA kernels are
  unaffected; the only callers passing `cluster_shape_m > 1` are
  hd256 forward and hd256 bwd dQ.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `dbb6c98` from `Johnsonms/persistent-cluster-hd256`
on top of `Johnsonms/seqused-k-hd256-v2`. Original branch was based on a
pre-merge snapshot of the hd256 tree; base commits were absorbed into
the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487
(`b21e204`).

**Stacked PR** — base branch should be `Johnsonms/seqused-k-hd256-v2`,
not `main`. Depends on that PR for the `seqused_k` + paged-KV plumbing.

## Change

Persistent scheduling amortizes CTA launch overhead by issuing a
grid-stride loop over tiles. It was hardcoded off in hd256 since the
kernel's inception because the static tile scheduler was cluster-
unaware and split 2CTA clusters across independent work tiles,
corrupting output.

### Cluster-aware fix

- `Sm100FmhaStaticTileSchedulerParams` gains a `cluster_shape_m`
  constexpr (default 1, so 1CTA kernels are unchanged).
- Grid is sized in cluster units:
  `max_ctas = (sm_count // cluster_shape_m) * cluster_shape_m`;
  problem size is multiplied by `cluster_shape_m` so `dsl_min`
  compares apples to apples.
- `num_persistent_clusters` replaces `num_persistent_sm` as the
  grid-stride step, so `advance_to_next_work` advances by one cluster
  per iteration instead of one CTA.
- In `get_current_work`, the CTA rank within the cluster is
  reconstructed from the launch `block_idx`
  (`cta_rank = blk_coord[0] % cluster_shape_m`) and spliced back into
  `mid = m_block * cluster_shape_m + cta_rank`, so both CTAs in a 2CTA
  cluster land on their half of the same tile.

### Enablement gate (work-per-tile heuristic)

Persistent's per-tile cost (cluster-rank reconstruction + grid-stride
state) only pays off when work-per-tile is small, i.e. short KV. On
B200, persistent wins at short seqlen but regresses dense prefill at
long seqlen because launch overhead is already amortized by the
hardware at high tile counts.

- **Gate:** `... and seqlen_k <= 2048`. Keeps the decode win and
  holds long-context prefill within noise of the pre-persistent
  baseline.
- Gate is folded into `interface.py`'s compile_key so configs
  straddling the threshold compile separate kernels with the correct
  persistent flag baked in.

### Files

- `flash_attn/cute/tile_scheduler.py`: add `cluster_shape_m` on
  `Sm100FmhaStaticTileSchedulerParams` and the matching updates to
  `Sm100FmhaStaticTileScheduler` and `compute_sm100_fmha_grid`.
- `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`: drop
  `self.is_persistent = False` hardcode (now honors constructor arg),
  pass `cluster_shape_mnk` to `compute_grid`, divide `blk_idx[0]` by
  `cluster_shape_mnk[0]` when constructing `FmhaStaticTileScheduler`.
- `flash_attn/cute/sm100_hd256_2cta_fmha_backward_dqkernel.py`:
  same fix applied to the bwd dQ kernel for consistency (its
  scheduler is the same `Sm100` static scheduler).
- `flash_attn/cute/interface.py`: compute `hd256_is_persistent` above
  the compile_key (gated on `seqlen_k <= 2048` in addition to
  causal/cu_seqlens_q/...); add it to the compile_key; use it at
  `fa_fwd` construction.

## Validation

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 `seqused_k` tests (inherited from parent PR), the 6 paged tests
(grandparent PR), and the existing d=256 dense subset:

```
-k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or
    seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or
    (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **90 passed, 78 skipped, 0 failed** — identical pass/skip
count to `origin/main`. No new tests added by this commit (scheduler
refactor tests are covered via the existing d=256 dense parametrize
space, which hits both gate-ON and gate-OFF configurations).

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

#### Short seqlen — **gate-ON path** (persistent scheduling active)

##### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 1k     | F | 985  |  987 | +0.2% |
| 2k     | F | 1307 | 1301 | −0.5% |
| 1k     | T | 592  |  592 |  0.0% |
| 2k     | T | 910  |  911 | +0.1% |

##### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 1k     | F | 1012 | 1015 | +0.3% |
| 2k     | F | 1329 | 1329 |  0.0% |
| 1k     | T |  600 |  600 |  0.0% |
| 2k     | T |  919 |  919 |  0.0% |

**Observation:** the persistent path is active but delivers essentially
**no speedup** in this bench config (32 Q-heads MHA, 32:2 GQA). The
commit message cites +10–25% wins at `seqlen_k=1024` on a different
config (8 Q-heads); with 32 Q-heads the tile count per batch is 4×
larger and launch overhead is already amortized by the hardware, so the
persistent loop's benefit is saturated out. Reproducibility is very
tight (variance <1 TFLOPS across 3 runs), so this is not run noise —
just "no regression" rather than "a win" at this config.

#### Long seqlen — **gate-OFF path** (should match main)

##### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1474 | 1479 | +0.3% |
| 8k     | F | 1582 | 1588 | +0.4% |
| 16k    | F | 1641 | 1636 | −0.3% |
| 32k    | F | 1450 | 1442 | −0.6% |
| 64k    | F | 1417 | 1447 | **+2.1%** |
| 128k   | F | 1398 | 1423 | +1.8% |
| 4k     | T | 1215 | 1218 | +0.2% |
| 8k     | T | 1411 | 1416 | +0.4% |
| 16k    | T | 1540 | 1540 |  0.0% |
| 32k    | T | 1552 | 1619 | **+4.3%** |
| 64k    | T | 1486 | 1469 | −1.1% |
| 128k   | T | 1363 | 1383 | +1.5% |

##### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1496 | 1504 | +0.5% |
| 8k     | F | 1601 | 1598 | −0.2% |
| 16k    | F | 1620 | 1640 | +1.2% |
| 32k    | F | 1482 | 1539 | **+3.8%** |
| 64k    | F | 1436 | 1460 | +1.7% |
| 128k   | F | 1389 | 1363 | −1.9% |
| 4k     | T | 1242 | 1244 | +0.2% |
| 8k     | T | 1434 | 1438 | +0.3% |
| 16k    | T | 1556 | 1557 | +0.1% |
| 32k    | T | 1624 | 1574 | **−3.1%** |
| 64k    | T | 1493 | 1493 |  0.0% |
| 128k   | T | 1373 | 1369 | −0.3% |

**Observation:** long-seqlen path is within ±2% on most cells. Same
long-seqlen noise pattern as the parent PRs at 32k/64k (batch-
quantization zone); deltas swing both directions, no systematic
regression.

## Caveats

- **Short-seqlen benefit is config-dependent.** The persistent path
  trades extra per-tile state for launch-overhead amortization; the
  trade only pays off when work-per-tile is small. At 8 Q-heads the
  commit's original bench showed +10–25% at 1k; at 32 Q-heads in my
  bench the benefit collapses to 0% because tile count is already
  high enough to amortize launches. If the use case is decode-style
  small-batch / few-head, the gate-ON path is expected to deliver the
  advertised win.
- Backward dQ kernel gets the same cluster-aware scheduler change
  for consistency but is not exercised by `benchmark_attn.py` in this
  PR's validation. Forward dense path is the regression-critical one
  and is covered above.
- `cluster_shape_m` currently defaults to 1 so 1CTA kernels are
  unaffected; the only callers passing `cluster_shape_m > 1` are
  hd256 forward and hd256 bwd dQ.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `94e63db` from `Johnsonms/seqused-k-hd256` on
top of `Johnsonms/paged-kv-hd256-v2`. Original branch was based on a
pre-merge snapshot of the hd256 tree; base commits were absorbed into
the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487
(`b21e204`).

**Stacked PR** — base branch should be `Johnsonms/paged-kv-hd256-v2`,
not `main`. Depends on that PR for the paged-KV kernel plumbing.

## Change

Enables variable per-batch KV sequence lengths via a `seqused_k`
tensor — needed for MLA-style decode (DeepSeek-V2 / V3 / R1), where
different batches have different KV cache occupancies. Works with
both dense and paged K/V.

### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`

- Drop the `mSeqUsedK` half of the `__call__` assertion; only
  `mSeqUsedQ` stays blocked for now.
- Build the `seqused_k` cute tensor in `__call__` alongside the
  `page_table` / paged-layout construction.
- Add `mSeqUsedK` to `kernel()` signature and pass `seqused_k` from
  `__call__`.
- Replace `seqlen_k` derivation in all 4 warp sections with a ternary:
  `mSeqUsedK[batch_coord] if set, else <dense/paged expression>`.
- Move `batch_coord` above `seqlen_k` in the MMA warp (second warp
  section) — it was declared later but now needs to be in scope for
  `seqused_k` indexing.
- **Zero-KV batch handling (`seqlen_k == 0`):** extend `continue_cond`
  in all four warp sections with `continue_cond or seqlen_k <= 0`, so
  load / MMA / correction / softmax warps skip in sync instead of
  deadlocking on `K0 / Vend / QK0 / PVend / first-stats` tiles.

### `flash_attn/cute/interface.py`

- Relax the `seqused_q is None and seqused_k is None` assertion to
  `seqused_q is None` (the kernel now handles `seqused_k`).
- Prefill zero-KV batches on the host with zero output and `-inf`
  LSE, so their output is defined even though the kernel skips them.

### `tests/cute/test_flash_attn.py`

- `test_flash_attn_seqused_k_hd256_sm100`: dense + `seqused_k`
  (padded K/V with per-batch valid lengths) bit-exact vs a
  `cu_seqlens_k` packed reference, parametrized over asymmetric
  per-batch lengths.
- `test_flash_attn_paged_seqused_k_hd256_sm100`: paged + `seqused_k`
  combined (MLA decode pattern), bit-exact vs packed reference.
- `test_flash_attn_seqused_k_zero_hd256_sm100`: `seqused_k = 0` for
  one batch, parametrized dense/paged. Verifies no deadlock, zero
  output, `-inf` LSE on the empty row, and finite output / LSE on
  the other row.

## Validation

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 new `seqused_k` tests, the 6 paged tests from the parent PR, and
the existing d=256 dense subset:

```
-k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or
    seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or
    (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **90 passed, 78 skipped, 0 failed** — 78 from the dense d=256
subset (identical pass/skip count to `origin/main`), 6 from the parent
PR's `paged_hd256_sm100_tma[_gqa]` tests, and **6 from the 3 new
`seqused_k_*_hd256_sm100` tests** introduced here.

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

#### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1474 | 1485 | +0.7% |
| 8k     | F | 1582 | 1595 | +0.8% |
| 16k    | F | 1641 | 1650 | +0.5% |
| 32k    | F | 1450 | 1488 | **+2.6%** |
| 64k    | F | 1417 | 1484 | **+4.7%** |
| 128k   | F | 1398 | 1492 | **+6.7%** |
| 4k     | T | 1215 | 1217 | +0.2% |
| 8k     | T | 1411 | 1415 | +0.3% |
| 16k    | T | 1540 | 1545 | +0.3% |
| 32k    | T | 1552 | 1615 | **+4.1%** |
| 64k    | T | 1486 | 1496 | +0.7% |
| 128k   | T | 1363 | 1370 | +0.5% |

#### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1496 | 1511 | +1.0% |
| 8k     | F | 1601 | 1612 | +0.7% |
| 16k    | F | 1620 | 1653 | **+2.0%** |
| 32k    | F | 1482 | 1464 | −1.2% |
| 64k    | F | 1436 | 1472 | **+2.5%** |
| 128k   | F | 1389 | 1469 | **+5.8%** |
| 4k     | T | 1242 | 1244 | +0.2% |
| 8k     | T | 1434 | 1439 | +0.3% |
| 16k    | T | 1556 | 1562 | +0.4% |
| 32k    | T | 1624 | 1564 | −3.7% |
| 64k    | T | 1493 | 1507 | +0.9% |
| 128k   | T | 1373 | 1368 | −0.4% |

Aggregated means: **MHA fwd +1.8%, GQA fwd +0.7%**.

## Caveat — unexpected long-seqlen non-causal speedup

The intent of this change is correctness only (adds a ternary, a
`continue_cond` extension, moves a variable declaration). None of
these should improve inner-loop throughput.

Yet we observe a **reproducible** +4–7% at long-seqlen non-causal
(MHA 128k F +6.7%, GQA 128k F +5.8%, MHA 64k F +4.7%, MHA 32k T
+4.1%). 3-run variance per cell is tight (typically <1%), so this is
not run-to-run noise. Bracketed against the parent `paged-kv-v2` PR,
which measured within ±0.3% of main on the same cells, the gain is
introduced specifically by this commit's source-level reorderings
(likely register-allocation or instruction-scheduling artifacts from
ptxas).

### Why this is a concern, not just free perf

1. **Unintended change** — we can't explain it from the diff, which
   means the compiler's decision hinges on something fragile (variable
   declaration order, kernel signature shape). A future unrelated edit
   could flip this back to `±0%`, or worse, regress it.
2. **Could mask a different regression.** If the reordering also
   subtly changed some other code path we don't benchmark (e.g. paged
   path with `seqused_k = None`), we wouldn't notice until production.
3. **Not portable guidance.** We can't tell future contributors
   "move `batch_coord` earlier to get +6%" because the mechanism isn't
   a deliberate optimization.

### TODO before opening / merging

- [ ] Compare SASS between this branch and `paged-kv-v2` for the
      long-seqlen non-causal dense path; identify which instructions
      changed and whether the gain is attributable to a specific
      scheduling/allocation difference.
- [ ] Confirm the paged path with `seqused_k = None` isn't regressed
      (benchmark_attn.py doesn't exercise paged; add a quick
      paged-bench harness or run the paged tests under timing).
- [ ] Decide whether to keep the reorderings (if the mechanism is
      understood) or revert the non-essential ones (declaration move)
      to isolate the correctness change from the accidental perf gain.

## Other caveat

- Host-side prefill in `interface.py` walks zero-KV batch indices in
  Python. For batch sizes in the thousands with many zero-KV rows
  this could show up in CPU profiles; current FA4 callers aren't in
  that regime.
- `seqused_q` is still asserted `None` — the kernel would need a
  similar ternary in the Q-side iteration, out of scope for this PR.
Johnsonms added a commit to Johnsonms/flash-attention that referenced this pull request Apr 23, 2026
Rebased cherry-pick of `dbb6c98` from `Johnsonms/persistent-cluster-hd256`
on top of `Johnsonms/seqused-k-hd256-v2`. Original branch was based on a
pre-merge snapshot of the hd256 tree; base commits were absorbed into
the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487
(`b21e204`).

**Stacked PR** — base branch should be `Johnsonms/seqused-k-hd256-v2`,
not `main`. Depends on that PR for the `seqused_k` + paged-KV plumbing.

## Change

Persistent scheduling amortizes CTA launch overhead by issuing a
grid-stride loop over tiles. It was hardcoded off in hd256 since the
kernel's inception because the static tile scheduler was cluster-
unaware and split 2CTA clusters across independent work tiles,
corrupting output.

### Cluster-aware fix

- `Sm100FmhaStaticTileSchedulerParams` gains a `cluster_shape_m`
  constexpr (default 1, so 1CTA kernels are unchanged).
- Grid is sized in cluster units:
  `max_ctas = (sm_count // cluster_shape_m) * cluster_shape_m`;
  problem size is multiplied by `cluster_shape_m` so `dsl_min`
  compares apples to apples.
- `num_persistent_clusters` replaces `num_persistent_sm` as the
  grid-stride step, so `advance_to_next_work` advances by one cluster
  per iteration instead of one CTA.
- In `get_current_work`, the CTA rank within the cluster is
  reconstructed from the launch `block_idx`
  (`cta_rank = blk_coord[0] % cluster_shape_m`) and spliced back into
  `mid = m_block * cluster_shape_m + cta_rank`, so both CTAs in a 2CTA
  cluster land on their half of the same tile.

### Enablement gate (work-per-tile heuristic)

Persistent's per-tile cost (cluster-rank reconstruction + grid-stride
state) only pays off when work-per-tile is small, i.e. short KV. On
B200, persistent wins at short seqlen but regresses dense prefill at
long seqlen because launch overhead is already amortized by the
hardware at high tile counts.

- **Gate:** `... and seqlen_k <= 2048`. Keeps the decode win and
  holds long-context prefill within noise of the pre-persistent
  baseline.
- Gate is folded into `interface.py`'s compile_key so configs
  straddling the threshold compile separate kernels with the correct
  persistent flag baked in.

### Files

- `flash_attn/cute/tile_scheduler.py`: add `cluster_shape_m` on
  `Sm100FmhaStaticTileSchedulerParams` and the matching updates to
  `Sm100FmhaStaticTileScheduler` and `compute_sm100_fmha_grid`.
- `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`: drop
  `self.is_persistent = False` hardcode (now honors constructor arg),
  pass `cluster_shape_mnk` to `compute_grid`, divide `blk_idx[0]` by
  `cluster_shape_mnk[0]` when constructing `FmhaStaticTileScheduler`.
- `flash_attn/cute/sm100_hd256_2cta_fmha_backward_dqkernel.py`:
  same fix applied to the bwd dQ kernel for consistency (its
  scheduler is the same `Sm100` static scheduler).
- `flash_attn/cute/interface.py`: compute `hd256_is_persistent` above
  the compile_key (gated on `seqlen_k <= 2048` in addition to
  causal/cu_seqlens_q/...); add it to the compile_key; use it at
  `fa_fwd` construction.

## Validation

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 `seqused_k` tests (inherited from parent PR), the 6 paged tests
(grandparent PR), and the existing d=256 dense subset:

```
-k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or
    seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or
    (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **90 passed, 78 skipped, 0 failed** — identical pass/skip
count to `origin/main`. No new tests added by this commit (scheduler
refactor tests are covered via the existing d=256 dense parametrize
space, which hits both gate-ON and gate-OFF configurations).

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

#### Short seqlen — **gate-ON path** (persistent scheduling active)

##### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 1k     | F | 985  |  987 | +0.2% |
| 2k     | F | 1307 | 1301 | −0.5% |
| 1k     | T | 592  |  592 |  0.0% |
| 2k     | T | 910  |  911 | +0.1% |

##### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 1k     | F | 1012 | 1015 | +0.3% |
| 2k     | F | 1329 | 1329 |  0.0% |
| 1k     | T |  600 |  600 |  0.0% |
| 2k     | T |  919 |  919 |  0.0% |

**Observation:** the persistent path is active but delivers essentially
**no speedup** in this bench config (32 Q-heads MHA, 32:2 GQA). The
commit message cites +10–25% wins at `seqlen_k=1024` on a different
config (8 Q-heads); with 32 Q-heads the tile count per batch is 4×
larger and launch overhead is already amortized by the hardware, so the
persistent loop's benefit is saturated out. Reproducibility is very
tight (variance <1 TFLOPS across 3 runs), so this is not run noise —
just "no regression" rather than "a win" at this config.

#### Long seqlen — **gate-OFF path** (should match main)

##### MHA 32:32

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1474 | 1479 | +0.3% |
| 8k     | F | 1582 | 1588 | +0.4% |
| 16k    | F | 1641 | 1636 | −0.3% |
| 32k    | F | 1450 | 1442 | −0.6% |
| 64k    | F | 1417 | 1447 | **+2.1%** |
| 128k   | F | 1398 | 1423 | +1.8% |
| 4k     | T | 1215 | 1218 | +0.2% |
| 8k     | T | 1411 | 1416 | +0.4% |
| 16k    | T | 1540 | 1540 |  0.0% |
| 32k    | T | 1552 | 1619 | **+4.3%** |
| 64k    | T | 1486 | 1469 | −1.1% |
| 128k   | T | 1363 | 1383 | +1.5% |

##### GQA 32:2

| seqlen | causal | main | this PR | Δ |
|-------:|:------:|-----:|--------:|--:|
| 4k     | F | 1496 | 1504 | +0.5% |
| 8k     | F | 1601 | 1598 | −0.2% |
| 16k    | F | 1620 | 1640 | +1.2% |
| 32k    | F | 1482 | 1539 | **+3.8%** |
| 64k    | F | 1436 | 1460 | +1.7% |
| 128k   | F | 1389 | 1363 | −1.9% |
| 4k     | T | 1242 | 1244 | +0.2% |
| 8k     | T | 1434 | 1438 | +0.3% |
| 16k    | T | 1556 | 1557 | +0.1% |
| 32k    | T | 1624 | 1574 | **−3.1%** |
| 64k    | T | 1493 | 1493 |  0.0% |
| 128k   | T | 1373 | 1369 | −0.3% |

**Observation:** long-seqlen path is within ±2% on most cells. Same
long-seqlen noise pattern as the parent PRs at 32k/64k (batch-
quantization zone); deltas swing both directions, no systematic
regression.

## Caveats

- **Short-seqlen benefit is config-dependent.** The persistent path
  trades extra per-tile state for launch-overhead amortization; the
  trade only pays off when work-per-tile is small. At 8 Q-heads the
  commit's original bench showed +10–25% at 1k; at 32 Q-heads in my
  bench the benefit collapses to 0% because tile count is already
  high enough to amortize launches. If the use case is decode-style
  small-batch / few-head, the gate-ON path is expected to deliver the
  advertised win.
- Backward dQ kernel gets the same cluster-aware scheduler change
  for consistency but is not exercised by `benchmark_attn.py` in this
  PR's validation. Forward dense path is the regression-critical one
  and is covered above.
- `cluster_shape_m` currently defaults to 1 so 1CTA kernels are
  unaffected; the only callers passing `cluster_shape_m > 1` are
  hd256 forward and hd256 bwd dQ.
@wangsiyu

Copy link
Copy Markdown
Contributor Author

Hi @wangsiyu @cherichy @dishengbin, this PR is approved and will be merged with the following follow-up items:

  1. The current implementation follows a CuTe DSL example style. It should be refactored to align with the FlashAttention codebase style (using flash_bwd_sm100.py as the reference) within the next four weeks.
  2. Performance is a critical aspect of this feature. During the refactor and any future enhancements, we should ensure that high FLOPS performance is maintained and continue exploring additional optimizations where appropriate.

Cc: @tridao @jayhshah @drisspg

Both fwd and bwd will be refactored.

Johnsonms added a commit that referenced this pull request Apr 28, 2026
…rformance gain) (#2488)

* [hd256] Improve forward kernel with exp2 FMA emulation

Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top
of merged main (hd256 PR #2412, `27b4eb9`). The original branch was
based on a pre-merge snapshot; the other five commits in that branch
were absorbed into the squash-merge, leaving this one novel change.

## Change

Replace a fraction of hardware `exp2` (SFU) instructions with a
polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed
FMA polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write-slot acquisition is moved earlier to overlap
any pipeline stall with the `exp2` compute.

Kernel-only change; no API change. Backward is untouched.

## Validation (this PR vs `origin/main` @ `b21e204` — includes #2412 hd256 base and #2487 post-merge cleanup)

B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked
clocks @ 1755 MHz, seqlens 4k..128k.

### FWD delta vs `origin/main` (TFLOPS mean, 3 runs each)

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  -0.2%    |  +0.7%   |
| 8k     |   F    |  +0.4%    |  +0.3%   |
| 16k    |   F    |  +0.7%    |  -1.0%   |
| 32k    |   F    |  +2.3%    |  -0.3%   |
| 64k    |   F    |  **+5.4%** |  **+5.2%** |
| 128k   |   F    |  **+7.4%** |  **+7.3%** |
| 4k     |   T    |  +0.3%    |  +0.9%   |
| 8k     |   T    |  +0.8%    |  +0.9%   |
| 16k    |   T    |  +0.8%    |  +1.1%   |
| 32k    |   T    |  **+5.2%** |  -0.4%   |
| 64k    |   T    |  +1.1%    |  +1.1%   |
| 128k   |   T    |  **+3.4%** |  **+2.6%** |

- 19 of 24 cells positive; 4 slightly negative, all within the
  batch-quantization noise band that `origin/main` itself already
  showed in our 3-run regression sweep.
- Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax
  SFU pressure is worst, consistent with the theory above.
- Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**.

### Correctness smoke

`pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k
"256-False-0-0.0-False-False"` on B200:
**78 passed, 78 skipped, 0 failed** — identical pass/skip count to
`origin/main`.

## Caveat

Exp2 FMA emulation introduces small numerical differences vs hardware
`exp2`. The existing test tolerances accept the delta.

* [hd256] Wire ex2_emu params through _TUNING_CONFIG with tuned values

The exp2 emulation knobs (ex2_emu_freq, ex2_emu_res, ex2_emu_start_frg)
and softmax register counts for the hd256 forward kernel were hardcoded
in BlackwellFusedMultiHeadAttentionForward.__init__, invisible to the
central _TUNING_CONFIG table used by all other kernel configs.

- flash_fwd_sm100.py: add hd256 entries to _TUNING_CONFIG (causal and
  non-causal; always 2cta, no sm103 variant). New ex2_emu_res field is
  hd256-specific; existing entries are unaffected. hd256 uses a fixed
  num_regs_other=32 (not derived from the 512-budget formula).
- sm100_hd256_2cta_fmha_forward.py: replace hardcoded self.* assignments
  with a _TUNING_CONFIG lookup.

Tuned values (B200, bf16, locked clocks): freq=14, res=6, start_frg=0
for both causal and non-causal. The inner loop steps k by 2, so k%freq
only takes even values; freq=14/res=6 gives ~43% emulation (3 out of 7
even k%14 steps), replacing the previous 50:50 split (freq=4/res=3).
Johnsonms added a commit that referenced this pull request May 1, 2026
* [hd256] Add TMA paged KV support to SM100 2CTA forward kernel

Rebased cherry-pick of `49fe257` from `Johnsonms/paged-kv-hd256` on top
of merged main (hd256 PR #2412, `27b4eb9` + post-merge cleanup #2487,
`b21e204`). Original branch was based on a pre-merge snapshot; its
base commits were absorbed into the squash-merge.

## Change

Adds paged KV support to the SM100 hd256 2CTA forward kernel. The paged
path reuses the dense TMA load path — logical KV blocks are remapped to
physical page indices through the page table at load time, so each page
maps to exactly one TMA tile.

**Constraint:** `page_size` must equal `tile_n = 128`.

### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`

- Conditional K/V tensor layout in `__call__`: dense
  `(s_k, d, ((h_r, h_k), b))` vs paged
  `(page_size, d, h_k, num_pages)` for K (and transposed for V).
- Conditional K/V TMA setup in the load warp: dense uses
  `domain_offset` + batch indexing; paged uses `head_kv` slicing and
  keeps `num_pages` as the outer mode for per-load `page_idx` lookup.
- Conditional per-load `page_idx`: K uses mode-2 subtile + mode-3 page;
  V uses mode-1 page.
- Plumb `mPageTable` + `max_seqlen_k` through the kernel signature.
  `seqlen_k` in each of the 4 warp sections now uses `max_seqlen_k`
  for the paged path.
- Store `qhead_per_kvhead` on `self` and derive `head_kv_coord` via
  integer divide (matches the `flash_fwd_sm100` convention for
  contiguous GQA grouping).
- Relax `mPageTable` / `paged_kv_non_tma` assertions.

### `flash_attn/cute/paged_kv.py`

- Extract `_flatten_smem_sm100` / `_copy_row_async` helpers from
  `load_KV` — pure refactor, no behavior change for existing callers.

### `tests/cute/test_flash_attn.py`

- `test_flash_attn_paged_hd256_sm100_tma`: bit-exact vs dense varlen
  reference + determinism check, parametrized over `seqlen_q`.
- `test_flash_attn_paged_hd256_sm100_tma_gqa`: same check for GQA with
  `nheads_kv in {2, 4, 8}` — exercises `qhead_per_kvhead > 1`, which
  a modulo-aliasing bug would fail.

## Validation (this PR vs `origin/main` @ `b21e204`)

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 new paged tests with the existing d=256 dense subset:

```
-k "paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **84 passed, 78 skipped, 0 failed** in 2 min — 78 from the
dense d=256 subset (identical pass/skip count to `origin/main`) and
**6 from the new `paged_hd256_sm100_tma[_gqa]` tests**.

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  +0.2%    |  +0.3%   |
| 8k     |   F    |   0.0%    |  -0.1%   |
| 16k    |   F    |  -0.2%    |   0.0%   |
| 32k    |   F    |  -1.0%    |  **+2.3%** |
| 64k    |   F    |  **+2.1%** |  +0.8%  |
| 128k   |   F    |  -0.8%    |  -1.7%   |
| 4k     |   T    |  +0.2%    |  +0.2%   |
| 8k     |   T    |  +0.1%    |  +0.1%   |
| 16k    |   T    |  +0.3%    |  +0.1%   |
| 32k    |   T    |  **+3.0%** |  -1.2%  |
| 64k    |   T    |  -0.3%    |  -0.1%   |
| 128k   |   T    |  +0.1%    |  -0.7%   |

- **22 of 24 cells within ±2%.**
- Two `> 2%` outliers are both **positive** and in the batch-
  quantization noise zone that `origin/main` itself showed run-spread
  in during our 3-run baseline sweep — not regressions.
- **Aggregated means: MHA +0.31%, GQA +0.00%.**
- Paged-KV path isn't exercised by `benchmark_attn.py` (which uses
  contiguous KV); dense-path perf parity is the regression-critical
  property and is preserved.

## Caveat

- **page_size == tile_n == 128 is a hard constraint.** Callers that
  want a different page size will need a separate path.
- The paged-KV path itself is correctness-tested by the two new
  `paged_hd256_sm100_tma` tests (bit-exact vs dense reference, with
  and without GQA). Perf of the paged path was not benchmarked.

* [hd256] Address review comments on TMA paged KV

- interface.py: assert max_seqlen_k % page_size == 0, page_table sized to
  exact seqlen, and page_table fully contiguous for hd256 paged path
- tests: add shuffled-page-table test; allclose for correctness checks
- paged_kv.py: trim _flatten_smem_sm100 docstring to one line
- sm100_hd256_2cta_fmha_forward.py: cut multi-line comment blocks

* [hd256] Prefetch page indices and eliminate redundant V page reads in TMA paged KV

K and V for the same KV block share the same physical page, so the
separate mPageTable read issued for V was always fetching the same
index already loaded for K.  Carry k_page_idx forward as
v_page_idx_prev and drop all V-side page-table reads.

Additionally, issue the next K page read immediately after K TMA
dispatch (while V TMA is being issued) so the ~25-cycle L2 latency
is hidden behind in-flight work.  Together these changes halve the
number of scalar GMEM page-table reads per kernel call.

NCU (B=4 S=8192 H=8 D=256):
  executed instructions  −0.4 %
  L2 elapsed cycles      −2.2 %  (overhead vs dense: +3.5 % → +1.2 %)

Benchmark — paged vs. dense latency overhead
GPU 0 locked 1965 MHz, non-causal, bf16, page_size=128:

  seqlen    B   before    after    delta
  ------   --   ------   ------   ------
    1024   32   +0.2 %   +0.4 %   −0.2 %
    2048   16   +0.4 %   +0.4 %    0.0 %
    4096    8   −0.1 %   +0.2 %   −0.3 %
    8192    4   +4.9 %   +1.8 %   −3.1 %
   16384    2   +7.7 %   +5.2 %   −2.5 %
   32768    1   +4.9 %   +0.4 %   −4.5 %
   65536    1   +0.9 %   −1.8 %   −2.7 %

No effect at short sequences (TMEM-bound); −2.5 to −4.5 % overhead
reduction at medium-to-long sequences where page-table reads were on
the producer warp's critical path.
ussoewwin pushed a commit to ussoewwin/flash-attention that referenced this pull request May 13, 2026
…ao-AILab#2412)

* [Feat] Support flash-attention head_dim 256 in CuteDSL

This PR adds head_dim=256 support to the FA4 FlashAttention implementation built with the CUTLASS CUTE DSL.

* Forward: uses a 2-CTA design and introduces a new pipeline to better hide memory latency; includes a TMEM-based design for intermediate storage.
* Backward: uses a 2-kernel approach and a 2-CTA design for the backward path.

No API changes for existing head dimensions. But coding style should be adjusted step by step.

This feature is authored by Siyu Wang, Shengbin Di, Yuxi Chi, Johnsonms,
Linfeng Zheng, Haoyan Huang, Lanbo Li, Yun Zhong, Man Yuan, Minmin Sun, Yong Li, Wei Lin.

* Fix ruff lint errors in head_dim=256 changes

Apply ruff check --fix and ruff format to bring the new hd256 files in
line with the project's pre-commit config (flash_attn/cute/*.py, minus
the excluded set in .pre-commit-config.yaml).

Manual fixes:
* mask.py: `Boolean(mask)` -> `cutlass.Boolean(mask)` (F821; other call
  sites in the file already use the qualified form).
* sm100_hd256_2cta_fmha_backward_dkdvkernel.py: drop duplicate
  `SM100_TMEM_CAPACITY_COLUMNS = 512` local definition that shadowed the
  import from tile_scheduler (F811); the values were identical.
* sm100_hd256_2cta_fmha_backward.py: both branches of the
  try/except ImportError imported the same two kernels once
  make_cotiled_copy/warp_reduction_sum were removed as unused; collapse
  to a single unconditional import.

Auto-fixes: 41 unused imports (F401) + 2 f-strings without placeholders
(F541) removed across sm100_hd256_2cta_fmha_{forward,backward,backward_dqkernel,
backward_dkdvkernel}.py, tile_scheduler.py, mask.py. ruff format
reformatted the 8 in-scope files touched by this PR.

Verified: `ruff check` and `ruff format --check` both clean on
flash_attn/cute/ (minus the pre-commit exclude list). Forward + varlen
smoke tests on B200 pass (150 passed, 35 skipped, 0 failed across
non-causal MHA, causal MHA, MQA/GQA, and varlen MHA at d=256).
Backward kernels not yet test-exercised; change is imports/whitespace
only and the kernels parse cleanly.

---------

Co-authored-by: Johnsonms <lizhaofu@gmail.com>
ussoewwin pushed a commit to ussoewwin/flash-attention that referenced this pull request May 13, 2026
…Lab#2487)

Follow-up polish on the freshly-merged hd256 feature (Dao-AILab#2412), sourced
from Copilot AI review comments on the original PR.

interface.py: drop duplicate `from cutlass import Int32` (already imported
at line 17) and unused `from flash_attn.cute.mask import Sm100MaskEnum as
MaskEnum`, which is never referenced.

mask.py: remove two dead `tidx, tidy, tidx = cute.arch.thread_idx()` lines
in Sm100FusedMask.apply_mask and apply_mask_via_causal_local. Neither
`tidx` nor `tidy` is ever read in the function bodies; these calls are
leftover debug scaffolding (consistent with the commented-out
`cute.printf("tidx = ...")` lines nearby at 490/525/665).

test_flash_attn.py: drop the stray "/SM110" from two TODO comments. The
skip guard is `IS_SM100` only (capability major == 10), and the hd256
2CTA kernel path is only taken when `arch // 10 == 10` (interface.py:573,
1310), never on SM110 (major == 11).
ussoewwin pushed a commit to ussoewwin/flash-attention that referenced this pull request May 13, 2026
…rformance gain) (Dao-AILab#2488)

* [hd256] Improve forward kernel with exp2 FMA emulation

Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `28faa77`). The original branch was
based on a pre-merge snapshot; the other five commits in that branch
were absorbed into the squash-merge, leaving this one novel change.

## Change

Replace a fraction of hardware `exp2` (SFU) instructions with a
polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile
computation. The key insight: SM100's SFU throughput is a bottleneck
for hdim=256 due to the large tile size. By substituting 3 out of
every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed
FMA polynomial approximation, we shift pressure onto the underutilized
FMA pipeline.

Additionally, the P write-slot acquisition is moved earlier to overlap
any pipeline stall with the `exp2` compute.

Kernel-only change; no API change. Backward is untouched.

## Validation (this PR vs `origin/main` @ `5bc2d52` — includes Dao-AILab#2412 hd256 base and Dao-AILab#2487 post-merge cleanup)

B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked
clocks @ 1755 MHz, seqlens 4k..128k.

### FWD delta vs `origin/main` (TFLOPS mean, 3 runs each)

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  -0.2%    |  +0.7%   |
| 8k     |   F    |  +0.4%    |  +0.3%   |
| 16k    |   F    |  +0.7%    |  -1.0%   |
| 32k    |   F    |  +2.3%    |  -0.3%   |
| 64k    |   F    |  **+5.4%** |  **+5.2%** |
| 128k   |   F    |  **+7.4%** |  **+7.3%** |
| 4k     |   T    |  +0.3%    |  +0.9%   |
| 8k     |   T    |  +0.8%    |  +0.9%   |
| 16k    |   T    |  +0.8%    |  +1.1%   |
| 32k    |   T    |  **+5.2%** |  -0.4%   |
| 64k    |   T    |  +1.1%    |  +1.1%   |
| 128k   |   T    |  **+3.4%** |  **+2.6%** |

- 19 of 24 cells positive; 4 slightly negative, all within the
  batch-quantization noise band that `origin/main` itself already
  showed in our 3-run regression sweep.
- Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax
  SFU pressure is worst, consistent with the theory above.
- Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**.

### Correctness smoke

`pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k
"256-False-0-0.0-False-False"` on B200:
**78 passed, 78 skipped, 0 failed** — identical pass/skip count to
`origin/main`.

## Caveat

Exp2 FMA emulation introduces small numerical differences vs hardware
`exp2`. The existing test tolerances accept the delta.

* [hd256] Wire ex2_emu params through _TUNING_CONFIG with tuned values

The exp2 emulation knobs (ex2_emu_freq, ex2_emu_res, ex2_emu_start_frg)
and softmax register counts for the hd256 forward kernel were hardcoded
in BlackwellFusedMultiHeadAttentionForward.__init__, invisible to the
central _TUNING_CONFIG table used by all other kernel configs.

- flash_fwd_sm100.py: add hd256 entries to _TUNING_CONFIG (causal and
  non-causal; always 2cta, no sm103 variant). New ex2_emu_res field is
  hd256-specific; existing entries are unaffected. hd256 uses a fixed
  num_regs_other=32 (not derived from the 512-budget formula).
- sm100_hd256_2cta_fmha_forward.py: replace hardcoded self.* assignments
  with a _TUNING_CONFIG lookup.

Tuned values (B200, bf16, locked clocks): freq=14, res=6, start_frg=0
for both causal and non-causal. The inner loop steps k by 2, so k%freq
only takes even values; freq=14/res=6 gives ~43% emulation (3 out of 7
even k%14 steps), replacing the previous 50:50 split (freq=4/res=3).
reubenconducts pushed a commit to reubenconducts/flash-attention that referenced this pull request Jun 2, 2026
…Lab#2489)

* [hd256] Add TMA paged KV support to SM100 2CTA forward kernel

Rebased cherry-pick of `49fe257` from `Johnsonms/paged-kv-hd256` on top
of merged main (hd256 PR Dao-AILab#2412, `28faa77` + post-merge cleanup Dao-AILab#2487,
`5bc2d52`). Original branch was based on a pre-merge snapshot; its
base commits were absorbed into the squash-merge.

## Change

Adds paged KV support to the SM100 hd256 2CTA forward kernel. The paged
path reuses the dense TMA load path — logical KV blocks are remapped to
physical page indices through the page table at load time, so each page
maps to exactly one TMA tile.

**Constraint:** `page_size` must equal `tile_n = 128`.

### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`

- Conditional K/V tensor layout in `__call__`: dense
  `(s_k, d, ((h_r, h_k), b))` vs paged
  `(page_size, d, h_k, num_pages)` for K (and transposed for V).
- Conditional K/V TMA setup in the load warp: dense uses
  `domain_offset` + batch indexing; paged uses `head_kv` slicing and
  keeps `num_pages` as the outer mode for per-load `page_idx` lookup.
- Conditional per-load `page_idx`: K uses mode-2 subtile + mode-3 page;
  V uses mode-1 page.
- Plumb `mPageTable` + `max_seqlen_k` through the kernel signature.
  `seqlen_k` in each of the 4 warp sections now uses `max_seqlen_k`
  for the paged path.
- Store `qhead_per_kvhead` on `self` and derive `head_kv_coord` via
  integer divide (matches the `flash_fwd_sm100` convention for
  contiguous GQA grouping).
- Relax `mPageTable` / `paged_kv_non_tma` assertions.

### `flash_attn/cute/paged_kv.py`

- Extract `_flatten_smem_sm100` / `_copy_row_async` helpers from
  `load_KV` — pure refactor, no behavior change for existing callers.

### `tests/cute/test_flash_attn.py`

- `test_flash_attn_paged_hd256_sm100_tma`: bit-exact vs dense varlen
  reference + determinism check, parametrized over `seqlen_q`.
- `test_flash_attn_paged_hd256_sm100_tma_gqa`: same check for GQA with
  `nheads_kv in {2, 4, 8}` — exercises `qhead_per_kvhead > 1`, which
  a modulo-aliasing bug would fail.

## Validation (this PR vs `origin/main` @ `5bc2d52`)

### Correctness smoke

`pytest tests/cute/test_flash_attn.py` on B200, filter combines the
6 new paged tests with the existing d=256 dense subset:

```
-k "paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)"
```

Result: **84 passed, 78 skipped, 0 failed** in 2 min — 78 from the
dense d=256 subset (identical pass/skip count to `origin/main`) and
**6 from the new `paged_hd256_sm100_tma[_gqa]` tests**.

### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each)

B200, bf16, hdim=256, locked clocks @ 1755 MHz.

| seqlen | causal | MHA 32:32 | GQA 32:2 |
|-------:|:------:|:---------:|:--------:|
| 4k     |   F    |  +0.2%    |  +0.3%   |
| 8k     |   F    |   0.0%    |  -0.1%   |
| 16k    |   F    |  -0.2%    |   0.0%   |
| 32k    |   F    |  -1.0%    |  **+2.3%** |
| 64k    |   F    |  **+2.1%** |  +0.8%  |
| 128k   |   F    |  -0.8%    |  -1.7%   |
| 4k     |   T    |  +0.2%    |  +0.2%   |
| 8k     |   T    |  +0.1%    |  +0.1%   |
| 16k    |   T    |  +0.3%    |  +0.1%   |
| 32k    |   T    |  **+3.0%** |  -1.2%  |
| 64k    |   T    |  -0.3%    |  -0.1%   |
| 128k   |   T    |  +0.1%    |  -0.7%   |

- **22 of 24 cells within ±2%.**
- Two `> 2%` outliers are both **positive** and in the batch-
  quantization noise zone that `origin/main` itself showed run-spread
  in during our 3-run baseline sweep — not regressions.
- **Aggregated means: MHA +0.31%, GQA +0.00%.**
- Paged-KV path isn't exercised by `benchmark_attn.py` (which uses
  contiguous KV); dense-path perf parity is the regression-critical
  property and is preserved.

## Caveat

- **page_size == tile_n == 128 is a hard constraint.** Callers that
  want a different page size will need a separate path.
- The paged-KV path itself is correctness-tested by the two new
  `paged_hd256_sm100_tma` tests (bit-exact vs dense reference, with
  and without GQA). Perf of the paged path was not benchmarked.

* [hd256] Address review comments on TMA paged KV

- interface.py: assert max_seqlen_k % page_size == 0, page_table sized to
  exact seqlen, and page_table fully contiguous for hd256 paged path
- tests: add shuffled-page-table test; allclose for correctness checks
- paged_kv.py: trim _flatten_smem_sm100 docstring to one line
- sm100_hd256_2cta_fmha_forward.py: cut multi-line comment blocks

* [hd256] Prefetch page indices and eliminate redundant V page reads in TMA paged KV

K and V for the same KV block share the same physical page, so the
separate mPageTable read issued for V was always fetching the same
index already loaded for K.  Carry k_page_idx forward as
v_page_idx_prev and drop all V-side page-table reads.

Additionally, issue the next K page read immediately after K TMA
dispatch (while V TMA is being issued) so the ~25-cycle L2 latency
is hidden behind in-flight work.  Together these changes halve the
number of scalar GMEM page-table reads per kernel call.

NCU (B=4 S=8192 H=8 D=256):
  executed instructions  −0.4 %
  L2 elapsed cycles      −2.2 %  (overhead vs dense: +3.5 % → +1.2 %)

Benchmark — paged vs. dense latency overhead
GPU 0 locked 1965 MHz, non-causal, bf16, page_size=128:

  seqlen    B   before    after    delta
  ------   --   ------   ------   ------
    1024   32   +0.2 %   +0.4 %   −0.2 %
    2048   16   +0.4 %   +0.4 %    0.0 %
    4096    8   −0.1 %   +0.2 %   −0.3 %
    8192    4   +4.9 %   +1.8 %   −3.1 %
   16384    2   +7.7 %   +5.2 %   −2.5 %
   32768    1   +4.9 %   +0.4 %   −4.5 %
   65536    1   +0.9 %   −1.8 %   −2.7 %

No effect at short sequences (TMEM-bound); −2.5 to −4.5 % overhead
reduction at medium-to-long sequences where page-table reads were on
the producer warp's critical path.
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.

8 participants