diff --git a/vllm/_aiter_ops.py b/vllm/_aiter_ops.py index a144d05eb44..ce4fc3cfbad 100644 --- a/vllm/_aiter_ops.py +++ b/vllm/_aiter_ops.py @@ -414,17 +414,32 @@ _AITER_HAS_FUSED_QK_RMSNORM: bool | None = None def check_aiter_fused_qk_rmsnorm() -> bool: - """Check if aiter provides fused_qk_rmsnorm (requires AITer >= PR #2442).""" + """Check if aiter provides fused_qk_rmsnorm. + + Supports both the new private name ``_fused_qk_rmsnorm`` + (AITER >= PR #2958) and the old public name ``fused_qk_rmsnorm`` + (AITER >= PR #2442). + + TODO(rbrugaro-amd): remove the legacy fused_qk_rmsnorm path once + AITER stabilizes the API (https://github.com/ROCm/aiter/issues/3207). + """ global _AITER_HAS_FUSED_QK_RMSNORM if _AITER_HAS_FUSED_QK_RMSNORM is None: try: from aiter.ops.fused_qk_norm_rope_cache_quant import ( # noqa: F401 - fused_qk_rmsnorm, + _fused_qk_rmsnorm, ) _AITER_HAS_FUSED_QK_RMSNORM = True except (ImportError, ModuleNotFoundError, AttributeError): - _AITER_HAS_FUSED_QK_RMSNORM = False + try: + from aiter.ops.fused_qk_norm_rope_cache_quant import ( # noqa: F401 + fused_qk_rmsnorm, + ) + + _AITER_HAS_FUSED_QK_RMSNORM = True + except (ImportError, ModuleNotFoundError, AttributeError): + _AITER_HAS_FUSED_QK_RMSNORM = False return _AITER_HAS_FUSED_QK_RMSNORM @@ -1066,21 +1081,42 @@ def _fused_mla_dual_rms_norm_impl( x2_epsilon: float, ) -> tuple[torch.Tensor, torch.Tensor]: try: - from aiter.ops.fused_qk_norm_rope_cache_quant import fused_qk_rmsnorm - except (ImportError, ModuleNotFoundError) as exc: + import aiter.ops.fused_qk_norm_rope_cache_quant as aiter_ops + except (ImportError, ModuleNotFoundError, AttributeError) as exc: raise ImportError( - "fused_qk_rmsnorm requires a newer AITer version " - "(>= PR #2442). Please upgrade aiter or disable the " + "fused_qk_rmsnorm requires AITer >= PR #2442. " + "Please upgrade aiter or disable the " "fuse_mla_dual_rms_norm pass." ) from exc - return fused_qk_rmsnorm( - q=x1, - q_weight=x1_weight, - q_eps=x1_epsilon, - k=x2, - k_weight=x2_weight, - k_eps=x2_epsilon, + if hasattr(aiter_ops, "_fused_qk_rmsnorm"): + return aiter_ops._fused_qk_rmsnorm( + q_out=None, + q=x1, + q_weight=x1_weight, + q_eps=x1_epsilon, + k_out=None, + k=x2, + k_weight=x2_weight, + k_eps=x2_epsilon, + ) + + # TODO(rbrugaro-amd): remove the legacy fused_qk_rmsnorm path once + # AITER stabilizes the API (https://github.com/ROCm/aiter/issues/3207). + if hasattr(aiter_ops, "fused_qk_rmsnorm"): + return aiter_ops.fused_qk_rmsnorm( + q=x1, + q_weight=x1_weight, + q_eps=x1_epsilon, + k=x2, + k_weight=x2_weight, + k_eps=x2_epsilon, + ) + + raise ImportError( + "fused_qk_rmsnorm requires AITer >= PR #2442. " + "Please upgrade aiter or disable the " + "fuse_mla_dual_rms_norm pass." ) diff --git a/vllm/compilation/passes/pass_manager.py b/vllm/compilation/passes/pass_manager.py index 5f5e252c79b..9c86518a946 100644 --- a/vllm/compilation/passes/pass_manager.py +++ b/vllm/compilation/passes/pass_manager.py @@ -7,7 +7,7 @@ from typing import Any, ParamSpec, TypeVar from torch import fx as fx from vllm import envs -from vllm._aiter_ops import rocm_aiter_ops +from vllm._aiter_ops import check_aiter_fused_qk_rmsnorm, rocm_aiter_ops from vllm.compilation.passes.utility.post_cleanup import PostCleanupPass from vllm.config import VllmConfig, set_current_vllm_config from vllm.logger import init_logger @@ -169,7 +169,11 @@ class PostGradPassManager(CustomGraphPass): # type: ignore[misc] if rocm_aiter_ops.is_enabled(): self.passes += [RocmAiterSiluMulFp8GroupQuantFusionPass(config)] - if self.pass_config.fuse_mla_dual_rms_norm and rocm_aiter_ops.is_enabled(): + if ( + self.pass_config.fuse_mla_dual_rms_norm + and rocm_aiter_ops.is_enabled() + and check_aiter_fused_qk_rmsnorm() + ): self.passes += [MLADualRMSNormFusionPass(config)] if self.pass_config.fuse_rope_kvcache: