Repository navigation
[AMD] Fix Load and Inference of MLA models with Quark PTPC FP8 attention on ROCm - #28734
Conversation
…quantization
Five fixes to support loading and inference of quark-quantized GLM-5.2
models (mxfp4 weights, per-token-per-channel FP8 attention) on ROCm/gfx95:
1. fp8_utils: channel_quant_to_tensor_quant - reshape 1D per-channel scale
[N] to [N,1] before multiplying against a 2D weight tensor, fixing a
broadcast shape mismatch during kv_b_proj post-load processing.
2. fp8_utils: apply_fp8_ptpc_linear - unpack pre-quantized input tuple as
input[0], input[1] to tolerate both (fp8, scale) 2-tuples and
(fp8, scale, bf16) 3-tuples produced by the DSA fused-RMSNorm+quant path.
3. quark_w8a8_fp8: apply_weights - handle pre-quantized tuple inputs from
fused RMSNorm+quant kernels on the aiter gfx95 path:
- 3-tuple (fp8, scale, bf16): use the pre-computed bf16 copy directly
(produced by fused_qkv_a_proj_with_mqa + DSA; avoids re-quantization).
- 2-tuple (fp8, group_scale): dequantize fp8 back to bf16 using the
group scale (group_size=128, handling both transposed and row-major
scale layouts), then fall through to apply_fp8_linear.
4. dsa_indexer: _get_q_k_bf16 - unwrap (fp8, scale, bf16) 3-tuple to its
bf16 element before passing hidden states to the unquantized wk layer.
5. dsa_indexer: _get_k_bf16 - same unwrap fix for the key-only forward path.
Validated: GSM8K 5-shot scores 93.93% flexible-extract / 93.71% strict-match
on GlmMoeDsaForCausalLM with tp_size=4 on MI355X.
There was a problem hiding this comment.
Code Review
This pull request introduces support for handling fused RMSNorm and FP8 group quantization outputs (which can be 2-tuples or 3-tuples containing unquantized bf16 tensors) across DSA indexing and linear layers. The review feedback highlights potential TypeError crashes in dsa_indexer.py when x is a 2-tuple, suggesting that it should be explicitly dequantized to bfloat16 rather than left as a tuple. Additionally, the reviewer recommends a more idiomatic PyTorch approach using unsqueeze(-1) in a loop to reshape the scale tensor in fp8_utils.py instead of using tuple shape arithmetic.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
|
/rerun-failed-ci |
|
All PR check failures are unrelated to PR chagnes:
@HaiShaw Could you please help take a look and merge when you get a chance, thanks! |
Arist12
left a comment
There was a problem hiding this comment.
Nice fix — we've been deploying MXFP4 GLM checkpoints on gfx950 recently and ended up in these same code paths, so I read through this. The three fixes look correct and minimal to me, and the test file passes on our side (MI355X, rocm/sgl-dev:v0.5.18-rocm724-mi35x-20260829). One small note inline.
| ) | ||
|
|
||
|
|
||
| def is_quark_w8a8_fp8_layer(layer) -> bool: |
There was a problem hiding this comment.
is_quark_w8a8_fp8_layer has no production caller on this branch — it appears only here, in __all__, and in the two comments above the assert; schemes/__init__.py doesn't re-export it and the symbol isn't on main. The only caller is the new test.
Not a bug: what actually keeps that assert from firing is _is_block_scale_fp8() in deepseek_common/utils.py, which already excludes per-channel FP8 from the fused kernel by weight_scale shape. So the behaviour is correct as written — it's just that the guard the comments describe isn't wired up, and if _is_block_scale_fp8 were ever relaxed the assert would fire with the intended protection still not connected.
Either hooking this into the fused-path gate, or dropping it and pointing the assert comment at _is_block_scale_fp8, would make the stated invariant match the code. Happy either way — just flagging it so it isn't lost.
There was a problem hiding this comment.
Good catch! Removed the outdated function and fixed references to it in comments/assert messages.
|
/rerun-failed-ci |
|
Hi @HaiShaw @kkHuang-amd, could you please help take a look and merge when you get a chance? All failing CI tests are unrelated to this PR's changes:
|
|
/rerun-failed-ci |
|
@ColinZ22 could you check if the `PR Test Base / base-c-test-8-gpu-b300 test failure is related or not? |
The failure is caused by Kimi-K3 MegaMoE scheduler abort during init on B300; The model ( Update: confirmed that this failure is caused by upstream changes and now fixed (PR #41138), pr-test-finish already passing after rebase and rerun. |
|
Hi @HaiShaw @kkHuang-amd, could you please help take a look again when you get a chance? Thanks! |
Motivation
Loading and running MLA models with quark-quantized MXFP4 weights + per-token-per-channel FP8 attention (like
amd/GLM-5.2-Quark-MXFP4-AttnFP8) on ROCm/gfx95 fails with multiple errors.The root cause is that the aiter gfx95
_gfx95_quant_format="fp8"path fuses RMSNorm + FP8 quantization into a single kernel (fused_rms_fp8_group_quant), passing pre-quantized tuples instead of plain bf16 tensors downstream. This pattern was already handled for compressed_tensors attention quantization across three prior PRs:apply_fp8_ptpc_linearrouting incompressed_tensors_w8a8_fp8.pyfor pre-quantized tuple inputsisinstance(input, tuple)branch inapply_fp8_ptpc_linearitself(fp8, scale, bf16)passthrough fromfused_rms_fp8_group_quantincommunicator.py, and the corresponding unwrap indsa_indexer._project_and_scale_head_gatesNone of these extended to quark attention quantization, and this PR closes that gap.
Modifications
fp8_utils.py:channel_quant_to_tensor_quantReshape 1D per-channel scale
[N]->[N, 1]before multiplying against the 2D weight. Quark'skv_b_projhas a per-channel FP8 weight with a 1D scale vector that failed to broadcast correctly against the[N, K]weight.fp8_utils.py:apply_fp8_ptpc_linearUnpack pre-quantized input tuple as
input[0], input[1]instead of destructuring, to tolerate both 2-tuples and 3-tuples (the 3-tuple was introduced in #22258 but the unpack assumed exactly 2 elements).quark_w8a8_fp8.py:apply_weightsMirror the handling that #12181 added to
compressed_tensors_w8a8_fp8.apply_weights, extended for quark's weight layout and both tuple formatsNote: unlike compressed_tensors which stores
shuffle(w), quark storesshuffle(w).t(), making the weight layout incompatible withapply_fp8_ptpc_linear/gemm_a8w8_bpreshuffle; So we dequantize to bf16 and use the existingapply_fp8_linearpath without requiring weight storage changes.dsa_indexer.py:_get_q_k_bf16and_get_k_bf16Mirror the unwrap added in #22258 for
_project_and_scale_head_gates, applied to the two remaining methods that also receivexand callself.wk.Validation
Accuracy is evaluated using lm_eval on MI355X.
Model:
amd/GLM-5.2-Quark-MXFP4-AttnFP8Reproduction
Launch server
SGLANG_USE_AITER=1 sglang serve \ --model-path amd/GLM-5.2-Quark-MXFP4-AttnFP8 \ --trust-remote-code --tp 8 \ --json-model-override-args '{"qk_rope_head_dim": 64}' \ --speculative-algorithm EAGLE \ --speculative-num-steps 5 \ --speculative-eagle-topk 1 \ --speculative-num-draft-tokens 6 \ --mem-fraction-static 0.8 \ --chunked-prefill-size 16384 \ --watchdog-timeout 1200 \ --model-loader-extra-config '{"enable_multithread_load": true, "num_threads": 8}' \ --host 127.0.0.1 --port 30000Accuracy
lm_eval --model local-completions \ --model_args model=amd/GLM-5.2-Quark-MXFP4-AttnFP8,base_url=http://localhost:30000/v1/completions,num_concurrent=128,max_retries =10,max_gen_toks=2048,timeout=60000 \ --batch_size auto --tasks gsm8k --num_fewshot 8Performance
Benchmarked on MI350X, tp=4,
sglang.bench_servingrandom workload (1k input / 1k output, 128 prompts, concurrency 64), averaged across 5 runs.CI States
Latest PR Test (Base): ✅ Run #36170599763
Latest PR Test (Extra): ❌ Run #36170599646
Latest PR Test (AMD ROCm 10): ❌ Run #36170600003