WIP: initial multimodal-gen support (#12484)
Co-authored-by: yhyang201 <yhyang201@gmail.com> Co-authored-by: yizhang2077 <1109276519@qq.com> Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: ispobock <ispobaoke@gmail.com> Co-authored-by: JiLi <leege233@gmail.com> Co-authored-by: CHEN Xi <78632976+RubiaCx@users.noreply.github.com> Co-authored-by: laixin <xielx@shanghaitech.edu.cn> Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com> Co-authored-by: jzhang38 <a1286225768@gmail.com> Co-authored-by: BrianChen1129 <yongqichcd@gmail.com> Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com> Co-authored-by: Edenzzzz <wtan45@wisc.edu> Co-authored-by: rlsu9 <r3su@ucsd.edu> Co-authored-by: Jinzhe Pan <48981407+eigensystem@users.noreply.github.com> Co-authored-by: foreverpiano <pianoqwz@qq.com> Co-authored-by: RandNMR73 <notomatthew31@gmail.com> Co-authored-by: PorridgeSwim <yz3883@columbia.edu> Co-authored-by: Jiali Chen <90408393+gary-chenjl@users.noreply.github.com>
This commit is contained in:
co-authored by
yhyang201
yizhang2077
Xinyuan Tong
ispobock
JiLi
CHEN Xi
laixin
SolitaryThinker
jzhang38
BrianChen1129
Kevin Lin
Edenzzzz
rlsu9
Jinzhe Pan
foreverpiano
RandNMR73
PorridgeSwim
Jiali Chen
parent
4fe53e5888
commit
7bc1dae095
@@ -0,0 +1,172 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/__init__.py
|
||||
|
||||
import traceback
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
# imported by other files, do not remove
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import ( # noqa: F401
|
||||
AttentionBackendEnum,
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def cuda_platform_plugin() -> str | None:
|
||||
is_cuda = False
|
||||
|
||||
try:
|
||||
from sglang.multimodal_gen.utils import import_pynvml
|
||||
|
||||
pynvml = import_pynvml() # type: ignore[no-untyped-call]
|
||||
pynvml.nvmlInit()
|
||||
try:
|
||||
# NOTE: Edge case: sgl_diffusion cpu build on a GPU machine.
|
||||
# Third-party pynvml can be imported in cpu build,
|
||||
# we need to check if sgl_diffusion is built with cpu too.
|
||||
# Otherwise, sgl_diffusion will always activate cuda plugin
|
||||
# on a GPU machine, even if in a cpu build.
|
||||
is_cuda = pynvml.nvmlDeviceGetCount() > 0
|
||||
finally:
|
||||
pynvml.nvmlShutdown()
|
||||
except Exception as e:
|
||||
if "nvml" not in e.__class__.__name__.lower():
|
||||
# If the error is not related to NVML, re-raise it.
|
||||
raise e
|
||||
|
||||
# CUDA is supported on Jetson, but NVML may not be.
|
||||
import os
|
||||
|
||||
def cuda_is_jetson() -> bool:
|
||||
return os.path.isfile("/etc/nv_tegra_release") or os.path.exists(
|
||||
"/sys/class/tegra-firmware"
|
||||
)
|
||||
|
||||
if cuda_is_jetson():
|
||||
is_cuda = True
|
||||
if is_cuda:
|
||||
logger.info("CUDA is available")
|
||||
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.cuda.CudaPlatform" if is_cuda else None
|
||||
)
|
||||
|
||||
|
||||
def mps_platform_plugin() -> str | None:
|
||||
"""Detect if MPS (Metal Performance Shaders) is available on macOS."""
|
||||
is_mps = False
|
||||
|
||||
try:
|
||||
import torch
|
||||
|
||||
if torch.backends.mps.is_available():
|
||||
is_mps = True
|
||||
logger.info("MPS (Metal Performance Shaders) is available")
|
||||
except Exception as e:
|
||||
logger.info("MPS detection failed: %s", e)
|
||||
|
||||
return "sglang.multimodal_gen.runtime.platforms.mps.MpsPlatform" if is_mps else None
|
||||
|
||||
|
||||
def cpu_platform_plugin() -> str | None:
|
||||
"""Detect if CPU platform should be used."""
|
||||
# CPU is always available as a fallback
|
||||
return "sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform"
|
||||
|
||||
|
||||
def rocm_platform_plugin() -> str | None:
|
||||
is_rocm = False
|
||||
|
||||
try:
|
||||
import amdsmi
|
||||
|
||||
amdsmi.amdsmi_init()
|
||||
try:
|
||||
if len(amdsmi.amdsmi_get_processor_handles()) > 0:
|
||||
is_rocm = True
|
||||
logger.info("ROCm platform is available")
|
||||
finally:
|
||||
amdsmi.amdsmi_shut_down()
|
||||
except Exception as e:
|
||||
logger.info("ROCm platform is unavailable: %s", e)
|
||||
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.rocm.RocmPlatform" if is_rocm else None
|
||||
)
|
||||
|
||||
|
||||
builtin_platform_plugins = {
|
||||
"cuda": cuda_platform_plugin,
|
||||
"rocm": rocm_platform_plugin,
|
||||
"mps": mps_platform_plugin,
|
||||
"cpu": cpu_platform_plugin,
|
||||
}
|
||||
|
||||
|
||||
def resolve_current_platform_cls_qualname() -> str:
|
||||
# TODO(will): if we need to support other platforms, we should consider if
|
||||
# vLLM's plugin architecture is suitable for our needs.
|
||||
|
||||
# Try MPS first on macOS
|
||||
platform_cls_qualname = mps_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to ROCm
|
||||
platform_cls_qualname = rocm_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to CUDA
|
||||
platform_cls_qualname = cuda_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to CPU as last resort
|
||||
platform_cls_qualname = cpu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
raise RuntimeError("No platform plugin found. Please check your " "installation.")
|
||||
|
||||
|
||||
_current_platform: Platform | None = None
|
||||
_init_trace: str = ""
|
||||
|
||||
if TYPE_CHECKING:
|
||||
current_platform: Platform
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
if name == "current_platform":
|
||||
# lazy init current_platform.
|
||||
# 1. out-of-tree platform plugins need `from sglang.multimodal_gen.runtime.platforms import
|
||||
# Platform` so that they can inherit `Platform` class. Therefore,
|
||||
# we cannot resolve `current_platform` during the import of
|
||||
# `sglang.multimodal_gen.runtime.platforms`.
|
||||
# 2. when users use out-of-tree platform plugins, they might run
|
||||
# `import sgl_diffusion`, some sgl_diffusion internal code might access
|
||||
# `current_platform` during the import, and we need to make sure
|
||||
# `current_platform` is only resolved after the plugins are loaded
|
||||
# (we have tests for this, if any developer violate this, they will
|
||||
# see the test failures).
|
||||
global _current_platform
|
||||
if _current_platform is None:
|
||||
platform_cls_qualname = resolve_current_platform_cls_qualname()
|
||||
_current_platform = resolve_obj_by_qualname(platform_cls_qualname)()
|
||||
global _init_trace
|
||||
_init_trace = "".join(traceback.format_stack())
|
||||
return _current_platform
|
||||
elif name in globals():
|
||||
return globals()[name]
|
||||
else:
|
||||
raise AttributeError(f"No attribute named '{name}' exists in {__name__}.")
|
||||
|
||||
|
||||
__all__ = ["Platform", "PlatformEnum", "current_platform", "_init_trace"]
|
||||
@@ -0,0 +1,61 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/cpu.py
|
||||
|
||||
import platform
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
CpuArchEnum,
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
|
||||
|
||||
class CpuPlatform(Platform):
|
||||
_enum = PlatformEnum.CPU
|
||||
device_name = "CPU"
|
||||
device_type = "cpu"
|
||||
dispatch_key = "CPU"
|
||||
|
||||
@classmethod
|
||||
def get_cpu_architecture(cls) -> CpuArchEnum:
|
||||
"""Get the CPU architecture."""
|
||||
machine = platform.machine().lower()
|
||||
if machine in ("x86_64", "amd64", "i386", "i686"):
|
||||
return CpuArchEnum.X86
|
||||
elif machine in ("arm64", "aarch64"):
|
||||
return CpuArchEnum.ARM
|
||||
else:
|
||||
return CpuArchEnum.UNSPECIFIED
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return platform.processor()
|
||||
|
||||
@classmethod
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
return platform.machine()
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
# This is a rough estimate for CPU memory
|
||||
# In practice, you might want to use psutil or similar
|
||||
return 0
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(
|
||||
cls, device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
# For CPU, we can't easily get memory usage without additional libraries
|
||||
return 0.0
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator.CpuCommunicator"
|
||||
@@ -0,0 +1,430 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/cuda.py
|
||||
"""Code inside this file can safely assume cuda platform, e.g. importing
|
||||
pynvml. However, it should not initialize cuda context.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from functools import lru_cache, wraps
|
||||
from typing import TypeVar
|
||||
|
||||
import torch
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
AttentionBackendEnum,
|
||||
DeviceCapability,
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.common import is_blackwell
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import import_pynvml
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
|
||||
pynvml = import_pynvml() # type: ignore[no-untyped-call]
|
||||
|
||||
# pytorch 2.5 uses cudnn sdpa by default, which will cause crash on some models
|
||||
# see https://github.com/huggingface/diffusers/issues/9704 for details
|
||||
torch.backends.cuda.enable_cudnn_sdp(False)
|
||||
|
||||
|
||||
def device_id_to_physical_device_id(device_id: int) -> int:
|
||||
if "CUDA_VISIBLE_DEVICES" in os.environ:
|
||||
device_ids = os.environ["CUDA_VISIBLE_DEVICES"].split(",")
|
||||
if device_ids == [""]:
|
||||
msg = (
|
||||
"CUDA_VISIBLE_DEVICES is set to empty string, which means"
|
||||
" GPU support is disabled. If you are using ray, please unset"
|
||||
" the environment variable `CUDA_VISIBLE_DEVICES` inside the"
|
||||
" worker/actor. "
|
||||
"Check https://github.com/vllm-project/vllm/issues/8402 for"
|
||||
" more information."
|
||||
)
|
||||
raise RuntimeError(msg)
|
||||
physical_device_id = device_ids[device_id]
|
||||
return int(physical_device_id)
|
||||
else:
|
||||
return device_id
|
||||
|
||||
|
||||
def with_nvml_context(fn: Callable[_P, _R]) -> Callable[_P, _R]:
|
||||
@wraps(fn)
|
||||
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
|
||||
pynvml.nvmlInit()
|
||||
try:
|
||||
return fn(*args, **kwargs)
|
||||
finally:
|
||||
pynvml.nvmlShutdown()
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
class CudaPlatformBase(Platform):
|
||||
_enum = PlatformEnum.CUDA
|
||||
device_name: str = "cuda"
|
||||
device_type: str = "cuda"
|
||||
dispatch_key: str = "CUDA"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
if enforce_eager:
|
||||
logger.warning(
|
||||
"To see benefits of async output processing, enable CUDA "
|
||||
"graph. Since, enforce-eager is enabled, async output "
|
||||
"processor cannot be used"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def is_full_nvlink(cls, device_ids: list[int]) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def log_warnings(cls) -> None:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(
|
||||
cls, device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls_str(
|
||||
cls,
|
||||
selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
) -> str:
|
||||
# TODO(will): maybe come up with a more general interface for local attention
|
||||
# if distributed is False, we always try to use Flash attn
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.sliding_tile_attn import ( # noqa: F401
|
||||
SlidingTileAttentionBackend,
|
||||
)
|
||||
|
||||
logger.info("Using Sliding Tile Attention backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.error(
|
||||
"Failed to import Sliding Tile Attention backend: %s", str(e)
|
||||
)
|
||||
raise ImportError(
|
||||
"Sliding Tile Attention backend is not installed. "
|
||||
) from e
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN:
|
||||
try:
|
||||
from sageattention import sageattn # noqa: F401
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn import ( # noqa: F401
|
||||
SageAttentionBackend,
|
||||
)
|
||||
|
||||
logger.info("Using Sage Attention backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn.SageAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
|
||||
try:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.sageattn.api import ( # noqa: F401
|
||||
sageattn_blackwell,
|
||||
)
|
||||
|
||||
logger.info("Using Sage Attention 3 backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn3.SageAttention3Backend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention 3 backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn import ( # noqa: F401
|
||||
VideoSparseAttentionBackend,
|
||||
)
|
||||
|
||||
logger.info("Using Video Sparse Attention backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.error(
|
||||
"Failed to import Video Sparse Attention backend: %s", str(e)
|
||||
)
|
||||
raise ImportError(
|
||||
"Video Sparse Attention backend is not installed. "
|
||||
) from e
|
||||
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
||||
try:
|
||||
from kernel.attn.vmoba_attn.vmoba import moba_attn_varlen # noqa: F401
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.vmoba import ( # noqa: F401
|
||||
VMOBAAttentionBackend,
|
||||
)
|
||||
|
||||
logger.info("Using Video MOBA Attention backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.vmoba.VMOBAAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.error(
|
||||
"Failed to import Video MoBA Attention backend: %s", str(e)
|
||||
)
|
||||
raise ImportError(
|
||||
"Video MoBA Attention backend is not installed. "
|
||||
) from e
|
||||
elif selected_backend == AttentionBackendEnum.AITER:
|
||||
logger.info("Using AITer backend.")
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.aiter.AITerBackend"
|
||||
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||
elif selected_backend == AttentionBackendEnum.FA3:
|
||||
if is_blackwell():
|
||||
raise ValueError("The 'fa3' backend is not supported on Blackwell GPUs")
|
||||
elif selected_backend:
|
||||
raise ValueError(f"Invalid attention backend for {cls.device_name}")
|
||||
else:
|
||||
if is_blackwell():
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
logger.debug(f"Use torch_sdpa as default backend")
|
||||
else:
|
||||
target_backend = AttentionBackendEnum.FA3
|
||||
logger.debug(f"Use fa3 as default backend")
|
||||
|
||||
if not cls.has_device_capability(80):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for Volta and Turing " "GPUs."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
elif dtype not in (torch.float16, torch.bfloat16):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for dtype other than "
|
||||
"torch.float16 or torch.bfloat16."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
# FlashAttn is valid for the model, checking if the package is
|
||||
# installed.
|
||||
if target_backend == AttentionBackendEnum.FA3:
|
||||
try:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import ( # noqa: F401
|
||||
FlashAttentionBackend,
|
||||
)
|
||||
|
||||
supported_sizes = FlashAttentionBackend.get_supported_head_sizes()
|
||||
if head_size not in supported_sizes:
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for head size %d.",
|
||||
head_size,
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
except ImportError:
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend because the "
|
||||
"flash_attn package is not found. "
|
||||
"Make sure that flash_attn was built and installed "
|
||||
"(on by default)."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
logger.info("Using fa3 backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
|
||||
|
||||
|
||||
# NVML utils
|
||||
# Note that NVML is not affected by `CUDA_VISIBLE_DEVICES`,
|
||||
# all the related functions work on real physical device ids.
|
||||
# the major benefit of using NVML is that it will not initialize CUDA
|
||||
class NvmlCudaPlatform(CudaPlatformBase):
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||
try:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
|
||||
major, minor = pynvml.nvmlDeviceGetCudaComputeCapability(handle)
|
||||
return DeviceCapability(major=major, minor=minor)
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def has_device_capability(
|
||||
cls,
|
||||
capability: tuple[int, int] | int,
|
||||
device_id: int = 0,
|
||||
) -> bool:
|
||||
try:
|
||||
return bool(super().has_device_capability(capability, device_id))
|
||||
except RuntimeError:
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
return cls._get_physical_device_name(physical_device_id)
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
|
||||
return str(pynvml.nvmlDeviceGetUUID(handle))
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=8)
|
||||
@with_nvml_context
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
physical_device_id = device_id_to_physical_device_id(device_id)
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(physical_device_id)
|
||||
return int(pynvml.nvmlDeviceGetMemoryInfo(handle).total)
|
||||
|
||||
@classmethod
|
||||
@with_nvml_context
|
||||
def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool:
|
||||
"""
|
||||
query if the set of gpus are fully connected by nvlink (1 hop)
|
||||
"""
|
||||
handles = [pynvml.nvmlDeviceGetHandleByIndex(i) for i in physical_device_ids]
|
||||
for i, handle in enumerate(handles):
|
||||
for j, peer_handle in enumerate(handles):
|
||||
if i < j:
|
||||
try:
|
||||
p2p_status = pynvml.nvmlDeviceGetP2PStatus(
|
||||
handle,
|
||||
peer_handle,
|
||||
pynvml.NVML_P2P_CAPS_INDEX_NVLINK,
|
||||
)
|
||||
if p2p_status != pynvml.NVML_P2P_STATUS_OK:
|
||||
return False
|
||||
except pynvml.NVMLError:
|
||||
logger.exception(
|
||||
"NVLink detection failed. This is normal if"
|
||||
" your machine has no NVLink equipped."
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def _get_physical_device_name(cls, device_id: int = 0) -> str:
|
||||
handle = pynvml.nvmlDeviceGetHandleByIndex(device_id)
|
||||
return str(pynvml.nvmlDeviceGetName(handle))
|
||||
|
||||
@classmethod
|
||||
@with_nvml_context
|
||||
def log_warnings(cls) -> None:
|
||||
device_ids: int = pynvml.nvmlDeviceGetCount()
|
||||
if device_ids > 1:
|
||||
device_names = [cls._get_physical_device_name(i) for i in range(device_ids)]
|
||||
if (
|
||||
len(set(device_names)) > 1
|
||||
and os.environ.get("CUDA_DEVICE_ORDER") != "PCI_BUS_ID"
|
||||
):
|
||||
logger.warning(
|
||||
"Detected different devices in the system: %s. Please"
|
||||
" make sure to set `CUDA_DEVICE_ORDER=PCI_BUS_ID` to "
|
||||
"avoid unexpected behavior.",
|
||||
", ".join(device_names),
|
||||
)
|
||||
|
||||
|
||||
class NonNvmlCudaPlatform(CudaPlatformBase):
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||
major, minor = torch.cuda.get_device_capability(device_id)
|
||||
return DeviceCapability(major=major, minor=minor)
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return str(torch.cuda.get_device_name(device_id))
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
device_props = torch.cuda.get_device_properties(device_id)
|
||||
return int(device_props.total_memory)
|
||||
|
||||
@classmethod
|
||||
def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool:
|
||||
logger.exception(
|
||||
"NVLink detection not possible, as context support was"
|
||||
" not found. Assuming no NVLink available."
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
# Autodetect either NVML-enabled or non-NVML platform
|
||||
# based on whether NVML is available.
|
||||
nvml_available = False
|
||||
try:
|
||||
try:
|
||||
pynvml.nvmlInit()
|
||||
nvml_available = True
|
||||
except Exception:
|
||||
# On Jetson, NVML is not supported.
|
||||
nvml_available = False
|
||||
finally:
|
||||
if nvml_available:
|
||||
pynvml.nvmlShutdown()
|
||||
|
||||
CudaPlatform = NvmlCudaPlatform if nvml_available else NonNvmlCudaPlatform
|
||||
|
||||
try:
|
||||
from sphinx.ext.autodoc.mock import _MockModule
|
||||
|
||||
if not isinstance(pynvml, _MockModule):
|
||||
CudaPlatform.log_warnings()
|
||||
except ModuleNotFoundError:
|
||||
CudaPlatform.log_warnings()
|
||||
@@ -0,0 +1,252 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/interface.py
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import random
|
||||
from typing import TYPE_CHECKING, NamedTuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import resolve_obj_by_qualname
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
AttentionImpl,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class AttentionBackendEnum(enum.Enum):
|
||||
FA3 = enum.auto()
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
AITER = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
def __str__(self):
|
||||
return self.name.lower()
|
||||
|
||||
|
||||
class PlatformEnum(enum.Enum):
|
||||
CUDA = enum.auto()
|
||||
ROCM = enum.auto()
|
||||
TPU = enum.auto()
|
||||
CPU = enum.auto()
|
||||
MPS = enum.auto()
|
||||
OOT = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
|
||||
|
||||
class CpuArchEnum(enum.Enum):
|
||||
X86 = enum.auto()
|
||||
ARM = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
|
||||
|
||||
class DeviceCapability(NamedTuple):
|
||||
major: int
|
||||
minor: int
|
||||
|
||||
def as_version_str(self) -> str:
|
||||
return f"{self.major}.{self.minor}"
|
||||
|
||||
def to_int(self) -> int:
|
||||
"""
|
||||
Express device capability as an integer ``<major><minor>``.
|
||||
|
||||
It is assumed that the minor version is always a single digit.
|
||||
"""
|
||||
assert 0 <= self.minor < 10
|
||||
return self.major * 10 + self.minor
|
||||
|
||||
|
||||
class Platform:
|
||||
_enum: PlatformEnum
|
||||
device_name: str
|
||||
device_type: str
|
||||
|
||||
# available dispatch keys:
|
||||
# check https://github.com/pytorch/pytorch/blob/313dac6c1ca0fa0cde32477509cce32089f8532a/torchgen/model.py#L134 # noqa
|
||||
# use "CPU" as a fallback for platforms not registered in PyTorch
|
||||
dispatch_key: str = "CPU"
|
||||
|
||||
# The torch.compile backend for compiling simple and
|
||||
# standalone functions. The default value is "inductor" to keep
|
||||
# the same behavior as PyTorch.
|
||||
# NOTE: for the forward part of the model, vLLM has another separate
|
||||
# compilation strategy.
|
||||
simple_compile_backend: str = "inductor"
|
||||
|
||||
supported_quantization: list[str] = []
|
||||
|
||||
def is_cuda(self) -> bool:
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
def is_rocm(self) -> bool:
|
||||
return self._enum == PlatformEnum.ROCM
|
||||
|
||||
def is_tpu(self) -> bool:
|
||||
return self._enum == PlatformEnum.TPU
|
||||
|
||||
def is_cpu(self) -> bool:
|
||||
return self._enum == PlatformEnum.CPU
|
||||
|
||||
def is_out_of_tree(self) -> bool:
|
||||
return self._enum == PlatformEnum.OOT
|
||||
|
||||
def is_cuda_alike(self) -> bool:
|
||||
"""Stateless version of :func:`torch.cuda.is_available`."""
|
||||
return self._enum in (PlatformEnum.CUDA, PlatformEnum.ROCM)
|
||||
|
||||
def is_mps(self) -> bool:
|
||||
return self._enum == PlatformEnum.MPS
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls_str(
|
||||
cls,
|
||||
selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
) -> str:
|
||||
"""Get the attention backend class of a device."""
|
||||
return ""
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(
|
||||
cls,
|
||||
device_id: int = 0,
|
||||
) -> DeviceCapability | None:
|
||||
"""Stateless version of :func:`torch.cuda.get_device_capability`."""
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def has_device_capability(
|
||||
cls,
|
||||
capability: tuple[int, int] | int,
|
||||
device_id: int = 0,
|
||||
) -> bool:
|
||||
"""
|
||||
Test whether this platform is compatible with a device capability.
|
||||
|
||||
The ``capability`` argument can either be:
|
||||
|
||||
- A tuple ``(major, minor)``.
|
||||
- An integer ``<major><minor>``. (See :meth:`DeviceCapability.to_int`)
|
||||
"""
|
||||
current_capability = cls.get_device_capability(device_id=device_id)
|
||||
if current_capability is None:
|
||||
return False
|
||||
|
||||
if isinstance(capability, tuple):
|
||||
return current_capability >= capability
|
||||
|
||||
return current_capability.to_int() >= capability
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
"""Get the name of a device."""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
"""Get the uuid of a device, e.g. the PCI bus ID."""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
"""Get the total memory of a device in bytes."""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
"""
|
||||
Check if the current platform supports async output.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def inference_mode(cls):
|
||||
"""A device-specific wrapper of `torch.inference_mode`.
|
||||
|
||||
This wrapper is recommended because some hardware backends such as TPU
|
||||
do not support `torch.inference_mode`. In such a case, they will fall
|
||||
back to `torch.no_grad` by overriding this method.
|
||||
"""
|
||||
return torch.inference_mode(mode=True)
|
||||
|
||||
@classmethod
|
||||
def seed_everything(cls, seed: int | None = None) -> None:
|
||||
"""
|
||||
Set the seed of each random module.
|
||||
`torch.manual_seed` will set seed on all devices.
|
||||
|
||||
Loosely based on: https://github.com/Lightning-AI/pytorch-lightning/blob/2.4.0/src/lightning/fabric/utilities/seed.py#L20
|
||||
"""
|
||||
if seed is not None:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
@classmethod
|
||||
def verify_model_arch(cls, model_arch: str) -> None:
|
||||
"""
|
||||
Verify whether the current platform supports the specified model
|
||||
architecture.
|
||||
|
||||
- This will raise an Error or Warning based on the model support on
|
||||
the current platform.
|
||||
- By default all models are considered supported.
|
||||
"""
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def verify_quantization(cls, quant: str) -> None:
|
||||
"""
|
||||
Verify whether the quantization is supported by the current platform.
|
||||
"""
|
||||
if cls.supported_quantization and quant not in cls.supported_quantization:
|
||||
raise ValueError(
|
||||
f"{quant} quantization is currently not supported in "
|
||||
f"{cls.device_name}."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(
|
||||
cls, device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
"""
|
||||
Return the memory usage in bytes.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
"""
|
||||
Get device specific communicator class for distributed communication.
|
||||
"""
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase" # noqa
|
||||
|
||||
@classmethod
|
||||
def get_cpu_architecture(cls) -> CpuArchEnum:
|
||||
"""Get the CPU architecture of the current platform."""
|
||||
return CpuArchEnum.UNSPECIFIED
|
||||
|
||||
def get_attn_backend(self, *args, **kwargs) -> AttentionImpl:
|
||||
attention_cls_str = self.get_attn_backend_cls_str(*args, **kwargs)
|
||||
return resolve_obj_by_qualname(attention_cls_str)
|
||||
|
||||
|
||||
class UnspecifiedPlatform(Platform):
|
||||
_enum = PlatformEnum.UNSPECIFIED
|
||||
device_type = ""
|
||||
@@ -0,0 +1,88 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
DeviceCapability,
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MpsPlatform(Platform):
|
||||
_enum = PlatformEnum.MPS
|
||||
device_name: str = "mps"
|
||||
device_type: str = "mps"
|
||||
dispatch_key: str = "MPS"
|
||||
device_control_env_var: str = "MPS_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_uuid(cls, device_id: int = 0) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
if enforce_eager:
|
||||
logger.warning(
|
||||
"To see benefits of async output processing, enable MPS "
|
||||
"graph. Since, enforce-eager is enabled, async output "
|
||||
"processor cannot be used"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(
|
||||
cls, device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
return 0.0
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls_str(
|
||||
cls,
|
||||
selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
) -> str:
|
||||
# MPS supports SDPA (Scaled Dot-Product Attention) which is the most compatible
|
||||
logger.info("Using Torch SDPA backend for MPS.")
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
# Use base communicator for MPS
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase"
|
||||
|
||||
@classmethod
|
||||
def seed_everything(cls, seed: int | None = None) -> None:
|
||||
"""Set the seed for MPS device."""
|
||||
if seed is not None:
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
# MPS doesn't have manual_seed_all like CUDA
|
||||
# The manual_seed above should be sufficient
|
||||
@@ -0,0 +1,138 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from rocm/vllm: https://github.com/ROCm/vllm/blob/v0.7.3%2Brocm/vllm/platforms/rocm.py
|
||||
"""
|
||||
This file is a platform abstraction for ROCm GPUs,
|
||||
adjusted to match the structure and interface of `cuda.py`.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.multimodal_gen.envs as envs
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
AttentionBackendEnum,
|
||||
DeviceCapability,
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# ROCm uses the same torch.cuda interface
|
||||
class RocmPlatform(Platform):
|
||||
_enum = PlatformEnum.ROCM
|
||||
device_name: str = "rocm"
|
||||
device_type: str = "cuda" # torch uses 'cuda' backend string
|
||||
dispatch_key: str = "CUDA"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||
major, minor = torch.cuda.get_device_capability(device_id)
|
||||
return DeviceCapability(major=major, minor=minor)
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return str(torch.cuda.get_device_name(device_id))
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
return torch.cuda.get_device_properties(device_id).total_memory
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
if enforce_eager:
|
||||
logger.warning(
|
||||
"To see benefits of async output processing, enable CUDA graph. "
|
||||
"Since enforce-eager is enabled, async output processor cannot be used"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def log_warnings(cls) -> None:
|
||||
pass # ROCm-specific warnings can be added here
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(cls, device: torch.device | None = None) -> float:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls_str(
|
||||
cls,
|
||||
selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
) -> str:
|
||||
logger.info(
|
||||
"Trying SGL_DIFFUSION_ATTENTION_BACKEND=%s",
|
||||
envs.SGL_DIFFUSION_ATTENTION_BACKEND,
|
||||
)
|
||||
|
||||
if selected_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
elif selected_backend in (AttentionBackendEnum.FA3, None):
|
||||
pass
|
||||
|
||||
elif selected_backend in (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
):
|
||||
raise ValueError(
|
||||
f"{selected_backend.name} is not supported on {cls.device_name}."
|
||||
)
|
||||
elif selected_backend:
|
||||
raise ValueError(
|
||||
f"Invalid attention backend for {cls.device_name}: {selected_backend}"
|
||||
)
|
||||
|
||||
target_backend = AttentionBackendEnum.FA3
|
||||
if dtype not in (torch.float16, torch.bfloat16):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention backend for dtype other than "
|
||||
"torch.float16 or torch.bfloat16."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == AttentionBackendEnum.FA3:
|
||||
try:
|
||||
import flash_attn # noqa: F401
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import ( # noqa: F401
|
||||
FlashAttentionBackend,
|
||||
)
|
||||
|
||||
supported_sizes = FlashAttentionBackend.get_supported_head_sizes()
|
||||
if head_size not in supported_sizes:
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for head size %d.",
|
||||
head_size,
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
except ImportError:
|
||||
logger.info(
|
||||
"Cannot use FlashAttention backend because the "
|
||||
"flash_attn package is not found. "
|
||||
"Make sure that flash_attn was built and installed "
|
||||
"(on by default)."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
logger.info("Using Flash Attention backend.")
|
||||
|
||||
return "sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator.CudaCommunicator" # works for ROCm too
|
||||
Reference in New Issue
Block a user