Feat([FA4][CUTE DSL]) Add head_dim=256 support (forward + backward) - #2412
Conversation
|
Thanks for the contribution! Can you say more about the new pipeline for fwd? |
Thank you for the feedback! We will try our best to clean up and reorganize the code to match this design. |
|
i see there's a separate test file for hdim 256. Do we still need that or does |
Yes, merging is possible. Currently in the process of organizing. |
|
Forward kernel performance is improved via STG trick. The benchmark is ready to refresh. |
|
@wangsiyu Here are the performance benchmark results on our side for commit PR #2412 — Summary (hdim=256, B200) Bottom line
|
|
@wangsiyu @cherichy Compared with last benchmark:
Key Takeaways:
|
|
@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. |
That makes sense with these shapes. We will provide more benchmark with longer sequences |
|
@Johnsonms Thanks for the updated benchmark!
Key Takeaways: Forward:
Backward:
We will update CLC enabled data later once the performance is optimized. |
Will benchmark soon, let's check the alignment |
My benchmark in fa3 on H100 SMX and fa4 on B200Cc: @tzadouri Key observations:
Comparison: the contributor vs. ours BenchmarksThe two benchmarks are highly consistent, with only minor numerical differences. Key Takeaways
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 |
@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. |
|
@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? |
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.
|
|
Hi @wangsiyu Here is the script I used, Cc: @tzadouri
bench_sm100_hd256.py |
I am refining interface and will triggered these tests. |
|
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. |
|
@Johnsonms All conflicts have been resolved and All relateive unit tests passed |
|
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 |
|
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. |
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).
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).
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.
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.
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.
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.
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.
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.
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.
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.
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.
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.
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.
Both fwd and bwd will be refactored. |
…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).
* [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.
…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>
…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).
…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).
…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.

















Summary
This PR adds
head_dim=256support to the FA4 FlashAttention implementation built with the CUTLASS CUTE DSL.What’s included
head_dim=256Motivation
head_dim=256is common in newer model variants, but FA4 CUTE-DSL coverage is currently limited.Implementation
Performance
Performance numbers will be added in a follow-up update once the benchmark suite and configurations are finalized.
Testing
Author
This feature is authored by @wangsiyu @dishengbin @cherichy @Johnsonms