Skip to content

[Triton/Gluon] [GFX950] Add split-k support for fp8 mqa logits - #5603

Merged
cagrikymk merged 8 commits into
mainfrom
cagri/fp8_mqa_logits_splitk
Sep 17, 2026
Merged

cagrikymk merged 8 commits into
mainfrom
cagri/fp8_mqa_logits_splitk

Conversation

@cagrikymk

Copy link
Copy Markdown
Contributor

Motivation

For agentic workloads, we might have very large kv length with very few queries. So, enabling splitk help such workloads by splitting the kv length.

Num splits defined as dynamic arg to reduce unexpected compilations during runtime.

Tests are extended to trigger split-k with better coverage, with some cleanup.

Triton 3.7.0+amd.rocm7.2.0.git89002410

workload glm5.x speedup glm5.x TFLOP/s dsv4 speedup dsv4 TFLOP/s
AgentX prefill (512 MB budget)
conc 1, dominant launch 1.843x 997 -> 1836 1.272x 1634 -> 2079
conc 1, remainder 2.280x 748 -> 1705 1.607x 1260 -> 2025
conc 2, dominant launch 1.728x 1027 -> 1775 1.258x 1649 -> 2075
conc 4, dominant launch 2.484x 713 -> 1771 1.698x 1220 -> 2072
conc 8, dominant launch 2.459x 722 -> 1775 1.707x 1220 -> 2083
Long context, few query rows
1 x 512 x 64K 2.277x 612 -> 1394 1.618x 1091 -> 1765
1 x 1K x 64K 1.516x 1042 -> 1580 1.188x 1652 -> 1963
1 x 2K x 32K 1.522x 1044 -> 1589 1.114x 1798 -> 2003
1 x 1K x 128K 1.762x 1007 -> 1775 1.227x 1678 -> 2059
Batched, short per-row window
4 x 256 x 16K 1.488x 947 -> 1409 1.145x 1610 -> 1844
8 x 128 x 8K 1.395x 890 -> 1241 1.093x 1503 -> 1642
Already-shipped shapes
1 x 4K x 4K 1.160x 1171 -> 1357 0.996x 1836 -> 1829
1 x 8K x 8K 0.999x 1767 -> 1765 0.998x 2138 -> 2134
2 x 8K x 8K 1.004x 1891 -> 1899 1.000x 2187 -> 2187
1 x 8K x 32K 1.011x 1908 -> 1929 1.025x 2219 -> 2276

Triton 3.8.0+amd.rocm7.2.0.git111ff227

workload glm5.x speedup glm5.x TFLOP/s dsv4 speedup dsv4 TFLOP/s
AgentX prefill (512 MB budget)
conc 1, dominant launch 1.497x 1017 -> 1523 1.052x 1848 -> 1945
conc 1, remainder 1.890x 747 -> 1412 1.321x 1381 -> 1824
conc 2, dominant launch 2.231x 665 -> 1485 1.500x 1208 -> 1812
conc 4, dominant launch 2.122x 718 -> 1524 1.444x 1313 -> 1897
conc 8, dominant launch 2.102x 722 -> 1519 1.425x 1331 -> 1897
Long context, few query rows
1 x 512 x 64K 2.482x 608 -> 1508 1.602x 1097 -> 1757
1 x 1K x 64K 1.403x 1046 -> 1468 0.927x 1826 -> 1692
1 x 2K x 32K 1.466x 969 -> 1421 1.029x 1785 -> 1836
1 x 1K x 128K 1.487x 1007 -> 1497 0.971x 1865 -> 1811
Batched, short per-row window
4 x 256 x 16K 1.273x 976 -> 1242 0.912x 1785 -> 1628
8 x 128 x 8K 1.194x 910 -> 1086 0.864x 1602 -> 1384
Already-shipped shapes
1 x 4K x 4K 1.162x 1067 -> 1239 0.989x 1627 -> 1609
1 x 8K x 8K 1.000x 1503 -> 1502 1.002x 2015 -> 2018
2 x 8K x 8K 1.003x 1686 -> 1690 1.000x 2107 -> 2108
1 x 8K x 32K 0.995x 1674 -> 1666 0.998x 2176 -> 2173

@cagrikymk
cagrikymk requested a review from a team September 16, 2026 20:33
@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 5603 --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 [TRITON][GLUON][GFX950] Add split-k support for fp8 mqa logits [Triton/Gluon] [GFX950] Add split-k support for fp8 mqa logits Sep 16, 2026
@cagrikymk
cagrikymk requested a review from vgokhale September 16, 2026 20:34
vgokhale
vgokhale previously approved these changes 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.

🟡 Changes recommended

Critical compatibility and correctness issues remain unresolved.

Get a fresh assessment by requesting another Copilot review.

Pull request overview

Adds dynamic split-K support for gfx950 FP8 MQA logits and expands large-shape correctness coverage.

Changes:

  • Adds gfx950 split-K dispatch and 2D launch geometry.
  • Partitions KV processing across runtime splits.
  • Extends split-K and large-shape tests.
File summaries
File Summary
op_tests/triton_tests/attention/test_fp8_mqa_logits.py Expands split-K correctness coverage.
aiter/ops/triton/attention/fp8_mqa_logits.py Selects split counts and launches updated kernels.
aiter/ops/triton/_gluon_kernels/gfx950/attention/fp8_mqa_logits.py Implements KV splitting and split-aware masking/stores.

Critical issues remain: gfx1250 receives an unsupported argument, and odd query lengths can cause duplicate writes and a race. Tuning thresholds are also hardcoded instead of using configuration.

Review details

Suppressed comments (1)

aiter/ops/triton/attention/fp8_mqa_logits.py:227

  • MIN_BLOCK_M2_WGS is another new gfx950-specific tuning threshold hardcoded in Python. Please put this threshold in the architecture tuning configuration and resolve it through the existing config machinery rather than baking it into the wrapper.
            MIN_BLOCK_M2_WGS = 1024
  • Files reviewed: 3/3 changed files
  • Comments generated: 3
  • Review effort level: Lite

💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread aiter/ops/triton/attention/fp8_mqa_logits.py
Comment thread aiter/ops/triton/attention/fp8_mqa_logits.py Outdated
Comment thread aiter/ops/triton/attention/fp8_mqa_logits.py
@zufayu
zufayu requested review from a team and vgokhale September 17, 2026 01:10
@cagrikymk
cagrikymk merged commit a84bd36 into main Sep 17, 2026
71 checks passed
@cagrikymk
cagrikymk deleted the cagri/fp8_mqa_logits_splitk branch September 17, 2026 03:42
vgokhale added a commit that referenced this pull request Sep 17, 2026
…ing fixes, fp8 MQA logits split-k and DSv4 tunings (#5573, #5295, #5558, #5603, #5627, #5485) (#5638)

Cherry-picks six already-merged `main` PRs onto `release/v0.1.22` for the `v0.1.22.post1` post release.

| PR | `main` commit | Backport commit | What |
|---|---|---|---|
| #5485 | `9252f4672` | `155534984` | Extend the DeepSeek-V4 a8w8 blockscale GEMM tunings for gfx950 (tuning CSV only) |
| #5573 | `22d2c7c91` | `6b23ba866` | Pad the MXFP4 A4W4 MoE sort extent to a block_size multiple (fixes a HIP illegal memory access) |
| #5295 | `972c8e1fd` | `7d68b0edb` | Skip invalid expert IDs in MoE sorting |
| #5558 | `3fdfca11e` | `dd83a9d17` | Fix MoE routing kernel compile failure |
| #5603 | `a84bd368c` | `a96461997` | Add split-k support for fp8 MQA logits on gfx950 |
| #5627 | `f5ed7dc54` | `a41214712` | Follow-up to #5603: drop chunking when summation folding is unavailable (fixes Triton 3.6 compile) |

Original PRs:
- #5485: #5485
- #5573: #5573
- #5295: #5295
- #5558: #5558
- #5603: #5603
- #5627: #5627

To be published as `v0.1.22.post1` once merged (tag on the merge commit, release automation builds the wheel set).
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.

4 participants