Skip to content

[FlyDSL][gfx950] Optimize FP4 MQA prefill and decode - #5707

Draft
jiacao-amd wants to merge 7 commits into
ROCm:mainfrom
jiacao-amd:jiacao/fp4-prefill-4wave
Draft

jiacao-amd wants to merge 7 commits into
ROCm:mainfrom
jiacao-amd:jiacao/fp4-prefill-4wave

Conversation

@jiacao-amd

@jiacao-amd jiacao-amd commented Sep 20, 2026 •

Copy link
Copy Markdown
Contributor

Summary

  • replace the original FP4 MQA prefill implementation with the validated 4-wave kernel in pa_mqa_logits_fp4_prefill.py
  • group four same-batch query rows per prefill CTA, with one row per wave and one LDS-staged KV page shared by all four waves
  • use one decode kernel for every next_n, with cooperative head waves and no decode-to-prefill routing
  • specialize internally scheduled D128 H64/H128 decode with direct global KV/scale loads and grid-aware 32- or 64-token tiles
  • retain the staged-LDS decode path for external/persistent schedules and unsupported layouts
  • exclude benchmark scripts, raw measurements, and experimental kernel variants from the PR

Dependency

Depends on #5518. This branch should be rebased onto the updated main after #5518 merges so it inherits the corresponding page-stride, 64-bit addressing, and page128 support without duplicating that dependency's history here.

Implementation

Prefill

  • 4 waves per CTA, one query row per wave
  • four waves cooperatively load four KV pages into a single LDS stage
  • LDS-resident KV is reused across the four query rows
  • 32x32x64 FP4 MFMA
  • next-chunk global loads overlap with computation on the final page of the current chunk
  • 17,408 bytes LDS and 126 VGPRs in the measured kernel

Decode

The decode kernel now handles all next_n values directly:

  • head waves cooperate on the head reduction within each CTA
  • internally scheduled D128 H64/H128 layouts read KV and scales directly from global memory into registers
  • a 32-token, 2-wave tile is selected when the complete direct grid is at most 1024 CTAs
  • larger direct grids use 64-token tiles to control launch size
  • the direct path uses a 2D launch and compile-time page/chunk addressing
  • scalar weight loads are selected for the larger grids where they reduce contention
  • external/persistent schedules and unsupported layouts retain the staged-LDS generic path
  • callers can still override block_k and num_warps

For the ordinary next_n=1/2 direct shapes, each CTA owns one token chunk. Extra LDS stages therefore do not form a cross-chunk pipeline, and next_n does not itself create additional waves. Their improvement comes from exposing more grid-level parallelism with 32/64-token tiles and bypassing KV LDS staging; head waves only parallelize the head reduction.

Performance

Measured on MI355X using median GPU kernel time. Decode results use nine paired CUDA Graph rounds with identical inputs and randomized Old/New execution order. Unless noted otherwise, rows use H64, D128, page64, and ragged context lengths averaging approximately half of max_ctx.

Prefill

Shape (B x Q x KV) Old (us) New (us) Speedup
AgentX trace 781.16 610.38 1.280x
1 x 512 x 64K 249.20 200.74 1.241x
1 x 1K x 64K 473.36 377.41 1.254x
1 x 2K x 32K 451.07 356.27 1.266x
1 x 1K x 128K 910.54 693.36 1.313x
4 x 256 x 16K 121.77 93.07 1.308x
8 x 128 x 8K 69.46 50.87 1.366x
1 x 4K x 4K 98.64 56.60 1.743x
1 x 8K x 8K 372.80 213.22 1.748x
2 x 8K x 8K 745.16 423.80 1.758x
1 x 8K x 32K 1519.69 1206.67 1.259x

Decode

Here, Old is the original 256-token/4-wave MXFP4 decode configuration and New is the production configuration selected by the current dispatcher.

Shape (B x next_n x max_ctx) Old MXFP4 (us) New (us) Speedup
2 x 1 x 512 2.510 1.942 1.292x
3 x 1 x 1K 2.539 2.013 1.261x
4 x 1 x 2K 2.628 2.027 1.297x
16 x 1 x 4352 3.303 2.392 1.381x
8 x 1 x 8K 2.730 2.416 1.130x
1 x 1 x 32K 2.641 2.056 1.285x
1 x 1 x 64K 2.697 2.270 1.188x
1 x 1 x 128K 3.372 2.384 1.414x
2 x 2 x 512 2.582 1.961 1.317x
2 x 1 x 768, H128 3.287 2.048 1.605x
32 x 4 x 32K, AgentX/DSv4 MTP4 47.598 23.658 2.012x

The first ten general decode rows have a 1.311x geometric-mean paired speedup. Compared with the original MXFP4 timings published in the first PR table, the current New timings are 1.432x faster geometrically. The AgentX/DSv4 MTP4 row contains 32 requests and 128 total query rows.

Validation

  • Python bytecode compilation passes
  • Black formatting check passes
  • Ruff lint passes
  • git diff --check passes
  • MI355X decode regression: 59 passed
  • coverage includes direct, persistent, and external schedules; next_n 1/2/3/4/5/8; H16-H128; D128/D256; page64/page128; zero and ragged contexts; strided pages/output; and CUDA Graph replay
  • all paired benchmark variants match the reference output
  • prefill fixed-query and variable-query correctness sweeps pass with cosine similarity 1.0

@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 5707 --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.

@jiacao-amd
jiacao-amd force-pushed the jiacao/fp4-prefill-4wave branch 3 times, most recently from c7af86a to 6b301d2 Compare September 20, 2026 23:33
@jiacao-amd jiacao-amd changed the title [FlyDSL][gfx950] Add 4-wave FP4 MQA prefill kernel [FlyDSL][gfx950] Optimize FP4 MQA prefill and decode Sep 20, 2026
@jiacao-amd
jiacao-amd force-pushed the jiacao/fp4-prefill-4wave branch from 6b301d2 to f6bf06e Compare September 20, 2026 23:46
Signed-off-by: jiacao-amd <jiahui.cao@amd.com>
@jiacao-amd
jiacao-amd force-pushed the jiacao/fp4-prefill-4wave branch 6 times, most recently from cd969cc to db82535 Compare September 21, 2026 06:49
Signed-off-by: jiacao-amd <jiahui.cao@amd.com>
@jiacao-amd
jiacao-amd force-pushed the jiacao/fp4-prefill-4wave branch from db82535 to 3ab6dfc Compare September 21, 2026 07:13

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant