diff --git a/vllm/compilation/passes/pass_manager.py b/vllm/compilation/passes/pass_manager.py index 4cc6bc9e5f9..5f5e252c79b 100644 --- a/vllm/compilation/passes/pass_manager.py +++ b/vllm/compilation/passes/pass_manager.py @@ -143,6 +143,11 @@ class PostGradPassManager(CustomGraphPass): # type: ignore[misc] if self.pass_config.fuse_gemm_comms: self.passes += [AsyncTPPass(config)] + if self.pass_config.fuse_act_padding and rocm_aiter_ops.is_enabled(): + # Run the more specific RMSNorm+router-pad fusion before + # AR+RMS, since both consume fused_add_rms_norm. + self.passes += [RocmAiterTritonAddRMSNormPadFusionPass(config)] + if self.pass_config.fuse_allreduce_rms: if rocm_aiter_ops.is_enabled(): self.passes += [RocmAiterAllReduceFusionPass(config)] @@ -164,9 +169,6 @@ class PostGradPassManager(CustomGraphPass): # type: ignore[misc] if rocm_aiter_ops.is_enabled(): self.passes += [RocmAiterSiluMulFp8GroupQuantFusionPass(config)] - if self.pass_config.fuse_act_padding and rocm_aiter_ops.is_enabled(): - self.passes += [RocmAiterTritonAddRMSNormPadFusionPass(config)] - if self.pass_config.fuse_mla_dual_rms_norm and rocm_aiter_ops.is_enabled(): self.passes += [MLADualRMSNormFusionPass(config)]