Skip to content

[AMD] GLM-5.3-Flash mHC decode: fuse FFN all-reduce into hc_post - #42020

Draft
Jacob0226 wants to merge 1 commit into
sgl-project:mainfrom
Jacob0226:jacob/glm53-fused-ar-mhc-post
Draft

Jacob0226 wants to merge 1 commit into
sgl-project:mainfrom
Jacob0226:jacob/glm53-fused-ar-mhc-post

Conversation

@Jacob0226

@Jacob0226 Jacob0226 commented Oct 1, 2026 •

Copy link
Copy Markdown
Contributor

Summary

On a TP decode batch, each GLM-5.3-Flash layer writes its FFN output back in two launches. First the custom all-reduce sums the partial outputs. Then mhc_post_kernel reads the sum back and mixes it into the four residual streams. aiter's fused_allreduce_mhc_post_only does both in one kernel. With SGLANG_ROCM_FUSED_AR_MHC_POST=1, the FFN skips its own all-reduce, and the layer exit passes the partial sum to that kernel. If the custom all-reduce can't take the input, the layer runs all-reduce + hc_post as before.

One HIP-graph decode forward, batch 4, rank 0 (launches × median duration):

Kernel baseline this PR
cross_device_reduce_1stage 91 × 7.24 µs 46 × 7.64 µs
mhc_post_kernel 45 × 4.52 µs —
allreduce_mhc_post_large_m_kernel — 45 × 7.48 µs
Sum 862.1 µs 688.0 µs (−174.1)

Scope

ROCm with aiter, TP decode, off by default. It runs only when the FFN write-back is a plain TP sum over this rank's rows. That needs an MHC residual and at most 16 tokens, with no attention DP, MoE-CP gather, A2A backend, or scattered input. All other cases keep all-reduce + hc_post. Today only glm5_next builds an MHCState. DeepSeek-V4 has its own fused all-reduce + hc_post (#41021, #41308) in its MoE, through MhcPostFusion. It serves hidden size 5120 only and is untouched. GLM-5.3-Flash keeps its residual in an MHCState, so this PR hooks the layer exit instead. CUDA is untouched.

Test plan

Baseline rocm/sgl-dev:v0.5.20-rocm10-mi35x-20260928, with sglang replaced by main ebdeba2fee plus this commit. Both builds add a router GEMM config (ROCm/aiter#5933). Only SGLANG_ROCM_FUSED_AR_MHC_POST differs. amd/GLM-5.3-Flash-Quark-MXFP4, MI355X TP4, in8192 / out1024.

Concurrency Mean TTFT base (ms) this PR Median TTFT base (ms) this PR Median TPOT base (ms) this PR Δ
4 1373.42 303.08 233.34 234.56 9.23 9.08 −1.6%
16 544.11 530.42 243.04 238.88 13.44 13.29 −1.1%
64 1372.08 1355.96 353.64 351.02 26.68 26.50 −0.7%

At concurrency 64, batches exceed 16 tokens, so the fusion is off; that row is noise. Concurrency-4 mean TTFT swings from 300 to 1400 ms between runs, due to first-wave prefill stalls. Median TTFT is flat. The full forward goes 1664 → 1619 launches and 9.78 → 9.61 ms of kernel time, all from the Summary kernels.

Accuracy: GSM8K 0.967 → 0.969, single scoring pass each — within what one pass resolves.

…v-gated)

An mHC layer's FFN write-back is two launches on a TP decode batch: the
FFN's all-reduce, then hc_post, which reads the reduced output back to
mix it into the hc_mult residual streams. aiter's
fused_allreduce_mhc_post_only does both in one kernel on the custom
all-reduce's peer buffers.

With SGLANG_ROCM_FUSED_AR_MHC_POST=1, the FFN exit decides before the
FFN runs that the FFN skips its all-reduce, and its completion hands
the partial sum to that kernel. It applies only where the write-back is
a plain TP sum that stays on this rank's rows: an MHC FFN update, no
output transform or move, no attention DP, MoE-CP gather, A2A backend
or input-scattered layout, a single post-experts all-reduce over the TP
group, and at most 16 tokens. When the custom all-reduce declines the
input, the completion runs the all-reduce and hc_post as before.

Only glm5_next builds an MHCState today. DeepSeek-V4.1's fused
all-reduce + hc_post (sgl-project#41021) is wired into its own MoE through
MhcPostFusion and serves hidden size 5120 only; it is untouched.

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