You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
Kimi-K3 packed-varlen prefill (BF16, D_QK=192, D_V=128, H=12, Sq=4096, Sk~8K–128K) is long-K bound on the current gfx942 ASM kernel. vLLM already calls aiter.flash_attn_varlen_func. This PR speeds that path up without changing the public varlen API.
Technical Details
New ASM kernel fmha_fwd_hd192x128_bf16_rtna_group_splitkv (splits 2–8) plus a HIP LSE combine.
fmha_v3_varlen_fwd takes a trailing num_splits=0. 0 is auto, 1 keeps the unsplit kernel, 2–8 force that split. Existing callers that omit it are unchanged.
Auto uses split-3 when the kernel contract matches, Sk ≥ 8192, and Q_tiles * heads ≤ 2 * CU. Other cases stay unsplit.
Forced split uses the caller’s real out=, causal, window, and padded cu_seqlens. Incompatible requests raise.
flash_attn_varlen_func is unchanged and still auto-selects via num_splits=0.
Test Plan
op_tests/test_fmha_gfx942_asm_splitkv.py on MI325X
Correctness vs a CPU torch reference, including public auto-select at 8191/8192, preallocated out=, return_lse=False, causal rejection, CUDA graph on the split path, and torch.compile(fullgraph=True).
Relevant target shape Sq=4096 Sk=42700 H=12 on MI325X:
All standard extended tests (excludes ci:atom_full)
Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 5592 --add-label <label>
PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to opt this PR out.
github-actionsBot
changed the title
[ASM][gfx942] Optimize long-context FMHA with split-KV
[ASM] [HIP] [JIT] [gfx942] Optimize long-context FMHA with split-KV
Sep 16, 2026
The reason will be displayed to describe this comment to others. Learn more.
🔵 Needs a closer look
Two moderate review findings remain unresolved.
Pull request overview
Adds gfx942 split-KV BF16 FMHA for long-context varlen attention while preserving the public API.
Changes:
Adds split-KV ASM dispatch and HIP LSE combination.
Registers internal operations, interfaces, manifests, and JIT sources.
Adds correctness tests and performance benchmarks.
Review findings:
Moderate (1 vote):asm_mha_varlen_fwd.cu selects split-3 for seqlen_k 8161–8191 despite the documented 8192 threshold; align the contract and add boundary coverage.
Moderate (1 vote): The forced split path rejects max_seqlen_k == 0 before empty-K handling; handle empty K first or skip the partition check.
File summaries
File
Summary
op_tests/test_fmha_gfx942_asm_splitkv.py
Tests split-KV correctness, dispatch, LSE, graphs, and compilation.
[verified] For seqlen_k values 8161–8191, this ceil expression produces 256 and auto-selects split-3, even though the PR contract says the gate begins at Sk >= 8192 and the comment describes a floor of 256 tiles. The added boundary test only covers 8160, so this earlier split of a partial tile is untested and contradicts the documented heuristic. Author must align the implementation and contract by using an explicit 8192 boundary (or revise the stated threshold) and add both sides of it.
When the internal forced path is called with max_seqlen_k == 0, this check runs before the existing empty-K branch below. Both tile counts are zero, so it raises instead of returning the zero output/LSE that the unsplit path supports. Skip this partition check for empty K (or handle the empty case before selecting split-KV).
The reason will be displayed to describe this comment to others. Learn more.
🟡 Changes recommended
Kernel status validation and fake-output dtype handling must be corrected before approval.
Get a fresh assessment by requesting another Copilot review.
Review details
Suppressed comments (1)
aiter/ops/mha.py:1170
[verified] This new fake delegates to gen_fmha_v3_varlen_fwd_fake_tensor, whose return_dropout_randval=False branch creates the empty third output with default float32; the real C++ implementation creates that empty p tensor with q.options() (BF16) at csrc/py_itfs_cu/asm_mha_varlen_fwd.cu:457. torch.compile can therefore infer the wrong dtype for this new op, and the current test does not inspect the third or fourth outputs. Author must make the fake p dtype match the real op (preferably in the shared fake) and validate all four outputs.
The reason will be displayed to describe this comment to others. Learn more.
🔵 Needs a closer look
Unresolved fake-tensor dtype and benchmark validation issues require human review.
Review details
Suppressed comments (3)
aiter/ops/mha.py:1174
The internal op's fake reuses gen_fmha_v3_varlen_fwd_fake_tensor, whose empty p tensor defaults to float32, but the real C++ implementation creates p with q.options() (BF16 for this split-KV path). torch.compile therefore sees a different dtype for the third return value and can fail or propagate an invalid graph when that value is consumed. Make the fake return an empty p with dtype=q.dtype (or correct the shared fake helper) and add an assertion covering all four outputs.
return gen_fmha_v3_varlen_fwd_fake_tensor(
q,
k,
v,
cu_seqlens_q,
This gate can pass invalid output: if either kernel returns NaN/Inf, cosine_difference becomes NaN and NaN >= 1e-4 is false; if both outputs have zero norm, the denominator is also zero. Check finiteness and require a finite, positive denominator before accepting the comparison, otherwise the benchmark can report timings for a broken kernel.
The description reports 26/26 tests passed, but this file defines 27 collected cases: 1 single case, 8 boundary cases, 7 split-count cases, 3 dispatch cases, and 8 additional single-case tests. Please correct the reported result or explain which case was not run so the validation claim is reproducible.
[verified] The default CLI runs only --lse 1, while the production default is return_lse=False; that path passes nullptr to the new combine kernel and is not compared with the independent CPU reference for non-empty K. The benchmark's split-vs-unsplit cosine check cannot catch a shared error. Author must add a non-LSE correctness case to the default test matrix.
[verified] cosine_difference is not cosine distance: this computes 1 - 2*dot/(||reference||² + ||actual||²), while cosine difference is 1 - dot/(||reference||*||actual||). If the two outputs have different norms, the correctness gate and the reported metric measure a different quantity, so the claimed value is misleading. Author must compute the actual cosine distance or rename this metric and retune its threshold.
Cache attention references across split-count tests
op_tests/test_fmha_gfx942_asm_splitkv.py:173
[verified] run_torch is executed at line 173 for every (shape, num_splits, return_lse) combination, even though the reference is independent of num_splits. For the 4096×42700 ticket shape this repeats the full CPU attention reference seven times (once for each split), adding a very large, avoidable test-time cost and risking the 60-minute op-test timeout. Author must cache the reference per shape/return_lse or restructure the loop so it is computed once before iterating over split counts.
#5592 now exposes num_splits on fmha_v3_varlen_fwd and has dropped the narrow
fmha_v3_varlen_splitkv_fwd op. A tuned split count therefore travels with the
caller's out=, mask and padding rather than around them, which closes the
finding that a tuned row left a preallocated out= buffer unwritten.
Dispatch still screens the call against the kernel contract first, because
forcing a split C++ cannot serve is a TORCH_CHECK and a tuned row is a hint.
The tuner measures its asm_v3 candidates through the same entry point, and the
CSV-override test now asserts the buffer the caller passed.
Co-authored-by: Cursor <cursoragent@cursor.com>
#5592 now exposes num_splits on fmha_v3_varlen_fwd and has dropped the narrow
fmha_v3_varlen_splitkv_fwd op. A tuned split count therefore travels with the
caller's out=, mask and padding rather than around them, which closes the
finding that a tuned row left a preallocated out= buffer unwritten.
Dispatch still screens the call against the kernel contract first, because
forcing a split C++ cannot serve is a TORCH_CHECK and a tuned row is a hint.
The tuner measures its asm_v3 candidates through the same entry point, and the
CSV-override test now asserts the buffer the caller passed.
mp_tuner, base_tuner's measurement_kwargs and _read_csv fix, the
FileBaton zombie check and their tests return to #5592's version. They
now live in #5846, where the other tuners can use them, and the MHA
tuner no longer needs them. base_tuner keeps only the gpu_model filter
in run_config that tuned MHA rows rely on.
Co-authored-by: Cursor <cursoragent@cursor.com>
mp_tuner, base_tuner's measurement_kwargs and _read_csv fix, the
FileBaton zombie check and their tests return to #5592's version. They
now live in #5846, where the other tuners can use them, and the MHA
tuner no longer needs them. base_tuner keeps only the gpu_model filter
in run_config that tuned MHA rows rely on.
Public flash_attn_varlen_func must stay unsplit below Sk 8192 and
match split-3 at the gate. Skip the sweep on MI308, which has no
split-KV code object.
assertAllclose now stops the run on a bad O or LSE. CUDA graph captures
the auto-selected split shape, and return_lse=False is checked against
the torch reference.
#5592 now exposes num_splits on fmha_v3_varlen_fwd and has dropped the narrow
fmha_v3_varlen_splitkv_fwd op. A tuned split count therefore travels with the
caller's out=, mask and padding rather than around them, which closes the
finding that a tuned row left a preallocated out= buffer unwritten.
Dispatch still screens the call against the kernel contract first, because
forcing a split C++ cannot serve is a TORCH_CHECK and a tuned row is a hint.
The tuner measures its asm_v3 candidates through the same entry point, and the
CSV-override test now asserts the buffer the caller passed.
mp_tuner, base_tuner's measurement_kwargs and _read_csv fix, the
FileBaton zombie check and their tests return to #5592's version. They
now live in #5846, where the other tuners can use them, and the MHA
tuner no longer needs them. base_tuner keeps only the gpu_model filter
in run_config that tuned MHA rows rely on.
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
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.
Motivation
Kimi-K3 packed-varlen prefill (BF16, D_QK=192, D_V=128, H=12, Sq=4096, Sk~8K–128K) is long-K bound on the current gfx942 ASM kernel. vLLM already calls
aiter.flash_attn_varlen_func. This PR speeds that path up without changing the public varlen API.Technical Details
fmha_fwd_hd192x128_bf16_rtna_group_splitkv(splits 2–8) plus a HIP LSE combine.fmha_v3_varlen_fwdtakes a trailingnum_splits=0.0is auto,1keeps the unsplit kernel,2–8force that split. Existing callers that omit it are unchanged.Q_tiles * heads ≤ 2 * CU. Other cases stay unsplit.out=, causal, window, and padded cu_seqlens. Incompatible requests raise.flash_attn_varlen_funcis unchanged and still auto-selects vianum_splits=0.Test Plan
op_tests/test_fmha_gfx942_asm_splitkv.pyon MI325Xop_tests/op_benchmarks/bench_fmha_gfx942_asm_splitkv.py --sq 4096 --sk 42700 --heads 12 --splits 3Test Result
Correctness vs a CPU torch reference, including public auto-select at 8191/8192, preallocated
out=,return_lse=False, causal rejection, CUDA graph on the split path, andtorch.compile(fullgraph=True).Relevant target shape Sq=4096 Sk=42700 H=12 on MI325X:
Submission Checklist