[AMD] Clean up vllm dependencies in moe_runner/triton.py (#11349)
Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -32,23 +32,25 @@ _is_hip = is_hip()
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_cpu_amx_available = cpu_has_amx_support()
|
_is_cpu_amx_available = cpu_has_amx_support()
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
_use_aiter = bool(int(os.getenv("SGLANG_MOE_USE_AITER", "0")))
|
_use_aiter = bool(int(os.getenv("SGLANG_USE_AITER", "0")))
|
||||||
_MOE_PADDING_SIZE = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0
|
_MOE_PADDING_SIZE = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0
|
||||||
|
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda or _is_hip:
|
||||||
from sgl_kernel import gelu_and_mul, silu_and_mul
|
from sgl_kernel import gelu_and_mul, silu_and_mul
|
||||||
elif _is_cpu and _is_cpu_amx_available:
|
|
||||||
pass
|
|
||||||
elif _is_hip:
|
|
||||||
from vllm import _custom_ops as vllm_ops # gelu_and_mul, silu_and_mul
|
|
||||||
|
|
||||||
|
if _is_hip:
|
||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
try:
|
try:
|
||||||
from aiter import moe_sum
|
from aiter import moe_sum
|
||||||
except ImportError:
|
except ImportError:
|
||||||
raise ImportError("aiter is required when SGLANG_USE_AITER is set to True")
|
raise ImportError(
|
||||||
|
"aiter is required when SGLANG_USE_AITER is set to True"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
from vllm import _custom_ops as vllm_ops # moe_sum
|
||||||
|
elif _is_cpu and _is_cpu_amx_available:
|
||||||
|
pass
|
||||||
|
|
||||||
if _is_cuda or _is_hip:
|
if _is_cuda or _is_hip:
|
||||||
from sgl_kernel import ( # noqa: F401
|
from sgl_kernel import ( # noqa: F401
|
||||||
@@ -206,7 +208,7 @@ class TritonRunnerCore(MoeRunnerCore):
|
|||||||
gemm1_alpha,
|
gemm1_alpha,
|
||||||
gemm1_limit,
|
gemm1_limit,
|
||||||
)
|
)
|
||||||
elif _is_cuda:
|
elif _is_cuda or _is_hip:
|
||||||
silu_and_mul(intermediate_cache1.view(-1, N), intermediate_cache2)
|
silu_and_mul(intermediate_cache1.view(-1, N), intermediate_cache2)
|
||||||
else:
|
else:
|
||||||
vllm_ops.silu_and_mul(
|
vllm_ops.silu_and_mul(
|
||||||
@@ -215,7 +217,7 @@ class TritonRunnerCore(MoeRunnerCore):
|
|||||||
elif activation == "gelu":
|
elif activation == "gelu":
|
||||||
assert gemm1_alpha is None, "gemm1_alpha is not supported for gelu"
|
assert gemm1_alpha is None, "gemm1_alpha is not supported for gelu"
|
||||||
assert gemm1_limit is None, "gemm1_limit is not supported for gelu"
|
assert gemm1_limit is None, "gemm1_limit is not supported for gelu"
|
||||||
if _is_cuda:
|
if _is_cuda or _is_hip:
|
||||||
gelu_and_mul(intermediate_cache1.view(-1, N), intermediate_cache2)
|
gelu_and_mul(intermediate_cache1.view(-1, N), intermediate_cache2)
|
||||||
else:
|
else:
|
||||||
vllm_ops.gelu_and_mul(
|
vllm_ops.gelu_and_mul(
|
||||||
|
|||||||
Reference in New Issue
Block a user