Skip to content

[ASM] [HIP] [JIT] [gfx942] Optimize long-context FMHA with split-KV - #5592

Open
amd-yashagar wants to merge 10 commits into
mainfrom
feat/gfx942-fmha-hd192-split-asm
Open

amd-yashagar wants to merge 10 commits into
mainfrom
feat/gfx942-fmha-hd192-split-asm

Conversation

@amd-yashagar

@amd-yashagar amd-yashagar commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

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

  • 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
  • op_tests/op_benchmarks/bench_fmha_gfx942_asm_splitkv.py --sq 4096 --sk 42700 --heads 12 --splits 3

Test 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, and torch.compile(fullgraph=True).

Relevant target shape Sq=4096 Sk=42700 H=12 on MI325X:

  • unsplit ASM 3.065 ms
  • split-3 2.269 ms (26% faster)
  • cosine difference 5.087e-6
  • combine is ~0.020 ms

Submission Checklist

@amd-yashagar
amd-yashagar requested review from a team and a lite review from Copilot September 16, 2026 14:02
@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs; main branch always runs both MI35X and MI300X
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all 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-actions github-actions Bot 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

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

🔵 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.
op_tests/op_benchmarks/bench_fmha_gfx942_asm_splitkv.py Benchmarks unsplit versus split-KV execution.
hsa/gfx942/fmha_v3_fwd/fmha_splitkv.csv Registers split-KV code-object metadata.
csrc/py_itfs_cu/asm_mha_varlen_fwd.cu Implements split-KV selection and dispatch.
csrc/kernels/fmha_fwd_v3_splitkv_combine.cu Combines partial outputs and LSE values.
csrc/include/torch/mha_v3_varlen_fwd.h Declares the internal split-KV operation.
csrc/include/rocm_ops.hpp Registers the internal operation.
csrc/include/mha_fwd.h Declares split-KV interfaces and constants.
csrc/cpp_itfs/mha_fwd.cu Launches the split-KV ASM kernel.
aiter/ops/mha.py Registers Python and fake-tensor bindings.
aiter/jit/optCompilerConfig.json Adds the combine kernel to JIT builds.
Review details

Suppressed comments (2)

csrc/py_itfs_cu/asm_mha_varlen_fwd.cu:29

  • [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.
    const int kv_tiles = (seqlen_k + kHd192SplitKvTile - 1) / kHd192SplitKvTile;
    if(kv_tiles < kHd192SplitKvAutoMinKvTiles)

csrc/py_itfs_cu/asm_mha_varlen_fwd.cu:503

  • 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).
    if(use_hd192_splitkv)
  • Files reviewed: 11/12 changed files
  • Comments generated: 0
  • Review effort level: Lite

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Copilot AI review requested due to automatic review settings September 16, 2026 14:17

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

🟡 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.
    return gen_fmha_v3_varlen_fwd_fake_tensor(
  • Files reviewed: 11/12 changed files
  • Comments generated: 1
  • Review effort level: Lite

Comment thread csrc/py_itfs_cu/asm_mha_varlen_fwd.cu
Copilot AI review requested due to automatic review settings September 16, 2026 14:51
@amd-yashagar
amd-yashagar force-pushed the feat/gfx942-fmha-hd192-split-asm branch 2 times, most recently from 29694f8 to 33d07c2 Compare September 16, 2026 14:59

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

🔵 Needs a closer look

Low-level ASM, HIP combination, and dispatch changes require final human review.

Review details
  • Files reviewed: 11/12 changed files
  • Comments generated: 0 new
  • Review effort level: Lite

Copilot AI review requested due to automatic review settings September 16, 2026 15:01

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

🔵 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,

op_tests/op_benchmarks/bench_fmha_gfx942_asm_splitkv.py:92

  • 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.
    cosine_difference = 1.0 - 2.0 * (
        reference.double() * actual.double()
    ).sum().item() / (
        reference.double().square().sum().item() + actual.double().square().sum().item()
    )

op_tests/test_fmha_gfx942_asm_splitkv.py:146

  • 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.
@pytest.mark.parametrize("num_splits", range(2, 9))
def test_splitkv_counts(num_splits):
    sq, sk, h = 129, 2048, 4
    q, k, v, cu_q, cu_k = _make_packed(sq, sk, h, seed=100 + num_splits)
  • Files reviewed: 11/12 changed files
  • Comments generated: 0 new
  • Review effort level: Lite

@zufayu
zufayu requested a review from amd-ruitang3 September 17, 2026 01:11
Copilot AI review requested due to automatic review settings September 22, 2026 08:43

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

The reviewed tests and benchmark need MI308 capability guards, and auto-selection coverage needs a deterministic dispatch assertion.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 2 Medium severity

Open (2)

Comment thread op_tests/test_fmha_gfx942_asm_splitkv.py Outdated
Comment thread op_tests/test_fmha_gfx942_asm_splitkv.py

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🔵 Needs a closer look

Unresolved moderate test-validation gaps affect split-KV graph, non-LSE correctness, and mismatch detection.

Review effort: Lite
Findings: None

Resolved since last review (2)
Previously missed (1)

In code that hasn't changed since last review

Medium severity Default tests omit non-LSE correctness coverage

op_tests/​test_fmha_gfx942_asm_splitkv.py:359

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

Copilot AI review requested due to automatic review settings September 23, 2026 09:06

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

Unresolved gradient validation and padded-cu, metric, and test-efficiency issues remain.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 2 High severity · 2 Medium severity

Open (4)
Previously missed (2)

In code that hasn't changed since last review

Medium severity Compute true cosine distance or rename the metric

op_tests/​op_benchmarks/​bench_fmha_gfx942_asm_splitkv.py:112

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

Medium severity 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.

Comment on lines +571 to +575
aiter::launch_fmha_fwd_v3_splitkv_combine(
o_parts.data_ptr(),
lse_parts.data_ptr(),
out.data_ptr(),
softmax_lse.numel() == 0 ? nullptr : softmax_lse.data_ptr(),
amd-bartgips added a commit that referenced this pull request Sep 23, 2026
#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>
amd-bartgips added a commit that referenced this pull request Sep 23, 2026
#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.
amd-bartgips added a commit that referenced this pull request Sep 25, 2026
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>
amd-bartgips added a commit that referenced this pull request Sep 28, 2026
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.
Gate auto-select on Sk >= 8192, skip the empty-K partition check, and
keep Triton imported before torch while satisfying ruff/black.
A failed producer must not write out from empty partial buffers.
Match the shared fake p dtype to the real op, reject non-finite bench
comparisons, and sweep the kernel with a torch reference plus markdown tables.
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.
Forced split now sees the caller's out, causal, and padding arguments.
Default 0 keeps existing callers on auto-select.
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.
The check helpers all build the same q/k/v, cu_seqlens, and scale tensors.
Copilot AI lite review requested due to automatic review settings September 30, 2026 08:57
@amd-yashagar
amd-yashagar force-pushed the feat/gfx942-fmha-hd192-split-asm branch from 2760b97 to d7c3e42 Compare September 30, 2026 08:57

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🔵 Needs a closer look

The unresolved test-efficiency issue and broad kernel/API changes require human review before approval.

Review effort: Lite
Findings: 2 High severity · 2 Medium severity

Open (4)

amd-bartgips added a commit that referenced this pull request Oct 1, 2026
#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.
amd-bartgips added a commit that referenced this pull request Oct 1, 2026
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.
Copilot AI lite review requested due to automatic review settings October 5, 2026 09:01

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

softmax_scale=scale,
causal=False,
return_lse=return_lse,
out=out,
Comment on lines +494 to +498
!q_descale_.has_value() && !return_dropout_randval &&
!cu_seqlens_q_padded.has_value() && !cu_seqlens_k_padded.has_value() &&
how_v3_bf16_cvt == 1;
TORCH_CHECK(!split_forced || splitkv_compatible,
"forced gfx942 hd192 split-KV request is incompatible with this input");
Copilot AI lite review requested due to automatic review settings October 6, 2026 07:41

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

@zufayu
zufayu requested review from shay-li77 and removed request for amd-ruitang3 October 9, 2026 05:47

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants