[ROCm][Bugfix] Fix fused_mla_dual_rms_norm for AITER API rename _fused_qk_rmsnorm (#42606)

Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
This commit is contained in:
Rita Brugarolas
2026-05-15 14:50:03 -06:00
committed by GitHub
parent de2d76f352
commit bd9dbe6060
2 changed files with 56 additions and 16 deletions
+50 -14
View File
@@ -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."
)
+6 -2
View File
@@ -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: