Skip to content

[Cute,hd256] Post-merge cleanup: dead code, duplicate imports - #2487

Merged
Johnsonms merged 1 commit into
Dao-AILab:mainfrom
Johnsonms:Johnsonms/hd256-fixes
Apr 23, 2026
Merged

Johnsonms merged 1 commit into
Dao-AILab:mainfrom
Johnsonms:Johnsonms/hd256-fixes

Conversation

@Johnsonms

Copy link
Copy Markdown
Collaborator

Follow-up Cleanup for hd256 Feature (#2412)

Follow-up polish for the newly merged hd256 feature (#2412), based on Copilot review comments. No kernel code was changed and there is no behavior change.

Changes

  • flash_attn/cute/interface.py

    • Removed duplicate Int32 import
    • Removed unused MaskEnum import
  • flash_attn/cute/mask.py

    • Removed two dead thread_idx() lines in:
      • Sm100FusedMask.apply_mask
      • apply_mask_via_causal_local
    • These were leftover debug scaffolding and were never used
  • tests/cute/test_flash_attn.py

    • Removed stray /SM110 from two TODO comments
    • The hd256 2CTA path is SM100-only, so SM110 was not applicable

Intentionally Unchanged

Two Copilot flags were reviewed and left unchanged because they are false positives:

  • warp_reduction in utils.py

    • Current logic is correct for all existing callers
  • CLC scheduler in tile_scheduler.py

    • The apparent axis mismatch is intentional; the scheduler is axis-agnostic by design

Validation

  • pre-commit run passed cleanly on all three files
  • Smoke test:
    • 78 passed, 78 skipped, 0 failed
  • Regression benchmark vs origin/main:
    • 45/48 cells within ±2%
    • The remaining 3 outliers (+3.4%, -2.4%, +3.6%) fall in an already noisy region
    • No systematic drift observed
    • Aggregated mean delta per category stayed within ±0.3%
      `

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).

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Pull request overview

Post-merge cleanup for the recently added SM100 hd256 (2-CTA) CUTE-DSL path, focusing on removing dead code and redundant/unused imports without changing runtime behavior.

Changes:

  • Removed duplicate/unused imports in flash_attn/cute/interface.py.
  • Removed leftover dead debug scaffolding (thread_idx() assignments) in mask application helpers in flash_attn/cute/mask.py.
  • Updated two TODO comments in tests/cute/test_flash_attn.py to remove the incorrect /SM110 reference.

Reviewed changes

Copilot reviewed 3 out of 3 changed files in this pull request and generated no comments.

File Description
flash_attn/cute/interface.py Drops duplicate/unused imports related to the hd256 follow-up cleanup.
flash_attn/cute/mask.py Removes dead thread_idx() assignments that had no effect on masking logic.
tests/cute/test_flash_attn.py Fixes TODO comments to reflect SM100-only applicability for the hd256 2-CTA path.

💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.

@Johnsonms
Johnsonms merged commit b21e204 into Dao-AILab:main Apr 23, 2026
4 checks passed
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
Johnsonms deleted the Johnsonms/hd256-fixes branch April 23, 2026 22:00
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.
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
…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.

3 participants