[Cute,hd256] Post-merge cleanup: dead code, duplicate imports - #2487
Merged
Merged
Conversation
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).
Contributor
There was a problem hiding this comment.
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 inflash_attn/cute/mask.py. - Updated two TODO comments in
tests/cute/test_flash_attn.pyto remove the incorrect/SM110reference.
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.
drisspg
approved these changes
Apr 23, 2026
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `e122e67` from `Johnsonms/exp2-emu-hd256` on top of merged main (hd256 PR Dao-AILab#2412, `27b4eb9`). The original branch was based on a pre-merge snapshot; the other five commits in that branch were absorbed into the squash-merge, leaving this one novel change. ## Change Replace a fraction of hardware `exp2` (SFU) instructions with a polynomial FMA emulation (`ex2_emulation_2`) in the softmax P-tile computation. The key insight: SM100's SFU throughput is a bottleneck for hdim=256 due to the large tile size. By substituting 3 out of every 4 `exp2` calls (`ex2_emu_freq=4`, `ex2_emu_res=3`) with packed FMA polynomial approximation, we shift pressure onto the underutilized FMA pipeline. Additionally, the P write-slot acquisition is moved earlier to overlap any pipeline stall with the `exp2` compute. Kernel-only change; no API change. Backward is untouched. ## Validation (this PR vs `origin/main` @ `b21e204` — includes Dao-AILab#2412 hd256 base and Dao-AILab#2487 post-merge cleanup) B200, bf16, hdim=256, MHA (32:32) and GQA (32:2), 3-run means, locked clocks @ 1755 MHz, seqlens 4k..128k. ### FWD delta vs `origin/main` (TFLOPS mean, 3 runs each) | seqlen | causal | MHA 32:32 | GQA 32:2 | |-------:|:------:|:---------:|:--------:| | 4k | F | -0.2% | +0.7% | | 8k | F | +0.4% | +0.3% | | 16k | F | +0.7% | -1.0% | | 32k | F | +2.3% | -0.3% | | 64k | F | **+5.4%** | **+5.2%** | | 128k | F | **+7.4%** | **+7.3%** | | 4k | T | +0.3% | +0.9% | | 8k | T | +0.8% | +0.9% | | 16k | T | +0.8% | +1.1% | | 32k | T | **+5.2%** | -0.4% | | 64k | T | +1.1% | +1.1% | | 128k | T | **+3.4%** | **+2.6%** | - 19 of 24 cells positive; 4 slightly negative, all within the batch-quantization noise band that `origin/main` itself already showed in our 3-run regression sweep. - Peak gain **+7.4%** (MHA 128k non-causal) — exactly where softmax SFU pressure is worst, consistent with the theory above. - Averages: **MHA fwd +2.3%**, **GQA fwd +1.5%**. ### Correctness smoke `pytest tests/cute/test_flash_attn.py::test_flash_attn_output -k "256-False-0-0.0-False-False"` on B200: **78 passed, 78 skipped, 0 failed** — identical pass/skip count to `origin/main`. ## Caveat Exp2 FMA emulation introduces small numerical differences vs hardware `exp2`. The existing test tolerances accept the delta.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `49fe257` from `Johnsonms/paged-kv-hd256` on top of merged main (hd256 PR Dao-AILab#2412, `27b4eb9` + post-merge cleanup Dao-AILab#2487, `b21e204`). Original branch was based on a pre-merge snapshot; its base commits were absorbed into the squash-merge. ## Change Adds paged KV support to the SM100 hd256 2CTA forward kernel. The paged path reuses the dense TMA load path — logical KV blocks are remapped to physical page indices through the page table at load time, so each page maps to exactly one TMA tile. **Constraint:** `page_size` must equal `tile_n = 128`. ### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py` - Conditional K/V tensor layout in `__call__`: dense `(s_k, d, ((h_r, h_k), b))` vs paged `(page_size, d, h_k, num_pages)` for K (and transposed for V). - Conditional K/V TMA setup in the load warp: dense uses `domain_offset` + batch indexing; paged uses `head_kv` slicing and keeps `num_pages` as the outer mode for per-load `page_idx` lookup. - Conditional per-load `page_idx`: K uses mode-2 subtile + mode-3 page; V uses mode-1 page. - Plumb `mPageTable` + `max_seqlen_k` through the kernel signature. `seqlen_k` in each of the 4 warp sections now uses `max_seqlen_k` for the paged path. - Store `qhead_per_kvhead` on `self` and derive `head_kv_coord` via integer divide (matches the `flash_fwd_sm100` convention for contiguous GQA grouping). - Relax `mPageTable` / `paged_kv_non_tma` assertions. ### `flash_attn/cute/paged_kv.py` - Extract `_flatten_smem_sm100` / `_copy_row_async` helpers from `load_KV` — pure refactor, no behavior change for existing callers. ### `tests/cute/test_flash_attn.py` - `test_flash_attn_paged_hd256_sm100_tma`: bit-exact vs dense varlen reference + determinism check, parametrized over `seqlen_q`. - `test_flash_attn_paged_hd256_sm100_tma_gqa`: same check for GQA with `nheads_kv in {2, 4, 8}` — exercises `qhead_per_kvhead > 1`, which a modulo-aliasing bug would fail. ## Validation (this PR vs `origin/main` @ `b21e204`) ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 new paged tests with the existing d=256 dense subset: ``` -k "paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **84 passed, 78 skipped, 0 failed** in 2 min — 78 from the dense d=256 subset (identical pass/skip count to `origin/main`) and **6 from the new `paged_hd256_sm100_tma[_gqa]` tests**. ### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each) B200, bf16, hdim=256, locked clocks @ 1755 MHz. | seqlen | causal | MHA 32:32 | GQA 32:2 | |-------:|:------:|:---------:|:--------:| | 4k | F | +0.2% | +0.3% | | 8k | F | 0.0% | -0.1% | | 16k | F | -0.2% | 0.0% | | 32k | F | -1.0% | **+2.3%** | | 64k | F | **+2.1%** | +0.8% | | 128k | F | -0.8% | -1.7% | | 4k | T | +0.2% | +0.2% | | 8k | T | +0.1% | +0.1% | | 16k | T | +0.3% | +0.1% | | 32k | T | **+3.0%** | -1.2% | | 64k | T | -0.3% | -0.1% | | 128k | T | +0.1% | -0.7% | - **22 of 24 cells within ±2%.** - Two `> 2%` outliers are both **positive** and in the batch- quantization noise zone that `origin/main` itself showed run-spread in during our 3-run baseline sweep — not regressions. - **Aggregated means: MHA +0.31%, GQA +0.00%.** - Paged-KV path isn't exercised by `benchmark_attn.py` (which uses contiguous KV); dense-path perf parity is the regression-critical property and is preserved. ## Caveat - **page_size == tile_n == 128 is a hard constraint.** Callers that want a different page size will need a separate path. - The paged-KV path itself is correctness-tested by the two new `paged_hd256_sm100_tma` tests (bit-exact vs dense reference, with and without GQA). Perf of the paged path was not benchmarked.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `94e63db` from `Johnsonms/seqused-k-hd256` on top of `Johnsonms/paged-kv-hd256-v2`. Original branch was based on a pre-merge snapshot of the hd256 tree; base commits were absorbed into the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487 (`b21e204`). **Stacked PR** — base branch should be `Johnsonms/paged-kv-hd256-v2`, not `main`. Depends on that PR for the paged-KV kernel plumbing. ## Change Enables variable per-batch KV sequence lengths via a `seqused_k` tensor — needed for MLA-style decode (DeepSeek-V2 / V3 / R1), where different batches have different KV cache occupancies. Works with both dense and paged K/V. ### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py` - Drop the `mSeqUsedK` half of the `__call__` assertion; only `mSeqUsedQ` stays blocked for now. - Build the `seqused_k` cute tensor in `__call__` alongside the `page_table` / paged-layout construction. - Add `mSeqUsedK` to `kernel()` signature and pass `seqused_k` from `__call__`. - Replace `seqlen_k` derivation in all 4 warp sections with a ternary: `mSeqUsedK[batch_coord] if set, else <dense/paged expression>`. - Move `batch_coord` above `seqlen_k` in the MMA warp (second warp section) — it was declared later but now needs to be in scope for `seqused_k` indexing. - **Zero-KV batch handling (`seqlen_k == 0`):** extend `continue_cond` in all four warp sections with `continue_cond or seqlen_k <= 0`, so load / MMA / correction / softmax warps skip in sync instead of deadlocking on `K0 / Vend / QK0 / PVend / first-stats` tiles. ### `flash_attn/cute/interface.py` - Relax the `seqused_q is None and seqused_k is None` assertion to `seqused_q is None` (the kernel now handles `seqused_k`). - Prefill zero-KV batches on the host with zero output and `-inf` LSE, so their output is defined even though the kernel skips them. ### `tests/cute/test_flash_attn.py` - `test_flash_attn_seqused_k_hd256_sm100`: dense + `seqused_k` (padded K/V with per-batch valid lengths) bit-exact vs a `cu_seqlens_k` packed reference, parametrized over asymmetric per-batch lengths. - `test_flash_attn_paged_seqused_k_hd256_sm100`: paged + `seqused_k` combined (MLA decode pattern), bit-exact vs packed reference. - `test_flash_attn_seqused_k_zero_hd256_sm100`: `seqused_k = 0` for one batch, parametrized dense/paged. Verifies no deadlock, zero output, `-inf` LSE on the empty row, and finite output / LSE on the other row. ## Validation ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 new `seqused_k` tests, the 6 paged tests from the parent PR, and the existing d=256 dense subset: ``` -k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **90 passed, 78 skipped, 0 failed** in 2m14s — 78 from the dense d=256 subset (identical pass/skip count to `origin/main`), 6 from the parent PR's `paged_hd256_sm100_tma[_gqa]` tests, and **6 from the 3 new `seqused_k_*_hd256_sm100` tests** introduced here. ### FWD perf Not re-measured on this branch. Expected to track the parent PR's numbers (within ±0.3% mean of `origin/main`) on the dense path, since the `seqused_k` plumbing only adds a scalar-tensor indirection per batch and a `continue_cond` check per warp section — no change to inner-loop throughput. ## Caveat - Host-side prefill in `interface.py` walks zero-KV batch indices in Python. For batch sizes in the thousands with many zero-KV rows this could show up in CPU profiles; the current FA4 callers aren't in that regime. - `seqused_q` is still asserted `None` — the kernel would need a similar ternary in the Q-side iteration, out of scope for this PR.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `94e63db` from `Johnsonms/seqused-k-hd256` on top of `Johnsonms/paged-kv-hd256-v2`. Original branch was based on a pre-merge snapshot of the hd256 tree; base commits were absorbed into the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487 (`b21e204`). **Stacked PR** — base branch should be `Johnsonms/paged-kv-hd256-v2`, not `main`. Depends on that PR for the paged-KV kernel plumbing. ## Change Enables variable per-batch KV sequence lengths via a `seqused_k` tensor — needed for MLA-style decode (DeepSeek-V2 / V3 / R1), where different batches have different KV cache occupancies. Works with both dense and paged K/V. ### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py` - Drop the `mSeqUsedK` half of the `__call__` assertion; only `mSeqUsedQ` stays blocked for now. - Build the `seqused_k` cute tensor in `__call__` alongside the `page_table` / paged-layout construction. - Add `mSeqUsedK` to `kernel()` signature and pass `seqused_k` from `__call__`. - Replace `seqlen_k` derivation in all 4 warp sections with a ternary: `mSeqUsedK[batch_coord] if set, else <dense/paged expression>`. - Move `batch_coord` above `seqlen_k` in the MMA warp (second warp section) — it was declared later but now needs to be in scope for `seqused_k` indexing. - **Zero-KV batch handling (`seqlen_k == 0`):** extend `continue_cond` in all four warp sections with `continue_cond or seqlen_k <= 0`, so load / MMA / correction / softmax warps skip in sync instead of deadlocking on `K0 / Vend / QK0 / PVend / first-stats` tiles. ### `flash_attn/cute/interface.py` - Relax the `seqused_q is None and seqused_k is None` assertion to `seqused_q is None` (the kernel now handles `seqused_k`). - Prefill zero-KV batches on the host with zero output and `-inf` LSE, so their output is defined even though the kernel skips them. ### `tests/cute/test_flash_attn.py` - `test_flash_attn_seqused_k_hd256_sm100`: dense + `seqused_k` (padded K/V with per-batch valid lengths) bit-exact vs a `cu_seqlens_k` packed reference, parametrized over asymmetric per-batch lengths. - `test_flash_attn_paged_seqused_k_hd256_sm100`: paged + `seqused_k` combined (MLA decode pattern), bit-exact vs packed reference. - `test_flash_attn_seqused_k_zero_hd256_sm100`: `seqused_k = 0` for one batch, parametrized dense/paged. Verifies no deadlock, zero output, `-inf` LSE on the empty row, and finite output / LSE on the other row. ## Validation ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 new `seqused_k` tests, the 6 paged tests from the parent PR, and the existing d=256 dense subset: ``` -k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **90 passed, 78 skipped, 0 failed** — 78 from the dense d=256 subset (identical pass/skip count to `origin/main`), 6 from the parent PR's `paged_hd256_sm100_tma[_gqa]` tests, and **6 from the 3 new `seqused_k_*_hd256_sm100` tests** introduced here. ### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each) B200, bf16, hdim=256, locked clocks @ 1755 MHz. #### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1474 | 1485 | +0.7% | | 8k | F | 1582 | 1595 | +0.8% | | 16k | F | 1641 | 1650 | +0.5% | | 32k | F | 1450 | 1488 | **+2.6%** | | 64k | F | 1417 | 1484 | **+4.7%** | | 128k | F | 1398 | 1492 | **+6.7%** | | 4k | T | 1215 | 1217 | +0.2% | | 8k | T | 1411 | 1415 | +0.3% | | 16k | T | 1540 | 1545 | +0.3% | | 32k | T | 1552 | 1615 | **+4.1%** | | 64k | T | 1486 | 1496 | +0.7% | | 128k | T | 1363 | 1370 | +0.5% | #### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1496 | 1511 | +1.0% | | 8k | F | 1601 | 1612 | +0.7% | | 16k | F | 1620 | 1653 | **+2.0%** | | 32k | F | 1482 | 1464 | −1.2% | | 64k | F | 1436 | 1472 | **+2.5%** | | 128k | F | 1389 | 1469 | **+5.8%** | | 4k | T | 1242 | 1244 | +0.2% | | 8k | T | 1434 | 1439 | +0.3% | | 16k | T | 1556 | 1562 | +0.4% | | 32k | T | 1624 | 1564 | −3.7% | | 64k | T | 1493 | 1507 | +0.9% | | 128k | T | 1373 | 1368 | −0.4% | Aggregated means: **MHA fwd +1.8%, GQA fwd +0.7%**. ## Caveat — unexpected long-seqlen non-causal speedup The intent of this change is correctness only (adds a ternary, a `continue_cond` extension, moves a variable declaration). None of these should improve inner-loop throughput. Yet we observe a **reproducible** +4–7% at long-seqlen non-causal (MHA 128k F +6.7%, GQA 128k F +5.8%, MHA 64k F +4.7%, MHA 32k T +4.1%). 3-run variance per cell is tight (typically <1%), so this is not run-to-run noise. Bracketed against the parent `paged-kv-v2` PR, which measured within ±0.3% of main on the same cells, the gain is introduced specifically by this commit's source-level reorderings (likely register-allocation or instruction-scheduling artifacts from ptxas). ### Why this is a concern, not just free perf 1. **Unintended change** — we can't explain it from the diff, which means the compiler's decision hinges on something fragile (variable declaration order, kernel signature shape). A future unrelated edit could flip this back to `±0%`, or worse, regress it. 2. **Could mask a different regression.** If the reordering also subtly changed some other code path we don't benchmark (e.g. paged path with `seqused_k = None`), we wouldn't notice until production. 3. **Not portable guidance.** We can't tell future contributors "move `batch_coord` earlier to get +6%" because the mechanism isn't a deliberate optimization. ### TODO before opening / merging - [ ] Compare SASS between this branch and `paged-kv-v2` for the long-seqlen non-causal dense path; identify which instructions changed and whether the gain is attributable to a specific scheduling/allocation difference. - [ ] Confirm the paged path with `seqused_k = None` isn't regressed (benchmark_attn.py doesn't exercise paged; add a quick paged-bench harness or run the paged tests under timing). - [ ] Decide whether to keep the reorderings (if the mechanism is understood) or revert the non-essential ones (declaration move) to isolate the correctness change from the accidental perf gain. ## Other caveat - Host-side prefill in `interface.py` walks zero-KV batch indices in Python. For batch sizes in the thousands with many zero-KV rows this could show up in CPU profiles; current FA4 callers aren't in that regime. - `seqused_q` is still asserted `None` — the kernel would need a similar ternary in the Q-side iteration, out of scope for this PR.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `dbb6c98` from `Johnsonms/persistent-cluster-hd256` on top of `Johnsonms/seqused-k-hd256-v2`. Original branch was based on a pre-merge snapshot of the hd256 tree; base commits were absorbed into the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487 (`b21e204`). **Stacked PR** — base branch should be `Johnsonms/seqused-k-hd256-v2`, not `main`. Depends on that PR for the `seqused_k` + paged-KV plumbing. ## Change Persistent scheduling amortizes CTA launch overhead by issuing a grid-stride loop over tiles. It was hardcoded off in hd256 since the kernel's inception because the static tile scheduler was cluster- unaware and split 2CTA clusters across independent work tiles, corrupting output. ### Cluster-aware fix - `Sm100FmhaStaticTileSchedulerParams` gains a `cluster_shape_m` constexpr (default 1, so 1CTA kernels are unchanged). - Grid is sized in cluster units: `max_ctas = (sm_count // cluster_shape_m) * cluster_shape_m`; problem size is multiplied by `cluster_shape_m` so `dsl_min` compares apples to apples. - `num_persistent_clusters` replaces `num_persistent_sm` as the grid-stride step, so `advance_to_next_work` advances by one cluster per iteration instead of one CTA. - In `get_current_work`, the CTA rank within the cluster is reconstructed from the launch `block_idx` (`cta_rank = blk_coord[0] % cluster_shape_m`) and spliced back into `mid = m_block * cluster_shape_m + cta_rank`, so both CTAs in a 2CTA cluster land on their half of the same tile. ### Enablement gate (work-per-tile heuristic) Persistent's per-tile cost (cluster-rank reconstruction + grid-stride state) only pays off when work-per-tile is small, i.e. short KV. On B200, persistent wins at short seqlen but regresses dense prefill at long seqlen because launch overhead is already amortized by the hardware at high tile counts. - **Gate:** `... and seqlen_k <= 2048`. Keeps the decode win and holds long-context prefill within noise of the pre-persistent baseline. - Gate is folded into `interface.py`'s compile_key so configs straddling the threshold compile separate kernels with the correct persistent flag baked in. ### Files - `flash_attn/cute/tile_scheduler.py`: add `cluster_shape_m` on `Sm100FmhaStaticTileSchedulerParams` and the matching updates to `Sm100FmhaStaticTileScheduler` and `compute_sm100_fmha_grid`. - `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`: drop `self.is_persistent = False` hardcode (now honors constructor arg), pass `cluster_shape_mnk` to `compute_grid`, divide `blk_idx[0]` by `cluster_shape_mnk[0]` when constructing `FmhaStaticTileScheduler`. - `flash_attn/cute/sm100_hd256_2cta_fmha_backward_dqkernel.py`: same fix applied to the bwd dQ kernel for consistency (its scheduler is the same `Sm100` static scheduler). - `flash_attn/cute/interface.py`: compute `hd256_is_persistent` above the compile_key (gated on `seqlen_k <= 2048` in addition to causal/cu_seqlens_q/...); add it to the compile_key; use it at `fa_fwd` construction. ## Validation ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 `seqused_k` tests (inherited from parent PR), the 6 paged tests (grandparent PR), and the existing d=256 dense subset: ``` -k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **90 passed, 78 skipped, 0 failed** — identical pass/skip count to `origin/main`. No new tests added by this commit (scheduler refactor tests are covered via the existing d=256 dense parametrize space, which hits both gate-ON and gate-OFF configurations). ### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each) B200, bf16, hdim=256, locked clocks @ 1755 MHz. #### Short seqlen — **gate-ON path** (persistent scheduling active) ##### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 1k | F | 985 | 987 | +0.2% | | 2k | F | 1307 | 1301 | −0.5% | | 1k | T | 592 | 592 | 0.0% | | 2k | T | 910 | 911 | +0.1% | ##### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 1k | F | 1012 | 1015 | +0.3% | | 2k | F | 1329 | 1329 | 0.0% | | 1k | T | 600 | 600 | 0.0% | | 2k | T | 919 | 919 | 0.0% | **Observation:** the persistent path is active but delivers essentially **no speedup** in this bench config (32 Q-heads MHA, 32:2 GQA). The commit message cites +10–25% wins at `seqlen_k=1024` on a different config (8 Q-heads); with 32 Q-heads the tile count per batch is 4× larger and launch overhead is already amortized by the hardware, so the persistent loop's benefit is saturated out. Reproducibility is very tight (variance <1 TFLOPS across 3 runs), so this is not run noise — just "no regression" rather than "a win" at this config. #### Long seqlen — **gate-OFF path** (should match main) ##### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1474 | 1479 | +0.3% | | 8k | F | 1582 | 1588 | +0.4% | | 16k | F | 1641 | 1636 | −0.3% | | 32k | F | 1450 | 1442 | −0.6% | | 64k | F | 1417 | 1447 | **+2.1%** | | 128k | F | 1398 | 1423 | +1.8% | | 4k | T | 1215 | 1218 | +0.2% | | 8k | T | 1411 | 1416 | +0.4% | | 16k | T | 1540 | 1540 | 0.0% | | 32k | T | 1552 | 1619 | **+4.3%** | | 64k | T | 1486 | 1469 | −1.1% | | 128k | T | 1363 | 1383 | +1.5% | ##### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1496 | 1504 | +0.5% | | 8k | F | 1601 | 1598 | −0.2% | | 16k | F | 1620 | 1640 | +1.2% | | 32k | F | 1482 | 1539 | **+3.8%** | | 64k | F | 1436 | 1460 | +1.7% | | 128k | F | 1389 | 1363 | −1.9% | | 4k | T | 1242 | 1244 | +0.2% | | 8k | T | 1434 | 1438 | +0.3% | | 16k | T | 1556 | 1557 | +0.1% | | 32k | T | 1624 | 1574 | **−3.1%** | | 64k | T | 1493 | 1493 | 0.0% | | 128k | T | 1373 | 1369 | −0.3% | **Observation:** long-seqlen path is within ±2% on most cells. Same long-seqlen noise pattern as the parent PRs at 32k/64k (batch- quantization zone); deltas swing both directions, no systematic regression. ## Caveats - **Short-seqlen benefit is config-dependent.** The persistent path trades extra per-tile state for launch-overhead amortization; the trade only pays off when work-per-tile is small. At 8 Q-heads the commit's original bench showed +10–25% at 1k; at 32 Q-heads in my bench the benefit collapses to 0% because tile count is already high enough to amortize launches. If the use case is decode-style small-batch / few-head, the gate-ON path is expected to deliver the advertised win. - Backward dQ kernel gets the same cluster-aware scheduler change for consistency but is not exercised by `benchmark_attn.py` in this PR's validation. Forward dense path is the regression-critical one and is covered above. - `cluster_shape_m` currently defaults to 1 so 1CTA kernels are unaffected; the only callers passing `cluster_shape_m > 1` are hd256 forward and hd256 bwd dQ.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `dbb6c98` from `Johnsonms/persistent-cluster-hd256` on top of `Johnsonms/seqused-k-hd256-v2`. Original branch was based on a pre-merge snapshot of the hd256 tree; base commits were absorbed into the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487 (`b21e204`). **Stacked PR** — base branch should be `Johnsonms/seqused-k-hd256-v2`, not `main`. Depends on that PR for the `seqused_k` + paged-KV plumbing. ## Change Persistent scheduling amortizes CTA launch overhead by issuing a grid-stride loop over tiles. It was hardcoded off in hd256 since the kernel's inception because the static tile scheduler was cluster- unaware and split 2CTA clusters across independent work tiles, corrupting output. ### Cluster-aware fix - `Sm100FmhaStaticTileSchedulerParams` gains a `cluster_shape_m` constexpr (default 1, so 1CTA kernels are unchanged). - Grid is sized in cluster units: `max_ctas = (sm_count // cluster_shape_m) * cluster_shape_m`; problem size is multiplied by `cluster_shape_m` so `dsl_min` compares apples to apples. - `num_persistent_clusters` replaces `num_persistent_sm` as the grid-stride step, so `advance_to_next_work` advances by one cluster per iteration instead of one CTA. - In `get_current_work`, the CTA rank within the cluster is reconstructed from the launch `block_idx` (`cta_rank = blk_coord[0] % cluster_shape_m`) and spliced back into `mid = m_block * cluster_shape_m + cta_rank`, so both CTAs in a 2CTA cluster land on their half of the same tile. ### Enablement gate (work-per-tile heuristic) Persistent's per-tile cost (cluster-rank reconstruction + grid-stride state) only pays off when work-per-tile is small, i.e. short KV. On B200, persistent wins at short seqlen but regresses dense prefill at long seqlen because launch overhead is already amortized by the hardware at high tile counts. - **Gate:** `... and seqlen_k <= 2048`. Keeps the decode win and holds long-context prefill within noise of the pre-persistent baseline. - Gate is folded into `interface.py`'s compile_key so configs straddling the threshold compile separate kernels with the correct persistent flag baked in. ### Files - `flash_attn/cute/tile_scheduler.py`: add `cluster_shape_m` on `Sm100FmhaStaticTileSchedulerParams` and the matching updates to `Sm100FmhaStaticTileScheduler` and `compute_sm100_fmha_grid`. - `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`: drop `self.is_persistent = False` hardcode (now honors constructor arg), pass `cluster_shape_mnk` to `compute_grid`, divide `blk_idx[0]` by `cluster_shape_mnk[0]` when constructing `FmhaStaticTileScheduler`. - `flash_attn/cute/sm100_hd256_2cta_fmha_backward_dqkernel.py`: same fix applied to the bwd dQ kernel for consistency (its scheduler is the same `Sm100` static scheduler). - `flash_attn/cute/interface.py`: compute `hd256_is_persistent` above the compile_key (gated on `seqlen_k <= 2048` in addition to causal/cu_seqlens_q/...); add it to the compile_key; use it at `fa_fwd` construction. ## Validation ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 `seqused_k` tests (inherited from parent PR), the 6 paged tests (grandparent PR), and the existing d=256 dense subset: ``` -k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **90 passed, 78 skipped, 0 failed** — identical pass/skip count to `origin/main`. No new tests added by this commit (scheduler refactor tests are covered via the existing d=256 dense parametrize space, which hits both gate-ON and gate-OFF configurations). ### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each) B200, bf16, hdim=256, locked clocks @ 1755 MHz. #### Short seqlen — **gate-ON path** (persistent scheduling active) ##### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 1k | F | 985 | 987 | +0.2% | | 2k | F | 1307 | 1301 | −0.5% | | 1k | T | 592 | 592 | 0.0% | | 2k | T | 910 | 911 | +0.1% | ##### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 1k | F | 1012 | 1015 | +0.3% | | 2k | F | 1329 | 1329 | 0.0% | | 1k | T | 600 | 600 | 0.0% | | 2k | T | 919 | 919 | 0.0% | **Observation:** the persistent path is active but delivers essentially **no speedup** in this bench config (32 Q-heads MHA, 32:2 GQA). The commit message cites +10–25% wins at `seqlen_k=1024` on a different config (8 Q-heads); with 32 Q-heads the tile count per batch is 4× larger and launch overhead is already amortized by the hardware, so the persistent loop's benefit is saturated out. Reproducibility is very tight (variance <1 TFLOPS across 3 runs), so this is not run noise — just "no regression" rather than "a win" at this config. #### Long seqlen — **gate-OFF path** (should match main) ##### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1474 | 1479 | +0.3% | | 8k | F | 1582 | 1588 | +0.4% | | 16k | F | 1641 | 1636 | −0.3% | | 32k | F | 1450 | 1442 | −0.6% | | 64k | F | 1417 | 1447 | **+2.1%** | | 128k | F | 1398 | 1423 | +1.8% | | 4k | T | 1215 | 1218 | +0.2% | | 8k | T | 1411 | 1416 | +0.4% | | 16k | T | 1540 | 1540 | 0.0% | | 32k | T | 1552 | 1619 | **+4.3%** | | 64k | T | 1486 | 1469 | −1.1% | | 128k | T | 1363 | 1383 | +1.5% | ##### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1496 | 1504 | +0.5% | | 8k | F | 1601 | 1598 | −0.2% | | 16k | F | 1620 | 1640 | +1.2% | | 32k | F | 1482 | 1539 | **+3.8%** | | 64k | F | 1436 | 1460 | +1.7% | | 128k | F | 1389 | 1363 | −1.9% | | 4k | T | 1242 | 1244 | +0.2% | | 8k | T | 1434 | 1438 | +0.3% | | 16k | T | 1556 | 1557 | +0.1% | | 32k | T | 1624 | 1574 | **−3.1%** | | 64k | T | 1493 | 1493 | 0.0% | | 128k | T | 1373 | 1369 | −0.3% | **Observation:** long-seqlen path is within ±2% on most cells. Same long-seqlen noise pattern as the parent PRs at 32k/64k (batch- quantization zone); deltas swing both directions, no systematic regression. ## Caveats - **Short-seqlen benefit is config-dependent.** The persistent path trades extra per-tile state for launch-overhead amortization; the trade only pays off when work-per-tile is small. At 8 Q-heads the commit's original bench showed +10–25% at 1k; at 32 Q-heads in my bench the benefit collapses to 0% because tile count is already high enough to amortize launches. If the use case is decode-style small-batch / few-head, the gate-ON path is expected to deliver the advertised win. - Backward dQ kernel gets the same cluster-aware scheduler change for consistency but is not exercised by `benchmark_attn.py` in this PR's validation. Forward dense path is the regression-critical one and is covered above. - `cluster_shape_m` currently defaults to 1 so 1CTA kernels are unaffected; the only callers passing `cluster_shape_m > 1` are hd256 forward and hd256 bwd dQ.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `94e63db` from `Johnsonms/seqused-k-hd256` on top of `Johnsonms/paged-kv-hd256-v2`. Original branch was based on a pre-merge snapshot of the hd256 tree; base commits were absorbed into the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487 (`b21e204`). **Stacked PR** — base branch should be `Johnsonms/paged-kv-hd256-v2`, not `main`. Depends on that PR for the paged-KV kernel plumbing. ## Change Enables variable per-batch KV sequence lengths via a `seqused_k` tensor — needed for MLA-style decode (DeepSeek-V2 / V3 / R1), where different batches have different KV cache occupancies. Works with both dense and paged K/V. ### `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py` - Drop the `mSeqUsedK` half of the `__call__` assertion; only `mSeqUsedQ` stays blocked for now. - Build the `seqused_k` cute tensor in `__call__` alongside the `page_table` / paged-layout construction. - Add `mSeqUsedK` to `kernel()` signature and pass `seqused_k` from `__call__`. - Replace `seqlen_k` derivation in all 4 warp sections with a ternary: `mSeqUsedK[batch_coord] if set, else <dense/paged expression>`. - Move `batch_coord` above `seqlen_k` in the MMA warp (second warp section) — it was declared later but now needs to be in scope for `seqused_k` indexing. - **Zero-KV batch handling (`seqlen_k == 0`):** extend `continue_cond` in all four warp sections with `continue_cond or seqlen_k <= 0`, so load / MMA / correction / softmax warps skip in sync instead of deadlocking on `K0 / Vend / QK0 / PVend / first-stats` tiles. ### `flash_attn/cute/interface.py` - Relax the `seqused_q is None and seqused_k is None` assertion to `seqused_q is None` (the kernel now handles `seqused_k`). - Prefill zero-KV batches on the host with zero output and `-inf` LSE, so their output is defined even though the kernel skips them. ### `tests/cute/test_flash_attn.py` - `test_flash_attn_seqused_k_hd256_sm100`: dense + `seqused_k` (padded K/V with per-batch valid lengths) bit-exact vs a `cu_seqlens_k` packed reference, parametrized over asymmetric per-batch lengths. - `test_flash_attn_paged_seqused_k_hd256_sm100`: paged + `seqused_k` combined (MLA decode pattern), bit-exact vs packed reference. - `test_flash_attn_seqused_k_zero_hd256_sm100`: `seqused_k = 0` for one batch, parametrized dense/paged. Verifies no deadlock, zero output, `-inf` LSE on the empty row, and finite output / LSE on the other row. ## Validation ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 new `seqused_k` tests, the 6 paged tests from the parent PR, and the existing d=256 dense subset: ``` -k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **90 passed, 78 skipped, 0 failed** — 78 from the dense d=256 subset (identical pass/skip count to `origin/main`), 6 from the parent PR's `paged_hd256_sm100_tma[_gqa]` tests, and **6 from the 3 new `seqused_k_*_hd256_sm100` tests** introduced here. ### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each) B200, bf16, hdim=256, locked clocks @ 1755 MHz. #### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1474 | 1485 | +0.7% | | 8k | F | 1582 | 1595 | +0.8% | | 16k | F | 1641 | 1650 | +0.5% | | 32k | F | 1450 | 1488 | **+2.6%** | | 64k | F | 1417 | 1484 | **+4.7%** | | 128k | F | 1398 | 1492 | **+6.7%** | | 4k | T | 1215 | 1217 | +0.2% | | 8k | T | 1411 | 1415 | +0.3% | | 16k | T | 1540 | 1545 | +0.3% | | 32k | T | 1552 | 1615 | **+4.1%** | | 64k | T | 1486 | 1496 | +0.7% | | 128k | T | 1363 | 1370 | +0.5% | #### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1496 | 1511 | +1.0% | | 8k | F | 1601 | 1612 | +0.7% | | 16k | F | 1620 | 1653 | **+2.0%** | | 32k | F | 1482 | 1464 | −1.2% | | 64k | F | 1436 | 1472 | **+2.5%** | | 128k | F | 1389 | 1469 | **+5.8%** | | 4k | T | 1242 | 1244 | +0.2% | | 8k | T | 1434 | 1439 | +0.3% | | 16k | T | 1556 | 1562 | +0.4% | | 32k | T | 1624 | 1564 | −3.7% | | 64k | T | 1493 | 1507 | +0.9% | | 128k | T | 1373 | 1368 | −0.4% | Aggregated means: **MHA fwd +1.8%, GQA fwd +0.7%**. ## Caveat — unexpected long-seqlen non-causal speedup The intent of this change is correctness only (adds a ternary, a `continue_cond` extension, moves a variable declaration). None of these should improve inner-loop throughput. Yet we observe a **reproducible** +4–7% at long-seqlen non-causal (MHA 128k F +6.7%, GQA 128k F +5.8%, MHA 64k F +4.7%, MHA 32k T +4.1%). 3-run variance per cell is tight (typically <1%), so this is not run-to-run noise. Bracketed against the parent `paged-kv-v2` PR, which measured within ±0.3% of main on the same cells, the gain is introduced specifically by this commit's source-level reorderings (likely register-allocation or instruction-scheduling artifacts from ptxas). ### Why this is a concern, not just free perf 1. **Unintended change** — we can't explain it from the diff, which means the compiler's decision hinges on something fragile (variable declaration order, kernel signature shape). A future unrelated edit could flip this back to `±0%`, or worse, regress it. 2. **Could mask a different regression.** If the reordering also subtly changed some other code path we don't benchmark (e.g. paged path with `seqused_k = None`), we wouldn't notice until production. 3. **Not portable guidance.** We can't tell future contributors "move `batch_coord` earlier to get +6%" because the mechanism isn't a deliberate optimization. ### TODO before opening / merging - [ ] Compare SASS between this branch and `paged-kv-v2` for the long-seqlen non-causal dense path; identify which instructions changed and whether the gain is attributable to a specific scheduling/allocation difference. - [ ] Confirm the paged path with `seqused_k = None` isn't regressed (benchmark_attn.py doesn't exercise paged; add a quick paged-bench harness or run the paged tests under timing). - [ ] Decide whether to keep the reorderings (if the mechanism is understood) or revert the non-essential ones (declaration move) to isolate the correctness change from the accidental perf gain. ## Other caveat - Host-side prefill in `interface.py` walks zero-KV batch indices in Python. For batch sizes in the thousands with many zero-KV rows this could show up in CPU profiles; current FA4 callers aren't in that regime. - `seqused_q` is still asserted `None` — the kernel would need a similar ternary in the Q-side iteration, out of scope for this PR.
Johnsonms
added a commit
to Johnsonms/flash-attention
that referenced
this pull request
Apr 23, 2026
Rebased cherry-pick of `dbb6c98` from `Johnsonms/persistent-cluster-hd256` on top of `Johnsonms/seqused-k-hd256-v2`. Original branch was based on a pre-merge snapshot of the hd256 tree; base commits were absorbed into the merged hd256 PR Dao-AILab#2412 (`27b4eb9`) and post-merge cleanup Dao-AILab#2487 (`b21e204`). **Stacked PR** — base branch should be `Johnsonms/seqused-k-hd256-v2`, not `main`. Depends on that PR for the `seqused_k` + paged-KV plumbing. ## Change Persistent scheduling amortizes CTA launch overhead by issuing a grid-stride loop over tiles. It was hardcoded off in hd256 since the kernel's inception because the static tile scheduler was cluster- unaware and split 2CTA clusters across independent work tiles, corrupting output. ### Cluster-aware fix - `Sm100FmhaStaticTileSchedulerParams` gains a `cluster_shape_m` constexpr (default 1, so 1CTA kernels are unchanged). - Grid is sized in cluster units: `max_ctas = (sm_count // cluster_shape_m) * cluster_shape_m`; problem size is multiplied by `cluster_shape_m` so `dsl_min` compares apples to apples. - `num_persistent_clusters` replaces `num_persistent_sm` as the grid-stride step, so `advance_to_next_work` advances by one cluster per iteration instead of one CTA. - In `get_current_work`, the CTA rank within the cluster is reconstructed from the launch `block_idx` (`cta_rank = blk_coord[0] % cluster_shape_m`) and spliced back into `mid = m_block * cluster_shape_m + cta_rank`, so both CTAs in a 2CTA cluster land on their half of the same tile. ### Enablement gate (work-per-tile heuristic) Persistent's per-tile cost (cluster-rank reconstruction + grid-stride state) only pays off when work-per-tile is small, i.e. short KV. On B200, persistent wins at short seqlen but regresses dense prefill at long seqlen because launch overhead is already amortized by the hardware at high tile counts. - **Gate:** `... and seqlen_k <= 2048`. Keeps the decode win and holds long-context prefill within noise of the pre-persistent baseline. - Gate is folded into `interface.py`'s compile_key so configs straddling the threshold compile separate kernels with the correct persistent flag baked in. ### Files - `flash_attn/cute/tile_scheduler.py`: add `cluster_shape_m` on `Sm100FmhaStaticTileSchedulerParams` and the matching updates to `Sm100FmhaStaticTileScheduler` and `compute_sm100_fmha_grid`. - `flash_attn/cute/sm100_hd256_2cta_fmha_forward.py`: drop `self.is_persistent = False` hardcode (now honors constructor arg), pass `cluster_shape_mnk` to `compute_grid`, divide `blk_idx[0]` by `cluster_shape_mnk[0]` when constructing `FmhaStaticTileScheduler`. - `flash_attn/cute/sm100_hd256_2cta_fmha_backward_dqkernel.py`: same fix applied to the bwd dQ kernel for consistency (its scheduler is the same `Sm100` static scheduler). - `flash_attn/cute/interface.py`: compute `hd256_is_persistent` above the compile_key (gated on `seqlen_k <= 2048` in addition to causal/cu_seqlens_q/...); add it to the compile_key; use it at `fa_fwd` construction. ## Validation ### Correctness smoke `pytest tests/cute/test_flash_attn.py` on B200, filter combines the 6 `seqused_k` tests (inherited from parent PR), the 6 paged tests (grandparent PR), and the existing d=256 dense subset: ``` -k "seqused_k_hd256_sm100 or paged_seqused_k_hd256_sm100 or seqused_k_zero_hd256_sm100 or paged_hd256_sm100_tma or (test_flash_attn_output and 256-False-0-0.0-False-False)" ``` Result: **90 passed, 78 skipped, 0 failed** — identical pass/skip count to `origin/main`. No new tests added by this commit (scheduler refactor tests are covered via the existing d=256 dense parametrize space, which hits both gate-ON and gate-OFF configurations). ### FWD perf delta vs `origin/main` (TFLOPS mean, 3 runs each) B200, bf16, hdim=256, locked clocks @ 1755 MHz. #### Short seqlen — **gate-ON path** (persistent scheduling active) ##### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 1k | F | 985 | 987 | +0.2% | | 2k | F | 1307 | 1301 | −0.5% | | 1k | T | 592 | 592 | 0.0% | | 2k | T | 910 | 911 | +0.1% | ##### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 1k | F | 1012 | 1015 | +0.3% | | 2k | F | 1329 | 1329 | 0.0% | | 1k | T | 600 | 600 | 0.0% | | 2k | T | 919 | 919 | 0.0% | **Observation:** the persistent path is active but delivers essentially **no speedup** in this bench config (32 Q-heads MHA, 32:2 GQA). The commit message cites +10–25% wins at `seqlen_k=1024` on a different config (8 Q-heads); with 32 Q-heads the tile count per batch is 4× larger and launch overhead is already amortized by the hardware, so the persistent loop's benefit is saturated out. Reproducibility is very tight (variance <1 TFLOPS across 3 runs), so this is not run noise — just "no regression" rather than "a win" at this config. #### Long seqlen — **gate-OFF path** (should match main) ##### MHA 32:32 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1474 | 1479 | +0.3% | | 8k | F | 1582 | 1588 | +0.4% | | 16k | F | 1641 | 1636 | −0.3% | | 32k | F | 1450 | 1442 | −0.6% | | 64k | F | 1417 | 1447 | **+2.1%** | | 128k | F | 1398 | 1423 | +1.8% | | 4k | T | 1215 | 1218 | +0.2% | | 8k | T | 1411 | 1416 | +0.4% | | 16k | T | 1540 | 1540 | 0.0% | | 32k | T | 1552 | 1619 | **+4.3%** | | 64k | T | 1486 | 1469 | −1.1% | | 128k | T | 1363 | 1383 | +1.5% | ##### GQA 32:2 | seqlen | causal | main | this PR | Δ | |-------:|:------:|-----:|--------:|--:| | 4k | F | 1496 | 1504 | +0.5% | | 8k | F | 1601 | 1598 | −0.2% | | 16k | F | 1620 | 1640 | +1.2% | | 32k | F | 1482 | 1539 | **+3.8%** | | 64k | F | 1436 | 1460 | +1.7% | | 128k | F | 1389 | 1363 | −1.9% | | 4k | T | 1242 | 1244 | +0.2% | | 8k | T | 1434 | 1438 | +0.3% | | 16k | T | 1556 | 1557 | +0.1% | | 32k | T | 1624 | 1574 | **−3.1%** | | 64k | T | 1493 | 1493 | 0.0% | | 128k | T | 1373 | 1369 | −0.3% | **Observation:** long-seqlen path is within ±2% on most cells. Same long-seqlen noise pattern as the parent PRs at 32k/64k (batch- quantization zone); deltas swing both directions, no systematic regression. ## Caveats - **Short-seqlen benefit is config-dependent.** The persistent path trades extra per-tile state for launch-overhead amortization; the trade only pays off when work-per-tile is small. At 8 Q-heads the commit's original bench showed +10–25% at 1k; at 32 Q-heads in my bench the benefit collapses to 0% because tile count is already high enough to amortize launches. If the use case is decode-style small-batch / few-head, the gate-ON path is expected to deliver the advertised win. - Backward dQ kernel gets the same cluster-aware scheduler change for consistency but is not exercised by `benchmark_attn.py` in this PR's validation. Forward dense path is the regression-critical one and is covered above. - `cluster_shape_m` currently defaults to 1 so 1CTA kernels are unaffected; the only callers passing `cluster_shape_m > 1` are hd256 forward and hd256 bwd dQ.
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).
3 of 4 tasks
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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.pyInt32importMaskEnumimportflash_attn/cute/mask.pythread_idx()lines in:Sm100FusedMask.apply_maskapply_mask_via_causal_localtests/cute/test_flash_attn.py/SM110from two TODO commentsIntentionally Unchanged
Two Copilot flags were reviewed and left unchanged because they are false positives:
warp_reductioninutils.pyCLC scheduler in
tile_scheduler.pyValidation
pre-commit runpassed cleanly on all three files78 passed, 78 skipped, 0 failedorigin/main:+3.4%,-2.4%,+3.6%) fall in an already noisy region`