forked from Karylab-cklius/vllm
[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:
+50
-14
@@ -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."
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user