Skip to content

[AMD] Fix Load and Inference of MLA models with Quark PTPC FP8 attention on ROCm - #28734

Merged
HaiShaw merged 27 commits into
sgl-project:mainfrom
ColinZ22:glm52-mxfp4-ptpcfp8-attn-fixes
Sep 29, 2026
Merged

HaiShaw merged 27 commits into
sgl-project:mainfrom
ColinZ22:glm52-mxfp4-ptpcfp8-attn-fixes

Conversation

@ColinZ22

@ColinZ22 ColinZ22 commented Jun 19, 2026 •

Copy link
Copy Markdown
Contributor

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:

None of these extended to quark attention quantization, and this PR closes that gap.

Modifications

fp8_utils.py: channel_quant_to_tensor_quant
Reshape 1D per-channel scale [N] -> [N, 1] before multiplying against the 2D weight. Quark's kv_b_proj has 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_linear
Unpack 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_weights
Mirror the handling that #12181 added to compressed_tensors_w8a8_fp8.apply_weights, extended for quark's weight layout and both tuple formats
Note: unlike compressed_tensors which stores shuffle(w), quark stores shuffle(w).t(), making the weight layout incompatible with apply_fp8_ptpc_linear/gemm_a8w8_bpreshuffle; So we dequantize to bf16 and use the existing apply_fp8_linear path without requiring weight storage changes.

dsa_indexer.py: _get_q_k_bf16 and _get_k_bf16
Mirror the unwrap added in #22258 for _project_and_scale_head_gates, applied to the two remaining methods that also receive x and call self.wk.

Validation

Accuracy is evaluated using lm_eval on MI355X.

Model: amd/GLM-5.2-Quark-MXFP4-AttnFP8

Task flexible-extract strict-match
GSM8K 93.93% ± 0.66% 93.71% ± 0.67%

Reproduction

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 30000

Accuracy

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 8

Performance

Benchmarked on MI350X, tp=4, sglang.bench_serving random workload (1k input / 1k output, 128 prompts, concurrency 64), averaged across 5 runs.

Checkpoint Config Output tput (tok/s) Median TPOT (ms) Median ITL (ms)
amd/GLM-5.2-MXFP4 EAGLE spec-decoding 2524 15.74 14.02
amd/GLM-5.2-Quark-MXFP4-AttnFP8 (this PR enables) EAGLE spec-decoding 2721 (+7.8%) 14.83 (-5.8%) 13.47 (-3.9%)
amd/GLM-5.2-MXFP4 No spec-decoding 1941 32.63 32.48
amd/GLM-5.2-Quark-MXFP4-AttnFP8 (this PR enables) No spec-decoding 2187 (+12.7%) 28.92 (-11.4%) 28.76 (-11.5%)

CI States

Latest PR Test (Base): ✅ Run #36170599763
Latest PR Test (Extra): ❌ Run #36170599646
Latest PR Test (AMD ROCm 10): ❌ Run #36170600003

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

@gemini-code-assist gemini-code-assist Bot 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.

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.

Comment thread python/sglang/srt/layers/attention/dsa/dsa_indexer.py Outdated
Comment thread python/sglang/srt/layers/attention/dsa/dsa_indexer.py Outdated
Comment thread python/sglang/srt/layers/quantization/fp8_utils.py Outdated
@ColinZ22 ColinZ22 changed the title [AMD] Fix GlmMoeDsaForCausalLM with quark mxfp4+ptpcfp8-attn quantization [AMD] Fix Load and Inference of MLA models with Quark PTPC FP8 attention on ROCm Jun 19, 2026
Comment thread python/sglang/srt/layers/attention/dsa/dsa_indexer.py
Comment thread python/sglang/srt/layers/quantization/quark/schemes/quark_w8a8_fp8.py Outdated
@ColinZ22

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@ColinZ22

Copy link
Copy Markdown
Contributor Author

All PR check failures are unrelated to PR chagnes:

  • base-c-test-acc-2/4/8-npu-a3 and multimodal-gen-test-1-npu-a3 are runner/container failures with no test assertion reached, unrelated since this PR contains zero NPU code;
  • stage-b-test-1-gpu-small/large: server timed out/failed initialization due to torch compile/dynamo environment fault on a BF16 MoE path, not related to PR changes;
  • base-b-test-1-gpu-small (3): flaky exact-equality assert for logprobs of unquantized Llama-3.1-8B on NVIDIA with no FP8/AMD/quark involvement, unrelated to PR changes.
  • stage-c-test-large-8-gpu-amd-mi35x (0): Qwen3.5-FP8 accuracy drop, both the fused and unfused fusion launches failed, unrelated to PR code changes.

@HaiShaw Could you please help take a look and merge when you get a chance, thanks!

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

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:

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.

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Good catch! Removed the outdated function and fixed references to it in comments/assert messages.

@ColinZ22

ColinZ22 commented Sep 3, 2026

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@ColinZ22

Copy link
Copy Markdown
Contributor Author

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:

  • base-b-test-1/4-npu-a3: runner/container launch failure ("Executing the custom container implementation failed"), no test reached, unrelated, this PR has no NPU code.

  • base-c-test-8-gpu-b300: test_kimi_k3_b300.py failing, Kimi-K3 MegaMoE rank-5 scheduler aborted during init on B300; a server-init crash on Kimi K3 path, unrelated the DSA/quark/fp8 paths this PR touches.

  • base-c-test-8-gpu-h20 and stage-b-...-mi35x-disaggregation: test_disaggregation_pp.py failing, PD-disaggregation server failed to come up (ConnectionRefused); unrelated to this PR.

  • stage-b-test-1-gpu-small-amd-mi35x: test_kimi_k3_kda_decode.py ModuleNotFoundError: aiter.ops.flydsl.utils, a missing-module/env issue in the aiter install, unrelated.

  • stage-c-test-4-gpu-amd (mi300): test_lora_gpt_oss_20b_logprob_diff.py failing, GPT-OSS-20B LoRA logprob KL flake (retried values range from 1.6 to 19.4 to 4.0). Unrelated, since this is GPT-OSS LoRA, not an MLA/DSA/quark path.

  • stage-c-test-large-8-gpu-amd-mi35x (2): test_deepseek_r1_mxfp4_8gpu.py prefill CUDA graph capture failed, attempted to call function marked as skipped (torch/dynamo). DeepSeek-R1 is not a DSA model and uses mxfp4, unrelated to this PR.

  • stage-c-dsv4-flash-fp4-fp8-amd-mi35x: test_deepseek_v4_flash_fp8_tbo.py failing, GSM8K accuracy 0.31 vs 0.91 on the brand-new DeepSeek-V4-flash FP8 path. The PR's only DSV4-reachable change (the DSA-indexer tuple unwrap) is a no-op for V4's plain-tensor inputs, so it can't cause an accuracy drop; this is likely a newly-landed V4-flash path unstability on main.

@ColinZ22

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@BowenBao

Copy link
Copy Markdown
Collaborator

@ColinZ22 could you check if the `PR Test Base / base-c-test-8-gpu-b300 test failure is related or not?

@ColinZ22

ColinZ22 commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor Author

@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 (moonshotai/Kimi-K3) loads as DeepseekV3ForCausalLM + ModelOpt on B300 with the trtllm DSA backend, not using quark path in this PR and not aiter/gfx95, also not using the only potentially shared changes (in dsa_indexer.py, guarded to only 3-tuple inputs, produced only by the aiter kernel fused_rms_fp8_group_quant), so failure is unrelated.

Update: confirmed that this failure is caused by upstream changes and now fixed (PR #41138), pr-test-finish already passing after rebase and rerun.

@ColinZ22

Copy link
Copy Markdown
Contributor Author

Hi @HaiShaw @kkHuang-amd, could you please help take a look again when you get a chance? Thanks!

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

Labels

amd blackwell SM100/SM120 bypass-fastfail deepseek dependencies Pull requests that update a dependency file diffusion SGLang Diffusion documentation Improvements or additions to documentation hicache Hierarchical Caching for SGLang jit-kernel lora model-gateway mthreads Multi-modal multi-modal language model npu quant LLM Quantization run-ci CI: run the baseline test suite on this PR sgl-kernel speculative-decoding

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants