Skip to content

[HIP] Add FP4 output mode to indexer_qk_rope_quant_and_cache - #5484

Merged
valarLip merged 1 commit into
mainfrom
feat/indexer-qk-rope-fp4-out
Sep 14, 2026
Merged

valarLip merged 1 commit into
mainfrom
feat/indexer-qk-rope-fp4-out

Conversation

@XiaobingSuper

@XiaobingSuper XiaobingSuper commented Sep 13, 2026 •

Copy link
Copy Markdown
Contributor

Extends indexer_qk_rope_quant_and_cache with an FP4 output mode, so a DSA sparse indexer can keep its keys as packed E2M1 plus e8m0 scales and have flydsl_pa_mqa_logits_fp4 (and its prefill variant) score them straight out of the paged cache. Before this, that op only emitted FP8, so a caller wanting the FP4 scorers had to skip the fusion and write the FP4 index cache itself.

API

FP4 mode is selected by supplying both new optional trailing arguments. Omitting them leaves the FP8 contract untouched.

indexer_qk_rope_quant_and_cache(
    ..., quant_block_size=32, scale_fmt="ue8m0", preshuffle=False,
    q_scale_out: Tensor,     # u8 [num_tokens, k_tiles, 4, 16, round_up(H // 16, 4)]
    kv_cache_scale: Tensor,  # u8 [num_blocks, k_tiles, 4, kv_block_size]
)
# q_out        u8 [num_tokens, n_heads, head_dim // 2]
# kv_cache     u8 [num_blocks, k_tiles, 4, kv_block_size, 16]
# weights_out  q.dtype [num_tokens, n_heads]

weights_out carries the plain weight under FP4, with neither q_scale nor weights_scale folded in. The FP8 path folds q_scale because its per-(token, head) scale has nowhere else to go; under FP4 the MFMA consumes the per-group e8m0 scale directly, so folding is impossible, and weights_scale is left out because the scorer applies its own fp32 scalar while folding it here would round through bf16.

Anything outside head_dim=128, rope_dim=64, group_size=32, kv_block_size=64, n_heads % 16 == 0 is rejected on the host with a specific message.

Verification

  • FP4 byte-exactness: 16/16 configs (tokens 8/32 x heads 32/64 x compute_all_q_rope x is_neox), zero mismatching bytes across all five outputs. Heads=32 exercises the qs_pad zero-padding path, heads=64 the unpadded one.
  • FP8 unchanged: rebuilt baseline and patched sources and diffed the outputs -- 576/576 tensors (68,063,232 elements) byte-identical across 192 configurations spanning both scale_fmt values, both preshuffle values, both is_neox, both compute_all_q_rope, and three token counts.
  • End-to-end: feeding the kernel's own bytes into flydsl_pa_mqa_logits_fp4 at GLM-5.2 shapes (H=32, D=128, kv_block=64, block_k=256) gives cosine 1.000000 and top-64 overlap 1.00 for next_n 1 and 4, max abs error 1.5e-8. Against independent bf16 ground truth the same pipeline scores cosine 0.980, the expected FP4 quantization loss.
  • black and ruff clean.

One deliberate deviation

The argmin oracle in op_tests/test_flydsl_pa_mqa_logits_fp4.py cannot be byte-reproduced by any AMD FP4 kernel: measured against the hardware v_cvt_scalef32_pk_fp4_f32 converter over 524,288 random bf16 values it differs on 5.87% of nibbles (27,774 are +0 vs -0, identical after dequant; 2,989 are exact midpoints where argmin rounds toward -inf and the hardware rounds half-to-even). AITER's own f32_to_mxfp4 matches the hardware on all 524,288. This implements the hardware semantics, matching _rope_rotate_activation_fp4quant, rmsnorm_rope_rotate_activation_fp4quant_kvcache and per_1x32_f4_quant_hip, so it emits -0 codes where that oracle produces +0.

Follow-up

The preshuffle offset arithmetic now exists in both cache_kernels.cu and dsv4_rotate_quant.cu under distinct names. That ABI-critical math living in two translation units is a drift hazard; a shared header would be the right cleanup.

The consumer side is ROCm/ATOM#perf-sparse-indexer-fp4, which depends on this branch.

@XiaobingSuper
XiaobingSuper requested review from a team and a lite review from Copilot September 13, 2026 12:48

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 was unable to review this pull request because the user who requested the review has reached their quota limit.

@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 5484 --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 Add FP4 output mode to indexer_qk_rope_quant_and_cache [HIP] Add FP4 output mode to indexer_qk_rope_quant_and_cache Sep 13, 2026
@github-actions github-actions Bot added the HIP label Sep 13, 2026
@zufayu
zufayu requested a review from junhaha666 September 14, 2026 01:43
Copilot AI review requested due to automatic review settings September 14, 2026 01:55
@XiaobingSuper
XiaobingSuper force-pushed the feat/indexer-qk-rope-fp4-out branch from 5de8056 to 68ccc1b Compare September 14, 2026 01:55

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 was unable to review this pull request because the user who requested the review has reached their quota limit.

The fused indexer op only emitted FP8, so callers wanting the FP4 paged
MQA-logits kernels had to skip the fusion and write the FP4 index cache
themselves. Supplying q_scale_out and kv_cache_scale now switches the op
to the packed e2m1 + e8m0 layout that flydsl_pa_mqa_logits_fp4 and its
prefill variant read directly.

FP4 mode keeps the weight unscaled: that kernel applies the per-group Q
scale inside the MFMA and takes weights_scale as its own fp32 scalar, so
folding either one here would be wrong and would round through bf16.

Omitting the two scale buffers leaves the FP8 contract byte for byte.
@XiaobingSuper
XiaobingSuper force-pushed the feat/indexer-qk-rope-fp4-out branch from 68ccc1b to 305a0f8 Compare September 14, 2026 02:12
Copilot AI review requested due to automatic review settings September 14, 2026 02:12

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 was unable to review this pull request because the user who requested the review has reached their quota limit.

@valarLip
valarLip merged commit ede5c6b into main Sep 14, 2026
56 checks passed
@valarLip
valarLip deleted the feat/indexer-qk-rope-fp4-out branch September 14, 2026 07:12
sammysun0711 pushed a commit to sammysun0711/aiter that referenced this pull request Sep 16, 2026
The fused indexer op only emitted FP8, so callers wanting the FP4 paged
MQA-logits kernels had to skip the fusion and write the FP4 index cache
themselves. Supplying q_scale_out and kv_cache_scale now switches the op
to the packed e2m1 + e8m0 layout that flydsl_pa_mqa_logits_fp4 and its
prefill variant read directly.

FP4 mode keeps the weight unscaled: that kernel applies the per-group Q
scale inside the MFMA and takes weights_scale as its own fp32 scalar, so
folding either one here would be wrong and would round through bf16.

Omitting the two scale buffers leaves the FP8 contract byte for byte.

Signed-off-by: Xiake Sun <xiake.sun@amd.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants