Repository navigation
Conversation
…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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_kernelreads the sum back and mixes it into the four residual streams. aiter'sfused_allreduce_mhc_post_onlydoes both in one kernel. WithSGLANG_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):
cross_device_reduce_1stagemhc_post_kernelallreduce_mhc_post_large_m_kernelScope
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_nextbuilds anMHCState. DeepSeek-V4 has its own fused all-reduce + hc_post (#41021, #41308) in its MoE, throughMhcPostFusion. It serves hidden size 5120 only and is untouched. GLM-5.3-Flash keeps its residual in anMHCState, 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 mainebdeba2feeplus this commit. Both builds add a router GEMM config (ROCm/aiter#5933). OnlySGLANG_ROCM_FUSED_AR_MHC_POSTdiffers.amd/GLM-5.3-Flash-Quark-MXFP4, MI355X TP4, in8192 / out1024.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.