[MUSA] Use MUSA-optimized operators in piecewise CUDA graph (#23633)
Signed-off-by: popsiclexu <zhenxuexu@gmail.com>
This commit is contained in:
+1
-1
@@ -123,7 +123,7 @@ srt_musa = [
|
|||||||
"sglang[runtime_common]",
|
"sglang[runtime_common]",
|
||||||
"torch",
|
"torch",
|
||||||
"torch_musa",
|
"torch_musa",
|
||||||
"torchada>=0.1.54",
|
"torchada>=0.1.55",
|
||||||
"mthreads-ml-py",
|
"mthreads-ml-py",
|
||||||
"mate>=0.2.0",
|
"mate>=0.2.0",
|
||||||
"deep-gemm>=0.1.3",
|
"deep-gemm>=0.1.3",
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ srt_musa = [
|
|||||||
"sglang[runtime_common]",
|
"sglang[runtime_common]",
|
||||||
"torch",
|
"torch",
|
||||||
"torch_musa",
|
"torch_musa",
|
||||||
"torchada>=0.1.54",
|
"torchada>=0.1.55",
|
||||||
"mthreads-ml-py",
|
"mthreads-ml-py",
|
||||||
"mate>=0.2.0",
|
"mate>=0.2.0",
|
||||||
"deep-gemm>=0.1.3",
|
"deep-gemm>=0.1.3",
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
# Copyright 2023-2024 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
import re
|
||||||
|
from dataclasses import replace as _dataclass_replace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.fx.graph as fx_graph
|
||||||
|
|
||||||
|
_DEVICE_REPR_RE = re.compile(r"\bdevice\(type='([^']+)'(?:,\s*index=(\d+))?\)")
|
||||||
|
|
||||||
|
|
||||||
|
def _replace_device_repr(m: re.Match) -> str:
|
||||||
|
dev_type = m.group(1)
|
||||||
|
dev_index = m.group(2)
|
||||||
|
if dev_index is not None:
|
||||||
|
return f"torch.device('{dev_type}:{dev_index}')"
|
||||||
|
return f"torch.device('{dev_type}')"
|
||||||
|
|
||||||
|
|
||||||
|
def patch_fx_custom_device() -> None:
|
||||||
|
"""
|
||||||
|
Fix FX codegen serialization for non-standard devices (e.g. torch_musa).
|
||||||
|
|
||||||
|
Root cause:
|
||||||
|
torch.device is registered as a custom builtin named 'device', imported
|
||||||
|
via 'from torch import device'. repr(torch.device('musa', 0)) produces
|
||||||
|
"device(type='musa', index=0)", which is syntactically valid but fails
|
||||||
|
at runtime because torch.device does not recognize 'musa' as a type when
|
||||||
|
invoked through the standard import path.
|
||||||
|
|
||||||
|
Fix:
|
||||||
|
Post-process the generated src string, replacing all occurrences of
|
||||||
|
device(type='x', index=N) with torch.device('x:N'), and ensure 'torch'
|
||||||
|
is present in the graph globals.
|
||||||
|
|
||||||
|
Note:
|
||||||
|
_get_repr is a closure inside _gen_python_code and cannot be patched
|
||||||
|
directly, so we wrap _gen_python_code and rewrite its output instead.
|
||||||
|
"""
|
||||||
|
original = fx_graph.CodeGen._gen_python_code
|
||||||
|
|
||||||
|
def patched(self, nodes, root_module, namespace, **kwargs):
|
||||||
|
result = original(self, nodes, root_module, namespace, **kwargs)
|
||||||
|
new_src = _DEVICE_REPR_RE.sub(_replace_device_repr, result.src)
|
||||||
|
if new_src is result.src:
|
||||||
|
return result
|
||||||
|
result.globals.setdefault("torch", torch)
|
||||||
|
if hasattr(result, "_replace"):
|
||||||
|
return result._replace(src=new_src)
|
||||||
|
return _dataclass_replace(result, src=new_src)
|
||||||
|
|
||||||
|
fx_graph.CodeGen._gen_python_code = patched
|
||||||
@@ -62,7 +62,14 @@ elif _is_xpu:
|
|||||||
elif _is_hip:
|
elif _is_hip:
|
||||||
from sgl_kernel import gelu_and_mul, gelu_quick, gelu_tanh_and_mul, silu_and_mul
|
from sgl_kernel import gelu_and_mul, gelu_quick, gelu_tanh_and_mul, silu_and_mul
|
||||||
elif _is_musa:
|
elif _is_musa:
|
||||||
from sgl_kernel import silu_and_mul
|
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||||
|
|
||||||
|
@register_fake_if_exists("aten::_fused_swiglu_forward")
|
||||||
|
def _(x):
|
||||||
|
d = x.shape[-1] // 2
|
||||||
|
output_shape = x.shape[:-1] + (d,)
|
||||||
|
return torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||||
|
|
||||||
|
|
||||||
if is_npu():
|
if is_npu():
|
||||||
import torch_npu
|
import torch_npu
|
||||||
@@ -106,9 +113,6 @@ class SiluAndMul(MultiPlatformOp):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
def forward_musa(self, x: torch.Tensor) -> torch.Tensor:
|
def forward_musa(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
if not get_global_server_args().disable_piecewise_cuda_graph:
|
|
||||||
return self.forward_native(x)
|
|
||||||
|
|
||||||
if not hasattr(self, "_musa_swish_glu"):
|
if not hasattr(self, "_musa_swish_glu"):
|
||||||
# XXX (MUSA): nn.SwishGLU seems to have better performance than silu_and_mul on MUSA, we can switch to it for now. We can consider implementing a silu_and_mul kernel for MUSA in the future if needed.
|
# XXX (MUSA): nn.SwishGLU seems to have better performance than silu_and_mul on MUSA, we can switch to it for now. We can consider implementing a silu_and_mul kernel for MUSA in the future if needed.
|
||||||
self._musa_swish_glu = nn.SwishGLU()
|
self._musa_swish_glu = nn.SwishGLU()
|
||||||
|
|||||||
@@ -344,9 +344,6 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
residual: Optional[torch.Tensor] = None,
|
residual: Optional[torch.Tensor] = None,
|
||||||
post_residual_addition: Optional[torch.Tensor] = None,
|
post_residual_addition: Optional[torch.Tensor] = None,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
if not get_global_server_args().disable_piecewise_cuda_graph:
|
|
||||||
return self.forward_native(x, residual, post_residual_addition)
|
|
||||||
|
|
||||||
if not x.is_contiguous():
|
if not x.is_contiguous():
|
||||||
x = x.contiguous()
|
x = x.contiguous()
|
||||||
|
|
||||||
|
|||||||
@@ -94,6 +94,24 @@ if _is_hip:
|
|||||||
# Fallback: vllm not available, will use native PyTorch implementation
|
# Fallback: vllm not available, will use native PyTorch implementation
|
||||||
_has_vllm = False
|
_has_vllm = False
|
||||||
|
|
||||||
|
if _is_musa:
|
||||||
|
|
||||||
|
@register_fake_if_exists("sgl_kernel::sgl_per_token_group_quant_8bit_v2")
|
||||||
|
def _(
|
||||||
|
input,
|
||||||
|
output_q,
|
||||||
|
output_s,
|
||||||
|
group_size,
|
||||||
|
eps,
|
||||||
|
fp8_min,
|
||||||
|
fp8_max,
|
||||||
|
scale_ue8m0,
|
||||||
|
fuse_silu_and_mul,
|
||||||
|
masked_m,
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
|
is_musa,
|
||||||
is_npu,
|
is_npu,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
require_gathered_buffer,
|
require_gathered_buffer,
|
||||||
@@ -73,6 +74,8 @@ logger = logging.getLogger(__name__)
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
_is_musa = is_musa()
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PrefillInputBuffers(ForwardInputBuffers):
|
class PrefillInputBuffers(ForwardInputBuffers):
|
||||||
@@ -147,6 +150,13 @@ def set_torch_compile_config():
|
|||||||
if hasattr(torch._dynamo.config, "cache_size_limit"):
|
if hasattr(torch._dynamo.config, "cache_size_limit"):
|
||||||
torch._dynamo.config.cache_size_limit = 1024
|
torch._dynamo.config.cache_size_limit = 1024
|
||||||
|
|
||||||
|
if _is_musa:
|
||||||
|
from sglang.srt.hardware_backend.musa.utils.patch_torch import (
|
||||||
|
patch_fx_custom_device,
|
||||||
|
)
|
||||||
|
|
||||||
|
patch_fx_custom_device()
|
||||||
|
|
||||||
|
|
||||||
class PiecewiseCudaGraphRunner:
|
class PiecewiseCudaGraphRunner:
|
||||||
"""A PiecewiseCudaGraphRunner runs the forward pass of a model with cuda graph and torch.compile."""
|
"""A PiecewiseCudaGraphRunner runs the forward pass of a model with cuda graph and torch.compile."""
|
||||||
|
|||||||
@@ -16,9 +16,10 @@ from typing import Callable, Union
|
|||||||
import torch
|
import torch
|
||||||
from torch.multiprocessing import reductions
|
from torch.multiprocessing import reductions
|
||||||
|
|
||||||
from sglang.srt.utils.common import is_npu, torch_release
|
from sglang.srt.utils.common import is_musa, is_npu, torch_release
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
_is_musa = is_musa()
|
||||||
|
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
from torch_npu.multiprocessing import reductions as npu_reductions
|
from torch_npu.multiprocessing import reductions as npu_reductions
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ requires = [
|
|||||||
"setuptools>=75.0",
|
"setuptools>=75.0",
|
||||||
"scikit-build-core>=0.10",
|
"scikit-build-core>=0.10",
|
||||||
"torch",
|
"torch",
|
||||||
"torchada>=0.1.54",
|
"torchada>=0.1.55",
|
||||||
"wheel",
|
"wheel",
|
||||||
]
|
]
|
||||||
build-backend = "setuptools.build_meta"
|
build-backend = "setuptools.build_meta"
|
||||||
|
|||||||
@@ -56,7 +56,7 @@ def cache_once(fn):
|
|||||||
|
|
||||||
@cache_once
|
@cache_once
|
||||||
def is_arch_support_pdl() -> bool:
|
def is_arch_support_pdl() -> bool:
|
||||||
if bool(torch.version.hip):
|
if getattr(torch.version, "hip", None) or getattr(torch.version, "musa", None):
|
||||||
return False
|
return False
|
||||||
try:
|
try:
|
||||||
device = torch.cuda.current_device()
|
device = torch.cuda.current_device()
|
||||||
|
|||||||
Reference in New Issue
Block a user