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:
Mick
2025-11-05 12:28:52 -08:00
committed by GitHub
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
249 changed files with 63750 additions and 11 deletions
@@ -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