[diffusion] feat: out of tree platform support (#37547)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
@@ -112,6 +112,10 @@ def _run_diffusion(request: ServeRequest) -> None:
|
||||
_print_diffusion_help(request)
|
||||
return
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import apply_plugin_hooks
|
||||
|
||||
apply_plugin_hooks()
|
||||
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
|
||||
add_multimodal_gen_serve_args,
|
||||
execute_serve_cmd,
|
||||
|
||||
@@ -1,8 +1,40 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
from sglang.multimodal_gen.configs.pipeline_configs import PipelineConfig
|
||||
from sglang.multimodal_gen.configs.sample import SamplingParams
|
||||
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from sglang.utils import LazyImport
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.multimodal_gen.configs.pipeline_configs import PipelineConfig
|
||||
from sglang.multimodal_gen.configs.sample import SamplingParams
|
||||
|
||||
__all__ = ["DiffGenerator", "PipelineConfig", "SamplingParams"]
|
||||
|
||||
DiffGenerator = LazyImport(
|
||||
"sglang.multimodal_gen.runtime.entrypoints.diffusion_generator",
|
||||
"DiffGenerator",
|
||||
)
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name == "PipelineConfig":
|
||||
from sglang.multimodal_gen.configs.pipeline_configs import PipelineConfig
|
||||
|
||||
value = PipelineConfig
|
||||
elif name == "SamplingParams":
|
||||
from sglang.multimodal_gen.configs.sample import SamplingParams
|
||||
|
||||
value = SamplingParams
|
||||
else:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
|
||||
def __dir__() -> list[str]:
|
||||
return sorted({*globals(), *__all__})
|
||||
|
||||
|
||||
# Trigger multimodal CI tests
|
||||
|
||||
@@ -21,13 +21,6 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
|
||||
ModelDeploymentConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
AttentionRequirements,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
get_attn_backend,
|
||||
get_global_forced_attn_backend,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
||||
LAYERWISE_OFFLOAD,
|
||||
)
|
||||
@@ -108,6 +101,10 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
||||
self, server_args
|
||||
) -> AttentionBackendEnum | None:
|
||||
"""Resolve the H3 DiT backend using the selector's precedence."""
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
get_global_forced_attn_backend,
|
||||
)
|
||||
|
||||
selected_backend = get_global_forced_attn_backend()
|
||||
if selected_backend is None:
|
||||
selected_backend, _ = server_args.resolve_component_attention_backend(
|
||||
@@ -301,6 +298,13 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
||||
"validated under torch.compile or the breakable CUDA "
|
||||
"graph; disable them or use --attention-backend fa."
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
AttentionRequirements,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
get_attn_backend,
|
||||
)
|
||||
|
||||
get_attn_backend(
|
||||
self.dit_config.arch_config.attention_head_dim,
|
||||
torch.bfloat16,
|
||||
|
||||
@@ -28,7 +28,6 @@ if TYPE_CHECKING:
|
||||
SGLANG_DIFFUSION_DEBUG_HOST_MEMORY: bool = False
|
||||
SGLANG_DIFFUSION_DEBUG_LAYERWISE_TIMING: bool = False
|
||||
SGLANG_DIFFUSION_DISABLE_LORA_MERGE_CACHE: bool = False
|
||||
SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork"
|
||||
SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda"
|
||||
SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = ""
|
||||
SGLANG_EXTERNAL_MODEL_PACKAGE: str = ""
|
||||
@@ -243,13 +242,8 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE": _lazy_int(
|
||||
"SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE", 4
|
||||
),
|
||||
# Use dedicated multiprocess context for workers.
|
||||
# Both spawn and fork work
|
||||
"SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD": _lazy_str(
|
||||
"SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD", "fork"
|
||||
),
|
||||
# Internal per-worker platform override used by disaggregated role launch.
|
||||
# Empty means normal platform auto-detection.
|
||||
# Select a built-in platform or an installed platform entry point.
|
||||
# Empty means automatic plugin activation followed by built-in detection.
|
||||
"SGLANG_DIFFUSION_PLATFORM_OVERRIDE": _lazy_str(
|
||||
"SGLANG_DIFFUSION_PLATFORM_OVERRIDE", ""
|
||||
),
|
||||
|
||||
@@ -11,6 +11,7 @@ import pickle
|
||||
from collections import namedtuple
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from dataclasses import dataclass
|
||||
from pkgutil import resolve_name
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
@@ -20,9 +21,6 @@ from torch.distributed import Backend, ProcessGroup
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator import (
|
||||
CpuCommunicator,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import all_gather_single
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
@@ -52,6 +50,16 @@ def get_local_torch_device() -> torch.device:
|
||||
return current_platform.get_local_torch_device()
|
||||
|
||||
|
||||
def _resolve_all_to_all_communicator_cls() -> type[DeviceCommunicatorBase]:
|
||||
qualname = current_platform.get_all_to_all_communicator_cls()
|
||||
communicator_cls = resolve_name(qualname)
|
||||
if not isinstance(communicator_cls, type) or not issubclass(
|
||||
communicator_cls, DeviceCommunicatorBase
|
||||
):
|
||||
raise TypeError(f"Expected a DeviceCommunicatorBase subclass: {qualname}")
|
||||
return communicator_cls
|
||||
|
||||
|
||||
def _get_unique_name(name: str) -> str:
|
||||
"""Get a unique name for the group.
|
||||
Example:
|
||||
@@ -162,7 +170,7 @@ class GroupCoordinator:
|
||||
cpu_group: ProcessGroup # group for CPU communication
|
||||
device_group: ProcessGroup # group for device communication
|
||||
use_device_communicator: bool # whether to use device communicator
|
||||
device_communicator: DeviceCommunicatorBase # device communicator
|
||||
device_communicator: DeviceCommunicatorBase # all_to_all_4D communicator
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -202,26 +210,13 @@ class GroupCoordinator:
|
||||
self.use_device_communicator = use_device_communicator
|
||||
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
|
||||
if use_device_communicator and self.world_size > 1:
|
||||
# Platform-aware device communicator selection
|
||||
if current_platform.is_cuda_alike():
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator,
|
||||
)
|
||||
|
||||
self.device_communicator = CudaCommunicator(
|
||||
cpu_group=self.cpu_group,
|
||||
device=self.device,
|
||||
device_group=self.device_group,
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
else:
|
||||
# For MPS and CPU, use the CPU communicator
|
||||
self.device_communicator = CpuCommunicator(
|
||||
cpu_group=self.cpu_group,
|
||||
device=self.device,
|
||||
device_group=self.device_group,
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
communicator_cls = _resolve_all_to_all_communicator_cls()
|
||||
self.device_communicator = communicator_cls(
|
||||
cpu_group=self.cpu_group,
|
||||
device=self.device,
|
||||
device_group=self.device_group,
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
|
||||
self.mq_broadcaster = None
|
||||
self.srt_custom_allreduce = None
|
||||
|
||||
@@ -252,17 +252,10 @@ def init_distributed_environment(
|
||||
"distributed environment"
|
||||
)
|
||||
|
||||
# For MPS, MUSA, and XPU, don't pass device_id as it doesn't support device indices
|
||||
extra_args = (
|
||||
{}
|
||||
if (
|
||||
current_platform.is_mps()
|
||||
or current_platform.is_musa()
|
||||
or current_platform.is_npu()
|
||||
or current_platform.is_cpu()
|
||||
or current_platform.is_xpu()
|
||||
)
|
||||
else dict(device_id=device_id)
|
||||
dict(device_id=device_id)
|
||||
if current_platform.supports_distributed_device_id()
|
||||
else {}
|
||||
)
|
||||
|
||||
if timeout is not None:
|
||||
|
||||
@@ -4,12 +4,20 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
|
||||
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcommand
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.generate import GenerateSubcommand
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import ServeSubcommand
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import apply_plugin_hooks
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
|
||||
|
||||
def generate_cmd_init() -> list[CLISubcommand]:
|
||||
# Command modules import the runtime graph. Activate plugins first so OOT
|
||||
# platforms can prepare that graph before its modules are evaluated.
|
||||
apply_plugin_hooks()
|
||||
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.generate import (
|
||||
GenerateSubcommand,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import ServeSubcommand
|
||||
|
||||
return [GenerateSubcommand(), ServeSubcommand()]
|
||||
|
||||
|
||||
@@ -21,6 +29,8 @@ def cmd_init() -> list[CLISubcommand]:
|
||||
|
||||
|
||||
def main() -> None:
|
||||
apply_plugin_hooks()
|
||||
|
||||
parser = FlexibleArgumentParser(description="sglang-diffusion CLI")
|
||||
parser.add_argument("-v", "--version", action="version", version="0.1.0")
|
||||
|
||||
|
||||
@@ -34,9 +34,9 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||
prepare_request,
|
||||
save_outputs,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.launch_server import launch_server
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import apply_plugin_hooks
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import sync_scheduler_client
|
||||
from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs
|
||||
from sglang.multimodal_gen.runtime.server_warmup import (
|
||||
@@ -58,15 +58,6 @@ from sglang.multimodal_gen.runtime.utils.trace_wrapper import (
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
try:
|
||||
# Set the start method to 'spawn' to avoid CUDA errors in forked processes.
|
||||
# This must be done at the top level of the module, before any CUDA context
|
||||
# or other processes are initialized.
|
||||
mp.set_start_method("spawn", force=True)
|
||||
except RuntimeError:
|
||||
# The start method can only be set once per program execution.
|
||||
pass
|
||||
|
||||
|
||||
def _replace_sampling_params_for_prompt(
|
||||
sampling_params_orig: SamplingParams,
|
||||
@@ -137,6 +128,10 @@ class DiffGenerator:
|
||||
|
||||
Priority level: Default pipeline config < User's pipeline config < User's kwargs
|
||||
"""
|
||||
# Not shared with from_server_args: the ServerArgs built below runs
|
||||
# Platform.apply_server_args_defaults, which hooks must precede.
|
||||
apply_plugin_hooks()
|
||||
|
||||
# If users also provide some kwargs, it will override the ServerArgs and PipelineConfig.
|
||||
|
||||
if (server_args := kwargs.get("server_args", None)) is not None:
|
||||
@@ -147,7 +142,7 @@ class DiffGenerator:
|
||||
else:
|
||||
server_args = ServerArgs.from_kwargs(**kwargs)
|
||||
|
||||
return cls.from_server_args(server_args, local_mode=local_mode)
|
||||
return cls._create(server_args, local_mode=local_mode)
|
||||
|
||||
@classmethod
|
||||
def from_server_args(
|
||||
@@ -162,6 +157,16 @@ class DiffGenerator:
|
||||
Returns:
|
||||
The created DiffGenerator
|
||||
"""
|
||||
apply_plugin_hooks()
|
||||
return cls._create(server_args, local_mode=local_mode)
|
||||
|
||||
@classmethod
|
||||
def _create(cls, server_args: ServerArgs, *, local_mode: bool) -> "DiffGenerator":
|
||||
"""Build and connect a generator, assuming hooks are already applied.
|
||||
|
||||
Each public constructor owns that step itself, so this shared body must
|
||||
not repeat it.
|
||||
"""
|
||||
globally_suppress_loggers()
|
||||
instance = cls(
|
||||
server_args=server_args,
|
||||
@@ -184,6 +189,9 @@ class DiffGenerator:
|
||||
self,
|
||||
) -> list[mp.Process]:
|
||||
"""Check if a local server is running; if not, start it and return the process handles."""
|
||||
# Not module scope: launch_server pulls in the whole worker graph.
|
||||
from sglang.multimodal_gen.runtime.launch_server import launch_server
|
||||
|
||||
# First, we need a client to test the server. Initialize it temporarily.
|
||||
sync_scheduler_client.initialize(self.server_args)
|
||||
|
||||
|
||||
@@ -14,11 +14,17 @@ from sglang.multimodal_gen.runtime.disaggregation.orchestrator import (
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.entrypoints.control_requests import ShutdownReq
|
||||
from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app
|
||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process
|
||||
from sglang.multimodal_gen.runtime.managers.worker_bootstrap import (
|
||||
SchedulerProcessSpec,
|
||||
ServerArgsPayload,
|
||||
bootstrap_http_server_process,
|
||||
bootstrap_scheduler_process,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import (
|
||||
configure_metrics,
|
||||
start_role_metrics_server,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import apply_plugin_hooks
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import SchedulerClient
|
||||
from sglang.multimodal_gen.runtime.server_args import (
|
||||
ServerArgs,
|
||||
@@ -28,7 +34,6 @@ from sglang.multimodal_gen.runtime.server_args import (
|
||||
from sglang.multimodal_gen.runtime.utils.common import is_port_available
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger
|
||||
from sglang.multimodal_gen.runtime.utils.process import (
|
||||
kill_itself_when_parent_died,
|
||||
kill_process_tree,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.trace_wrapper import init_diffusion_tracing
|
||||
@@ -97,11 +102,6 @@ def _kill_alive_processes(processes, timeout_s: float) -> None:
|
||||
_join_processes_with_deadline(alive, timeout_s)
|
||||
|
||||
|
||||
def _run_http_server_process(server_args: ServerArgs) -> None:
|
||||
kill_itself_when_parent_died()
|
||||
launch_http_server_only(server_args)
|
||||
|
||||
|
||||
def _request_monolithic_scheduler_shutdown(server_args: ServerArgs) -> None:
|
||||
if server_args.disagg_role != RoleType.MONOLITHIC:
|
||||
return
|
||||
@@ -138,6 +138,7 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
Args:
|
||||
launch_http_server: False for offline local mode
|
||||
"""
|
||||
apply_plugin_hooks()
|
||||
configure_logger(server_args)
|
||||
|
||||
# Start a new server with multiple worker processes
|
||||
@@ -156,17 +157,28 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
rank_offset = node_rank * local_num_gpus
|
||||
processes = []
|
||||
|
||||
# A local spawn context makes the worker boundary deterministic even when
|
||||
# an embedding application selected a different global start method.
|
||||
worker_context = mp.get_context("spawn")
|
||||
server_args_payload = ServerArgsPayload.capture(server_args)
|
||||
|
||||
# Launch this node's local worker processes.
|
||||
scheduler_pipe_readers = []
|
||||
scheduler_pipe_writers = []
|
||||
|
||||
for i in range(local_num_gpus):
|
||||
rank = rank_offset + i
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
reader, writer = worker_context.Pipe(duplex=False)
|
||||
scheduler_pipe_writers.append(writer)
|
||||
process = mp.Process(
|
||||
target=run_scheduler_process,
|
||||
args=(i, rank, server_args, writer),
|
||||
spec = SchedulerProcessSpec(
|
||||
local_rank=i,
|
||||
rank=rank,
|
||||
server_args=server_args_payload,
|
||||
pipe_writer=writer,
|
||||
)
|
||||
process = worker_context.Process(
|
||||
target=bootstrap_scheduler_process,
|
||||
args=(spec,),
|
||||
name=f"sglang-diffusionWorker-{rank}",
|
||||
daemon=True,
|
||||
)
|
||||
@@ -226,9 +238,9 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
logger.info("Starting FastAPI server.")
|
||||
if server_args.webui:
|
||||
logger.info("Launch FastAPI server in another process because of webui.")
|
||||
http_server_process = mp.Process(
|
||||
target=_run_http_server_process,
|
||||
args=(server_args,),
|
||||
http_server_process = worker_context.Process(
|
||||
target=bootstrap_http_server_process,
|
||||
args=(server_args_payload,),
|
||||
name="sglang-diffusion-webui",
|
||||
daemon=True,
|
||||
)
|
||||
@@ -370,6 +382,7 @@ def launch_pool_disagg_server(
|
||||
base_dict.update(role_overrides)
|
||||
base_dict.pop("pipeline_config", None)
|
||||
role_args = ServerArgs.from_kwargs(**base_dict)
|
||||
role_args_payload = ServerArgsPayload.capture(role_args)
|
||||
|
||||
pool_ctx = mp.get_context("spawn")
|
||||
inst_readers = []
|
||||
@@ -379,9 +392,18 @@ def launch_pool_disagg_server(
|
||||
reader, writer = pool_ctx.Pipe(duplex=False)
|
||||
gpu_id = gpu_ids[rank_idx]
|
||||
|
||||
# Physical GPU index as local_rank: torch.cuda.set_device() must
|
||||
# not depend on CUDA_VISIBLE_DEVICES remapping, which can be
|
||||
# stale if CUDA was already initialized in the parent.
|
||||
spec = SchedulerProcessSpec(
|
||||
local_rank=gpu_id,
|
||||
rank=rank_idx,
|
||||
server_args=role_args_payload,
|
||||
pipe_writer=writer,
|
||||
)
|
||||
process = pool_ctx.Process(
|
||||
target=_run_disagg_role_process,
|
||||
args=(gpu_id, rank_idx, role_args, writer),
|
||||
target=bootstrap_scheduler_process,
|
||||
args=(spec,),
|
||||
name=f"sglang-pool-{role_type.value}-{inst_idx}-r{rank_idx}",
|
||||
daemon=True,
|
||||
)
|
||||
@@ -454,27 +476,6 @@ def launch_pool_disagg_server(
|
||||
return all_processes
|
||||
|
||||
|
||||
def _run_disagg_role_process(
|
||||
gpu_id: int,
|
||||
rank: int,
|
||||
server_args: ServerArgs,
|
||||
pipe_writer: mp.connection.Connection,
|
||||
):
|
||||
"""Entry point for a disagg role process.
|
||||
|
||||
Uses the physical GPU index (gpu_id) as local_rank so that
|
||||
torch.cuda.set_device(local_rank) selects the correct GPU.
|
||||
This avoids relying on CUDA_VISIBLE_DEVICES remapping, which
|
||||
may not work if CUDA was pre-initialized in the parent process.
|
||||
"""
|
||||
run_scheduler_process(
|
||||
local_rank=gpu_id,
|
||||
rank=rank,
|
||||
server_args=server_args,
|
||||
pipe_writer=pipe_writer,
|
||||
)
|
||||
|
||||
|
||||
def launch_http_server_only(server_args):
|
||||
init_diffusion_tracing(server_args, "DiffHTTPServer")
|
||||
|
||||
@@ -702,6 +703,7 @@ def launch_disagg_role(server_args: ServerArgs):
|
||||
base_dict.update(role_overrides)
|
||||
base_dict.pop("pipeline_config", None)
|
||||
role_args = ServerArgs.from_kwargs(**base_dict)
|
||||
role_args_payload = ServerArgsPayload.capture(role_args)
|
||||
|
||||
# Spawn GPU worker processes
|
||||
# NOTE: All ranks must be spawned before waiting for ready signals,
|
||||
@@ -716,9 +718,18 @@ def launch_disagg_role(server_args: ServerArgs):
|
||||
reader, writer = pool_ctx.Pipe(duplex=False)
|
||||
gpu_id = base_gpu_id + rank_idx
|
||||
|
||||
# Physical GPU index as local_rank: torch.cuda.set_device() must not
|
||||
# depend on CUDA_VISIBLE_DEVICES remapping, which can be stale if CUDA
|
||||
# was already initialized in the parent.
|
||||
spec = SchedulerProcessSpec(
|
||||
local_rank=gpu_id,
|
||||
rank=rank_idx,
|
||||
server_args=role_args_payload,
|
||||
pipe_writer=writer,
|
||||
)
|
||||
process = pool_ctx.Process(
|
||||
target=_run_disagg_role_process,
|
||||
args=(gpu_id, rank_idx, role_args, writer),
|
||||
target=bootstrap_scheduler_process,
|
||||
args=(spec,),
|
||||
name=f"sglang-{role_type.value}-r{rank_idx}",
|
||||
daemon=True,
|
||||
)
|
||||
@@ -764,6 +775,8 @@ def launch_disagg_role(server_args: ServerArgs):
|
||||
|
||||
def dispatch_launch(server_args: ServerArgs):
|
||||
"""Route to the correct launch function based on --disagg-role."""
|
||||
apply_plugin_hooks()
|
||||
|
||||
if "NCCL_NVLS_ENABLE" not in os.environ or server_args.enable_nccl_nvls:
|
||||
os.environ["NCCL_NVLS_ENABLE"] = str(int(server_args.enable_nccl_nvls))
|
||||
|
||||
@@ -779,6 +792,7 @@ def dispatch_launch(server_args: ServerArgs):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
apply_plugin_hooks()
|
||||
server_args = prepare_server_args(sys.argv[1:])
|
||||
|
||||
try:
|
||||
|
||||
@@ -4,16 +4,16 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/custom_op.py
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
from functools import partial
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
import sglang.multimodal_gen.runtime.platforms as platforms
|
||||
from sglang.kernels.kernel_api_logging import debug_kernel_api
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_is_cuda = current_platform.is_cuda()
|
||||
|
||||
|
||||
class CustomOp(nn.Module):
|
||||
@@ -22,6 +22,15 @@ class CustomOp(nn.Module):
|
||||
Dispatches the forward method to the appropriate backend.
|
||||
"""
|
||||
|
||||
_oot_forward_registry: ClassVar[dict[str, dict[type["CustomOp"], Callable]]] = {}
|
||||
|
||||
@staticmethod
|
||||
def register_oot_forward(
|
||||
op_cls: type["CustomOp"], *, fn: Callable, platform_key: str
|
||||
) -> None:
|
||||
"""Register ``fn`` for an exact op class and behavioral dispatch key."""
|
||||
CustomOp._oot_forward_registry.setdefault(platform_key, {})[op_cls] = fn
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._forward_method = self.dispatch_forward()
|
||||
@@ -72,18 +81,45 @@ class CustomOp(nn.Module):
|
||||
def forward_xpu(self, *args, **kwargs) -> Any:
|
||||
return self.forward_native(*args, **kwargs)
|
||||
|
||||
def _defined_forward(self, method_name: str) -> Callable | None:
|
||||
"""Return an implementation defined below ``CustomOp`` in the MRO."""
|
||||
for op_cls in type(self).__mro__:
|
||||
if op_cls is CustomOp:
|
||||
return None
|
||||
if method_name in op_cls.__dict__:
|
||||
return getattr(self, method_name)
|
||||
return None
|
||||
|
||||
def dispatch_forward(self) -> Callable:
|
||||
if _is_cuda:
|
||||
platform = platforms.current_platform
|
||||
if platform.is_out_of_tree():
|
||||
# An empty key would silently skip the platform forward below and
|
||||
# dispatch everything to forward_oot instead.
|
||||
platform_key = platform.get_dispatch_key_name().strip()
|
||||
if not platform_key:
|
||||
raise ValueError(
|
||||
"Out-of-tree diffusion platforms must return a non-empty "
|
||||
"get_dispatch_key_name()"
|
||||
)
|
||||
forward = self._oot_forward_registry.get(platform_key, {}).get(type(self))
|
||||
if forward is not None:
|
||||
return partial(forward, self)
|
||||
if platform_key.isidentifier():
|
||||
platform_forward = self._defined_forward(f"forward_{platform_key}")
|
||||
if platform_forward is not None:
|
||||
return platform_forward
|
||||
return self.forward_oot
|
||||
elif platform.is_cuda():
|
||||
return self.forward_cuda
|
||||
elif current_platform.is_hip():
|
||||
elif platform.is_hip():
|
||||
return self.forward_hip
|
||||
elif current_platform.is_npu():
|
||||
elif platform.is_npu():
|
||||
return self.forward_npu
|
||||
elif current_platform.is_xpu():
|
||||
elif platform.is_xpu():
|
||||
return self.forward_xpu
|
||||
elif current_platform.is_musa():
|
||||
elif platform.is_musa():
|
||||
return self.forward_musa
|
||||
elif current_platform.is_cpu():
|
||||
elif platform.is_cpu():
|
||||
return self.forward_cpu
|
||||
else:
|
||||
return self.forward_native
|
||||
|
||||
@@ -10,8 +10,8 @@ from typing import Any, Callable, List, Optional
|
||||
import torch
|
||||
from torch.library import Library
|
||||
|
||||
import sglang.multimodal_gen.runtime.platforms as platforms
|
||||
from sglang.kernels.kernel_api_logging import debug_torch_op
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
|
||||
|
||||
def get_group_size(group) -> int:
|
||||
@@ -45,7 +45,7 @@ def direct_register_custom_op(
|
||||
"""
|
||||
`torch.library.custom_op` can have significant overhead because it
|
||||
needs to consider complicated dispatching logic. This function
|
||||
directly registers a custom op and dispatches it to the CUDA backend.
|
||||
directly registers a custom op for the active platform's dispatch key.
|
||||
See https://gist.github.com/youkaichao/ecbea9ec9fc79a45d2adce1784d7a9a5
|
||||
for more details.
|
||||
|
||||
@@ -90,7 +90,9 @@ def direct_register_custom_op(
|
||||
try:
|
||||
my_lib.define(op_name + schema_str)
|
||||
my_lib.impl(
|
||||
op_name, op_func, "CUDA" if not current_platform.is_npu() else "PrivateUse1"
|
||||
op_name,
|
||||
op_func,
|
||||
platforms.current_platform.get_torch_library_dispatch_key(),
|
||||
)
|
||||
if fake_impl is not None:
|
||||
my_lib._register_fake(op_name, fake_impl)
|
||||
|
||||
@@ -78,7 +78,10 @@ from sglang.multimodal_gen.runtime.pipelines_core import (
|
||||
build_pipeline,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.platforms import (
|
||||
current_platform,
|
||||
initialize_current_platform,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.post_training.gpu_worker_post_training_mixin import (
|
||||
GPUWorkerPostTrainingMixin,
|
||||
)
|
||||
@@ -1573,9 +1576,14 @@ def run_scheduler_process(
|
||||
pipe_writer: mp.connection.Connection,
|
||||
) -> None:
|
||||
"""Run a rank's scheduler and report readiness to the launching process."""
|
||||
# Idempotent safeguard for direct callers; process bootstraps already
|
||||
# initialized the platform before this module was imported.
|
||||
initialize_current_platform()
|
||||
|
||||
kill_itself_when_parent_died()
|
||||
configure_logger(server_args)
|
||||
globally_suppress_loggers()
|
||||
|
||||
if current_platform.is_cuda():
|
||||
set_cuda_arch()
|
||||
elif current_platform.is_musa():
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
"""Import-safe specifications and entry points for diffusion child processes.
|
||||
|
||||
``multiprocessing`` unpickles a target's arguments before calling the target.
|
||||
Runtime objects therefore cannot cross this boundary directly: merely passing
|
||||
``ServerArgs`` used to import the diffusion configuration graph before worker
|
||||
bootstrap began. ``ServerArgsPayload`` keeps that object graph opaque until the
|
||||
child reaches the explicit runtime-activation phase.
|
||||
|
||||
``spawn`` re-executes the launching script's module scope earlier still, before
|
||||
any argument is unpickled, so an offline script may bind only the
|
||||
``DiffGenerator`` proxy and ``_PRE_ACTIVATION_MODULES`` at module scope; every
|
||||
other diffusion import belongs inside its ``if __name__ == "__main__":`` guard.
|
||||
A violation is reported rather than silently tolerated.
|
||||
|
||||
Keep module scope limited to the standard library and import-neutral types.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import pickle
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from multiprocessing.connection import Connection
|
||||
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
_DIFFUSION_PREFIX = "sglang.multimodal_gen."
|
||||
_RUNTIME_NAMESPACES = (
|
||||
"sglang.multimodal_gen.runtime",
|
||||
"sglang.multimodal_gen.runtime.managers",
|
||||
)
|
||||
# What may legitimately be imported this early: bootstrap's own imports, and the
|
||||
# platform and plugin modules that every plugin loads and the contract keeps
|
||||
# import-safe. Listing these rather than their complement keeps the check
|
||||
# complete as subpackages are added.
|
||||
_PRE_ACTIVATION_MODULES = (
|
||||
"sglang.multimodal_gen.envs",
|
||||
"sglang.multimodal_gen.runtime.platforms",
|
||||
"sglang.multimodal_gen.runtime.utils",
|
||||
"sglang.multimodal_gen.runtime.managers.worker_bootstrap",
|
||||
)
|
||||
_MAX_REPORTED_MODULES = 5
|
||||
|
||||
|
||||
def _warn_if_runtime_imported_early() -> None:
|
||||
"""Name the modules that this child imported ahead of its own lifecycle."""
|
||||
early = sorted(
|
||||
name
|
||||
for name in list(sys.modules)
|
||||
if name.startswith(_DIFFUSION_PREFIX)
|
||||
and name not in _RUNTIME_NAMESPACES
|
||||
and not name.startswith(_PRE_ACTIVATION_MODULES)
|
||||
)
|
||||
if not early:
|
||||
return
|
||||
|
||||
listed = ", ".join(early[:_MAX_REPORTED_MODULES])
|
||||
if len(early) > _MAX_REPORTED_MODULES:
|
||||
listed += f" (+{len(early) - _MAX_REPORTED_MODULES} more)"
|
||||
# In a spawned child __main__ is the re-executed launching script, which is
|
||||
# the file whose imports have to move.
|
||||
script = vars(sys.modules["__main__"]).get("__file__", "the launching script")
|
||||
logging.getLogger(__name__).warning(
|
||||
"Diffusion runtime modules were imported before this worker initialized "
|
||||
"its platform: %s. spawn re-executes %s at module scope in every child, "
|
||||
"so these were built ahead of platform initialization and hook "
|
||||
"application, and the classes and registrations they created are "
|
||||
'already past reach. Move the import inside if __name__ == "__main__": '
|
||||
"or into the function that uses it.",
|
||||
listed,
|
||||
script,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ServerArgsPayload:
|
||||
"""A deferred ``ServerArgs`` snapshot safe to unpickle before bootstrap."""
|
||||
|
||||
_pickle: bytes = field(repr=False)
|
||||
|
||||
@classmethod
|
||||
def capture(cls, server_args: ServerArgs) -> ServerArgsPayload:
|
||||
return cls(pickle.dumps(server_args, protocol=pickle.HIGHEST_PROTOCOL))
|
||||
|
||||
def materialize(self) -> ServerArgs:
|
||||
# Importing ServerArgs pulls in pipeline configuration modules. This
|
||||
# method must only be called after the process lifecycle is initialized.
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
server_args = pickle.loads(self._pickle)
|
||||
if not isinstance(server_args, ServerArgs):
|
||||
raise TypeError("Bootstrap payload did not contain diffusion ServerArgs")
|
||||
return server_args
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SchedulerProcessSpec:
|
||||
"""Everything a scheduler child needs, without eagerly importing runtime state."""
|
||||
|
||||
local_rank: int
|
||||
rank: int
|
||||
server_args: ServerArgsPayload
|
||||
pipe_writer: Connection
|
||||
|
||||
|
||||
def bootstrap_scheduler_process(spec: SchedulerProcessSpec) -> None:
|
||||
"""Initialize a child in dependency order, then invoke its worker."""
|
||||
# Arm PDEATHSIG before any vendor code runs: everything below can block,
|
||||
# and a child that hangs there would outlive a dead launcher.
|
||||
from sglang.multimodal_gen.runtime.utils.process import (
|
||||
kill_itself_when_parent_died,
|
||||
)
|
||||
|
||||
kill_itself_when_parent_died()
|
||||
|
||||
# Every rank re-executes the same script, so one rank reporting is enough.
|
||||
if spec.rank == 0:
|
||||
_warn_if_runtime_imported_early()
|
||||
|
||||
# Platform initialization is the first extensible runtime action. In
|
||||
# particular it precedes plugin callbacks and hook target resolution, both
|
||||
# of which may import arbitrary runtime modules.
|
||||
from sglang.multimodal_gen.runtime.platforms import initialize_current_platform
|
||||
|
||||
initialize_current_platform()
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import (
|
||||
apply_plugin_hooks,
|
||||
load_plugins,
|
||||
)
|
||||
|
||||
load_plugins()
|
||||
apply_plugin_hooks()
|
||||
|
||||
server_args = spec.server_args.materialize()
|
||||
|
||||
# Resolve the function from its module after hook application. A ``from``
|
||||
# binding created earlier would retain the unpatched callable.
|
||||
from sglang.multimodal_gen.runtime.managers import gpu_worker
|
||||
|
||||
gpu_worker.run_scheduler_process(
|
||||
local_rank=spec.local_rank,
|
||||
rank=spec.rank,
|
||||
server_args=server_args,
|
||||
pipe_writer=spec.pipe_writer,
|
||||
)
|
||||
|
||||
|
||||
def bootstrap_http_server_process(server_args: ServerArgsPayload) -> None:
|
||||
from sglang.multimodal_gen.runtime.utils.process import (
|
||||
kill_itself_when_parent_died,
|
||||
)
|
||||
|
||||
kill_itself_when_parent_died()
|
||||
|
||||
_warn_if_runtime_imported_early()
|
||||
|
||||
# No initialize_current_platform() here: this child serves HTTP and never
|
||||
# touches the device, so it has no reason to bring up a vendor backend.
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import (
|
||||
apply_plugin_hooks,
|
||||
load_plugins,
|
||||
)
|
||||
|
||||
load_plugins()
|
||||
apply_plugin_hooks()
|
||||
|
||||
from sglang.multimodal_gen.runtime import launch_server
|
||||
|
||||
launch_server.launch_http_server_only(server_args.materialize())
|
||||
@@ -9,7 +9,6 @@ This package contains diffusion pipelines for generating videos and images.
|
||||
|
||||
from typing import cast
|
||||
|
||||
from sglang.multimodal_gen.registry import get_model_info
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
@@ -43,6 +42,8 @@ def build_pipeline(
|
||||
2. verify the model config and directory
|
||||
3. based on the config, determine the pipeline class
|
||||
"""
|
||||
from sglang.multimodal_gen.registry import get_model_info
|
||||
|
||||
model_path = server_args.model_path
|
||||
|
||||
# Check if pipeline class is explicitly specified
|
||||
|
||||
@@ -45,8 +45,7 @@ from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.torch_compile import (
|
||||
ActiveTargetCompiledCallable,
|
||||
build_torch_compile_kwargs,
|
||||
resolve_torch_compile_mode,
|
||||
resolve_torch_compile_kwargs,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -185,17 +184,16 @@ class DecodingStage(PipelineStage):
|
||||
compiled_callable.target_id != id(vae)
|
||||
or compiled_callable.compiled_module is None
|
||||
)
|
||||
compile_kwargs, mode = resolve_torch_compile_kwargs(
|
||||
"SGLANG_VAE_TORCH_COMPILE_MODE",
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
default="default",
|
||||
module=vae,
|
||||
)
|
||||
if current_platform.is_npu():
|
||||
compile_kwargs = build_torch_compile_kwargs(mode=None)
|
||||
if will_compile:
|
||||
logger.info("Compiling VAE decode with torchair backend on NPU")
|
||||
else:
|
||||
mode = resolve_torch_compile_mode(
|
||||
"SGLANG_VAE_TORCH_COMPILE_MODE",
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
default="default",
|
||||
)
|
||||
compile_kwargs = build_torch_compile_kwargs(mode=mode)
|
||||
if will_compile:
|
||||
logger.info("Compiling VAE decode with mode: %s", mode)
|
||||
|
||||
|
||||
@@ -171,9 +171,7 @@ from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
||||
from sglang.multimodal_gen.runtime.utils.torch_compile import (
|
||||
CompiledModuleRegistry,
|
||||
build_torch_compile_kwargs,
|
||||
maybe_enable_inductor_compute_comm_overlap,
|
||||
resolve_torch_compile_mode,
|
||||
resolve_torch_compile_kwargs,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -553,18 +551,17 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
||||
if self._torch_compile_registry.is_compiled(module):
|
||||
return
|
||||
|
||||
dit_config = getattr(self.server_args.pipeline_config, "dit_config", None)
|
||||
compile_kwargs, mode = resolve_torch_compile_kwargs(
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
config=dit_config,
|
||||
default="max-autotune-no-cudagraphs",
|
||||
module=module,
|
||||
enable_inductor_compute_comm_overlap=True,
|
||||
)
|
||||
if current_platform.is_npu():
|
||||
compile_kwargs = build_torch_compile_kwargs(mode=None)
|
||||
logger.info("Compiling transformer with torchair backend on NPU")
|
||||
else:
|
||||
maybe_enable_inductor_compute_comm_overlap()
|
||||
dit_config = getattr(self.server_args.pipeline_config, "dit_config", None)
|
||||
mode = resolve_torch_compile_mode(
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
config=dit_config,
|
||||
default="max-autotune-no-cudagraphs",
|
||||
)
|
||||
compile_kwargs = build_torch_compile_kwargs(mode=mode, module=module)
|
||||
logger.info(f"Compiling transformer with mode: {mode}")
|
||||
|
||||
if getattr(self.server_args, "regional_compile", False):
|
||||
|
||||
+10
-19
@@ -13,7 +13,6 @@ from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import inspect
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -73,7 +72,9 @@ from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
||||
from sglang.srt.utils.common import get_compiler_backend
|
||||
from sglang.multimodal_gen.runtime.utils.torch_compile import (
|
||||
resolve_torch_compile_kwargs,
|
||||
)
|
||||
|
||||
_is_npu = current_platform.is_npu()
|
||||
logger = init_logger(__name__)
|
||||
@@ -265,29 +266,19 @@ class MOVADenoisingStage(PipelineStage):
|
||||
module.__class__.__name__,
|
||||
)
|
||||
return
|
||||
compile_kwargs: dict[str, object] = {"fullgraph": False, "dynamic": None}
|
||||
|
||||
compile_kwargs, mode = resolve_torch_compile_kwargs(
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
config=model_config,
|
||||
default="max-autotune-no-cudagraphs",
|
||||
module=module,
|
||||
enable_inductor_compute_comm_overlap=True,
|
||||
)
|
||||
if current_platform.is_npu():
|
||||
backend = get_compiler_backend()
|
||||
compile_kwargs["backend"] = backend
|
||||
compile_kwargs["dynamic"] = False
|
||||
logger.info(
|
||||
"Compiling %s with torchair backend on NPU",
|
||||
module.__class__.__name__,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
import torch._inductor.config as _inductor_cfg
|
||||
|
||||
_inductor_cfg.reorder_for_compute_comm_overlap = True
|
||||
except ImportError:
|
||||
pass
|
||||
mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE") or getattr(
|
||||
model_config,
|
||||
"torch_compile_mode",
|
||||
"max-autotune-no-cudagraphs",
|
||||
)
|
||||
compile_kwargs["mode"] = mode
|
||||
logger.info("Compiling %s with mode: %s", module.__class__.__name__, mode)
|
||||
|
||||
# TODO(triple-mu): support customized fullgraph and dynamic in the future
|
||||
|
||||
@@ -5,8 +5,14 @@
|
||||
|
||||
import os
|
||||
import traceback
|
||||
from collections import Counter
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from importlib.metadata import EntryPoint, entry_points
|
||||
from pkgutil import resolve_name
|
||||
|
||||
from sglang.multimodal_gen import envs
|
||||
|
||||
# imported by other files, do not remove
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import ( # noqa: F401
|
||||
AttentionBackendEnum,
|
||||
@@ -18,6 +24,18 @@ from sglang.multimodal_gen.third_party import pynvml
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
PLATFORM_PLUGINS_GROUP = "sglang.multimodal_gen.platforms"
|
||||
_BUILTIN_PLATFORM_QUALNAMES = {
|
||||
"cpu": "sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform",
|
||||
"cuda": "sglang.multimodal_gen.runtime.platforms.cuda.CudaPlatform",
|
||||
"rocm": "sglang.multimodal_gen.runtime.platforms.rocm.RocmPlatform",
|
||||
"xpu": "sglang.multimodal_gen.runtime.platforms.xpu.XpuPlatform",
|
||||
"mps": "sglang.multimodal_gen.runtime.platforms.mps.MpsPlatform",
|
||||
"npu": "sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase",
|
||||
"musa": "sglang.multimodal_gen.runtime.platforms.musa.MusaPlatform",
|
||||
}
|
||||
BUILTIN_PLATFORM_NAMES = frozenset(_BUILTIN_PLATFORM_QUALNAMES)
|
||||
|
||||
|
||||
def cuda_platform_plugin() -> str | None:
|
||||
is_cuda = False
|
||||
@@ -72,9 +90,7 @@ def cuda_platform_plugin() -> str | None:
|
||||
if is_cuda:
|
||||
logger.debug("CUDA is available")
|
||||
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.cuda.CudaPlatform" if is_cuda else None
|
||||
)
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["cuda"] if is_cuda else None
|
||||
|
||||
|
||||
def mps_platform_plugin() -> str | None:
|
||||
@@ -90,13 +106,13 @@ def mps_platform_plugin() -> str | None:
|
||||
except Exception as e:
|
||||
logger.debug("MPS detection failed: %s", e)
|
||||
|
||||
return "sglang.multimodal_gen.runtime.platforms.mps.MpsPlatform" if is_mps else None
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["mps"] if is_mps else None
|
||||
|
||||
|
||||
def cpu_platform_plugin() -> str | None:
|
||||
def cpu_platform_plugin() -> str:
|
||||
"""Detect if CPU platform should be used."""
|
||||
# CPU is always available as a fallback
|
||||
return "sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform"
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["cpu"]
|
||||
|
||||
|
||||
def rocm_platform_plugin() -> str | None:
|
||||
@@ -115,9 +131,7 @@ def rocm_platform_plugin() -> str | None:
|
||||
except Exception as e:
|
||||
logger.debug("ROCm platform is unavailable: %s", e)
|
||||
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.rocm.RocmPlatform" if is_rocm else None
|
||||
)
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["rocm"] if is_rocm else None
|
||||
|
||||
|
||||
def npu_platform_plugin() -> str | None:
|
||||
@@ -131,11 +145,7 @@ def npu_platform_plugin() -> str | None:
|
||||
logger.debug("NPU is available")
|
||||
except Exception as e:
|
||||
logger.debug("NPU detection failed: %s", e)
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase"
|
||||
if is_npu
|
||||
else None
|
||||
)
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["npu"] if is_npu else None
|
||||
|
||||
|
||||
def musa_platform_plugin() -> str | None:
|
||||
@@ -152,9 +162,7 @@ def musa_platform_plugin() -> str | None:
|
||||
except Exception as e:
|
||||
logger.debug("MUSA platform is unavailable: %s", e)
|
||||
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.musa.MusaPlatform" if is_musa else None
|
||||
)
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["musa"] if is_musa else None
|
||||
|
||||
|
||||
def xpu_platform_plugin() -> str | None:
|
||||
@@ -175,104 +183,273 @@ def xpu_platform_plugin() -> str | None:
|
||||
except Exception as e:
|
||||
logger.info("Intel XPU platform is unavailable: %s", e)
|
||||
|
||||
return "sglang.multimodal_gen.runtime.platforms.xpu.XpuPlatform" if is_xpu else None
|
||||
return _BUILTIN_PLATFORM_QUALNAMES["xpu"] if is_xpu else None
|
||||
|
||||
|
||||
builtin_platform_plugins = {
|
||||
"cuda": cuda_platform_plugin,
|
||||
"rocm": rocm_platform_plugin,
|
||||
"xpu": xpu_platform_plugin,
|
||||
"mps": mps_platform_plugin,
|
||||
"cpu": cpu_platform_plugin,
|
||||
"xpu": xpu_platform_plugin,
|
||||
"rocm": rocm_platform_plugin,
|
||||
"cuda": cuda_platform_plugin,
|
||||
"npu": npu_platform_plugin,
|
||||
"musa": musa_platform_plugin,
|
||||
"cpu": cpu_platform_plugin,
|
||||
}
|
||||
|
||||
|
||||
def resolve_current_platform_cls_qualname() -> str:
|
||||
forced_platform = os.environ.get("SGLANG_DIFFUSION_PLATFORM_OVERRIDE", "").strip()
|
||||
if forced_platform:
|
||||
forced_map = {
|
||||
"cpu": "sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform",
|
||||
"cuda": "sglang.multimodal_gen.runtime.platforms.cuda.CudaPlatform",
|
||||
"rocm": "sglang.multimodal_gen.runtime.platforms.rocm.RocmPlatform",
|
||||
"mps": "sglang.multimodal_gen.runtime.platforms.mps.MpsPlatform",
|
||||
"npu": "sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase",
|
||||
"musa": "sglang.multimodal_gen.runtime.platforms.musa.MusaPlatform",
|
||||
"xpu": "sglang.multimodal_gen.runtime.platforms.xpu.XpuPlatform",
|
||||
}
|
||||
qualname = forced_map.get(forced_platform.lower())
|
||||
if qualname is None:
|
||||
raise ValueError(
|
||||
f"Unsupported SGLANG_DIFFUSION_PLATFORM_OVERRIDE={forced_platform!r}"
|
||||
)
|
||||
return qualname
|
||||
@dataclass(frozen=True)
|
||||
class _PlatformSelection:
|
||||
qualname: str
|
||||
plugin_name: str | None = None
|
||||
distribution_name: str | None = None
|
||||
|
||||
# TODO(will): if we need to support other platforms, we should consider if
|
||||
# vLLM's plugin architecture is suitable for our needs.
|
||||
@property
|
||||
def is_external(self) -> bool:
|
||||
return self.plugin_name is not None
|
||||
|
||||
# Try MPS first on macOS
|
||||
platform_cls_qualname = mps_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Try Intel XPU
|
||||
platform_cls_qualname = xpu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
def _select_current_platform() -> _PlatformSelection:
|
||||
selected = envs.SGLANG_DIFFUSION_PLATFORM_OVERRIDE.strip()
|
||||
if selected:
|
||||
builtin_name = selected.lower()
|
||||
if builtin_name in BUILTIN_PLATFORM_NAMES:
|
||||
return _PlatformSelection(_BUILTIN_PLATFORM_QUALNAMES[builtin_name])
|
||||
return _resolve_selected_platform(_discover_platform_plugin_entries(), selected)
|
||||
|
||||
# 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 NPU
|
||||
platform_cls_qualname = npu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to MUSA
|
||||
platform_cls_qualname = musa_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
|
||||
platform_selection = _resolve_automatic_platform(
|
||||
_discover_platform_plugin_entries()
|
||||
)
|
||||
if platform_selection is not None:
|
||||
return platform_selection
|
||||
|
||||
for detect in builtin_platform_plugins.values():
|
||||
platform_cls_qualname = detect()
|
||||
if platform_cls_qualname is not None:
|
||||
return _PlatformSelection(platform_cls_qualname)
|
||||
raise RuntimeError("No platform plugin found. Please check your installation.")
|
||||
|
||||
|
||||
def resolve_current_platform_cls_qualname() -> str:
|
||||
"""Resolve the selected class name without mutating singleton state."""
|
||||
return _select_current_platform().qualname
|
||||
|
||||
|
||||
def _discover_platform_plugin_entries() -> tuple[EntryPoint, ...]:
|
||||
entries = tuple(entry_points(group=PLATFORM_PLUGINS_GROUP))
|
||||
if entries:
|
||||
logger.info("Available diffusion platform plugins:")
|
||||
for entry_point in entries:
|
||||
logger.info(" - %s -> %s", entry_point.name, entry_point.value)
|
||||
return entries
|
||||
|
||||
|
||||
def _reject_platform_names(names: Iterable[str], *, reason: str) -> None:
|
||||
# Sorted so the message does not depend on entry-point iteration order.
|
||||
offenders = sorted(names)
|
||||
if offenders:
|
||||
raise RuntimeError(f"{reason}: " + ", ".join(repr(name) for name in offenders))
|
||||
|
||||
|
||||
def _validate_platform_entries(entries: tuple[EntryPoint, ...]) -> None:
|
||||
counts = Counter(entry_point.name for entry_point in entries)
|
||||
_reject_platform_names(
|
||||
(name for name, count in counts.items() if count > 1),
|
||||
reason="Diffusion platform entry-point names must be unique",
|
||||
)
|
||||
_reject_platform_names(
|
||||
(name for name in counts if name.lower() in BUILTIN_PLATFORM_NAMES),
|
||||
reason="Diffusion platform entry points cannot use built-in names",
|
||||
)
|
||||
|
||||
|
||||
def _platform_selection(
|
||||
entry_point: EntryPoint, qualname: object
|
||||
) -> _PlatformSelection | None:
|
||||
if qualname is None:
|
||||
return None
|
||||
# activate() is third-party, so a bad return is named rather than left to
|
||||
# surface as an AttributeError from some later attribute access.
|
||||
selected = qualname.strip() if isinstance(qualname, str) else ""
|
||||
if not selected:
|
||||
raise TypeError(
|
||||
f"Diffusion platform plugin {entry_point.name!r} must return a "
|
||||
"non-empty class qualname or None"
|
||||
)
|
||||
return _PlatformSelection(
|
||||
qualname=selected,
|
||||
plugin_name=entry_point.name,
|
||||
distribution_name=entry_point.dist.name if entry_point.dist else None,
|
||||
)
|
||||
|
||||
|
||||
def _resolve_selected_platform(
|
||||
entries: tuple[EntryPoint, ...],
|
||||
selected: str,
|
||||
) -> _PlatformSelection:
|
||||
matches = tuple(
|
||||
entry_point for entry_point in entries if entry_point.name == selected
|
||||
)
|
||||
if not matches:
|
||||
available = ", ".join(repr(entry_point.name) for entry_point in entries)
|
||||
raise ValueError(
|
||||
f"Unsupported SGLANG_DIFFUSION_PLATFORM_OVERRIDE={selected!r}; "
|
||||
"entry point not found in group "
|
||||
f"{PLATFORM_PLUGINS_GROUP!r} (available: "
|
||||
f"{available or 'none'})."
|
||||
)
|
||||
|
||||
_validate_platform_entries(matches)
|
||||
logger.info(
|
||||
"Selecting platform plugin %s via SGLANG_DIFFUSION_PLATFORM_OVERRIDE",
|
||||
selected,
|
||||
)
|
||||
selection = _platform_selection(matches[0], matches[0].load()())
|
||||
if selection is None:
|
||||
raise RuntimeError(
|
||||
f"Platform plugin {selected!r} is installed but activate() "
|
||||
"returned None (hardware not available on this machine?)."
|
||||
)
|
||||
logger.info("OOT platform plugin activated: %s -> %s", selected, selection.qualname)
|
||||
return selection
|
||||
|
||||
|
||||
def _resolve_automatic_platform(
|
||||
entries: tuple[EntryPoint, ...],
|
||||
) -> _PlatformSelection | None:
|
||||
_validate_platform_entries(entries)
|
||||
activated: list[_PlatformSelection] = []
|
||||
for entry_point in entries:
|
||||
# A raising activate() propagates: silently skipping it would fall back
|
||||
# to a built-in platform and run the whole job on the wrong hardware.
|
||||
selection = _platform_selection(entry_point, entry_point.load()())
|
||||
if selection is not None:
|
||||
activated.append(selection)
|
||||
logger.info(
|
||||
"OOT platform plugin activated: %s -> %s",
|
||||
entry_point.name,
|
||||
selection.qualname,
|
||||
)
|
||||
|
||||
if not activated:
|
||||
return None
|
||||
if len(activated) == 1:
|
||||
return activated[0]
|
||||
names = ", ".join(repr(selection.plugin_name) for selection in activated)
|
||||
raise RuntimeError(
|
||||
f"Multiple platform plugins activated: {names}. "
|
||||
"Set SGLANG_DIFFUSION_PLATFORM_OVERRIDE to select one."
|
||||
)
|
||||
|
||||
|
||||
def _load_platform_class(
|
||||
qualname: str, *, external: bool | None = None
|
||||
) -> type[Platform]:
|
||||
platform_cls = resolve_name(qualname)
|
||||
if not isinstance(platform_cls, type) or not issubclass(platform_cls, Platform):
|
||||
raise TypeError(f"Expected a Platform subclass: {qualname}")
|
||||
if external is None:
|
||||
external = qualname not in _BUILTIN_PLATFORM_QUALNAMES.values()
|
||||
if external and platform_cls._enum is not PlatformEnum.OOT:
|
||||
raise TypeError(
|
||||
f"External diffusion platform {qualname} must set "
|
||||
"_enum = sglang.multimodal_gen.runtime.platforms.PlatformEnum.OOT"
|
||||
)
|
||||
return platform_cls
|
||||
|
||||
|
||||
_current_platform: Platform | None = None
|
||||
_current_platform_selection: _PlatformSelection | None = None
|
||||
_init_trace: str = ""
|
||||
|
||||
_backend_init_done = False
|
||||
_backend_init_error: BaseException | None = None
|
||||
|
||||
current_platform: Platform
|
||||
|
||||
|
||||
def _resolve_current_platform() -> Platform:
|
||||
# Platform plugins import this module to subclass Platform, so resolution
|
||||
# must remain lazy.
|
||||
global _current_platform, _current_platform_selection, _init_trace
|
||||
|
||||
if _current_platform is not None:
|
||||
return _current_platform
|
||||
|
||||
selection = _select_current_platform()
|
||||
platform_cls = _load_platform_class(
|
||||
selection.qualname, external=selection.is_external
|
||||
)
|
||||
platform = platform_cls()
|
||||
if selection.is_external:
|
||||
for attribute in ("device_name", "device_type"):
|
||||
value = getattr(platform, attribute, None)
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise TypeError(
|
||||
f"External diffusion platform {selection.qualname} must "
|
||||
f"define a non-empty {attribute}"
|
||||
)
|
||||
|
||||
# Publish the instance and its provenance together, only after the
|
||||
# complete external contract has passed validation.
|
||||
_current_platform_selection = selection
|
||||
_current_platform = platform
|
||||
_init_trace = "".join(traceback.format_stack())
|
||||
return platform
|
||||
|
||||
|
||||
def get_selected_platform_dist() -> str | None:
|
||||
_resolve_current_platform()
|
||||
assert _current_platform_selection is not None
|
||||
return _current_platform_selection.distribution_name
|
||||
|
||||
|
||||
def initialize_current_platform() -> None:
|
||||
"""Run backend initialization once per process.
|
||||
|
||||
Only worker entry points call this, so a launcher never marks itself
|
||||
initialized and every worker starts from clean module state.
|
||||
|
||||
A failed initialization is terminal for that process: retrying arbitrary
|
||||
backend side effects can duplicate registrations and leave a worker in a
|
||||
state that reflects neither attempt.
|
||||
"""
|
||||
global _backend_init_error, _backend_init_done
|
||||
|
||||
if _backend_init_done:
|
||||
if _backend_init_error is not None:
|
||||
raise RuntimeError(
|
||||
"Diffusion platform backend initialization previously failed: "
|
||||
f"{_backend_init_error}"
|
||||
) from _backend_init_error
|
||||
return
|
||||
|
||||
try:
|
||||
_resolve_current_platform().init_backend()
|
||||
except BaseException as exc:
|
||||
# BaseException too: the finally below marks this attempt done, so an
|
||||
# unrecorded interrupt would let the next call report success.
|
||||
_backend_init_error = exc
|
||||
raise
|
||||
finally:
|
||||
_backend_init_done = True
|
||||
|
||||
|
||||
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`.
|
||||
global _current_platform
|
||||
if _current_platform is None:
|
||||
platform_cls_qualname = resolve_current_platform_cls_qualname()
|
||||
_current_platform = resolve_name(platform_cls_qualname)()
|
||||
global _init_trace
|
||||
_init_trace = "".join(traceback.format_stack())
|
||||
return _current_platform
|
||||
return _resolve_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"]
|
||||
__all__ = [
|
||||
"BUILTIN_PLATFORM_NAMES",
|
||||
"PLATFORM_PLUGINS_GROUP",
|
||||
"Platform",
|
||||
"PlatformEnum",
|
||||
"current_platform",
|
||||
"get_selected_platform_dist",
|
||||
"initialize_current_platform",
|
||||
"_init_trace",
|
||||
]
|
||||
|
||||
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
AttentionImpl,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args.server_args import ServerArgs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -112,15 +113,14 @@ class DeviceCapability(NamedTuple):
|
||||
|
||||
|
||||
class Platform:
|
||||
_enum: PlatformEnum
|
||||
_enum: PlatformEnum = PlatformEnum.UNSPECIFIED
|
||||
device_name: str
|
||||
device_type: str
|
||||
device: torch.device | None = None # Dummy attribute for compatibility
|
||||
|
||||
# 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"
|
||||
dispatch_key: str = ""
|
||||
|
||||
# The torch.compile backend for compiling simple and
|
||||
# standalone functions. The default value is "inductor" to keep
|
||||
@@ -131,6 +131,17 @@ class Platform:
|
||||
|
||||
supported_quantization: list[str] = []
|
||||
|
||||
def init_backend(self) -> None:
|
||||
"""One-time backend initialization, in each worker; raising aborts startup.
|
||||
|
||||
Where out-of-tree platforms register their custom-op forwards.
|
||||
"""
|
||||
pass
|
||||
|
||||
def apply_server_args_defaults(self, server_args: ServerArgs) -> None:
|
||||
"""Apply defaults before argument normalization and validation."""
|
||||
pass
|
||||
|
||||
def get_compile_backend(self, mode: str | None = None) -> str:
|
||||
"""Return the backend used to compile diffusion modules."""
|
||||
return self.simple_compile_backend
|
||||
@@ -139,6 +150,27 @@ class Platform:
|
||||
"""Return backend-specific options for a diffusion module."""
|
||||
return None
|
||||
|
||||
def get_dispatch_key_name(self) -> str:
|
||||
"""Return the behavioral dispatch key used by :class:`CustomOp`.
|
||||
|
||||
This is intentionally separate from ``dispatch_key``, which names a
|
||||
PyTorch dispatcher key such as ``PrivateUse1``. An out-of-tree backend
|
||||
can return an existing key such as ``cuda`` to reuse compatible
|
||||
``forward_cuda`` implementations, or a vendor key backed by registered
|
||||
forwards and ``forward_<key>`` methods.
|
||||
"""
|
||||
return "native"
|
||||
|
||||
def get_torch_library_dispatch_key(self) -> str:
|
||||
"""Return the key used for direct ``torch.library`` registrations."""
|
||||
if self.is_out_of_tree():
|
||||
if not self.dispatch_key:
|
||||
raise NotImplementedError(
|
||||
"Out-of-tree diffusion platforms must define dispatch_key"
|
||||
)
|
||||
return self.dispatch_key
|
||||
return "PrivateUse1" if self.is_npu() else "CUDA"
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_cuda(self) -> bool:
|
||||
return self.is_cuda_static()
|
||||
@@ -178,11 +210,11 @@ class Platform:
|
||||
|
||||
@classmethod
|
||||
def is_cuda_static(cls) -> bool:
|
||||
return getattr(cls, "_enum", None) == PlatformEnum.CUDA
|
||||
return cls._enum == PlatformEnum.CUDA
|
||||
|
||||
@classmethod
|
||||
def is_rocm_static(cls) -> bool:
|
||||
return getattr(cls, "_enum", None) == PlatformEnum.ROCM
|
||||
return cls._enum == PlatformEnum.ROCM
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_hpu(self) -> bool:
|
||||
@@ -190,11 +222,19 @@ class Platform:
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_xpu(self) -> bool:
|
||||
return hasattr(torch, "xpu") and torch.xpu.is_available()
|
||||
return (
|
||||
not self.is_out_of_tree()
|
||||
and hasattr(torch, "xpu")
|
||||
and torch.xpu.is_available()
|
||||
)
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_npu(self) -> bool:
|
||||
return hasattr(torch, "npu") and torch.npu.is_available()
|
||||
return (
|
||||
not self.is_out_of_tree()
|
||||
and hasattr(torch, "npu")
|
||||
and torch.npu.is_available()
|
||||
)
|
||||
|
||||
def is_out_of_tree(self) -> bool:
|
||||
return self._enum == PlatformEnum.OOT
|
||||
@@ -215,7 +255,11 @@ class Platform:
|
||||
@lru_cache(maxsize=1)
|
||||
def is_musa(self):
|
||||
try:
|
||||
return hasattr(torch, "musa") and torch.musa.is_available()
|
||||
return (
|
||||
not self.is_out_of_tree()
|
||||
and hasattr(torch, "musa")
|
||||
and torch.musa.is_available()
|
||||
)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
|
||||
@@ -312,6 +356,10 @@ class Platform:
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_device(self, local_rank: int) -> torch.device:
|
||||
if self.is_out_of_tree():
|
||||
raise NotImplementedError(
|
||||
"Out-of-tree diffusion platforms must implement get_device()"
|
||||
)
|
||||
if self.is_cuda() or self.is_rocm():
|
||||
return torch.device("cuda", local_rank)
|
||||
elif self.is_npu():
|
||||
@@ -344,6 +392,17 @@ class Platform:
|
||||
"No Accelerators(AMD/NV/MTT GPU, AMD MI instinct accelerators) available"
|
||||
)
|
||||
|
||||
def supports_distributed_device_id(self) -> bool:
|
||||
"""Whether torch.distributed accepts this platform's device ID."""
|
||||
return not (
|
||||
self.is_out_of_tree()
|
||||
or self.is_mps()
|
||||
or self.is_musa()
|
||||
or self.is_npu()
|
||||
or self.is_cpu()
|
||||
or self.is_xpu()
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
"""
|
||||
@@ -421,11 +480,23 @@ class Platform:
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
"""
|
||||
Get device specific communicator class for distributed communication.
|
||||
"""
|
||||
"""Return the platform's default device communicator class."""
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase" # noqa
|
||||
|
||||
@classmethod
|
||||
def get_all_to_all_communicator_cls(cls) -> str:
|
||||
"""Return the communicator used by ``all_to_all_4D``."""
|
||||
qualname = cls.get_device_communicator_cls()
|
||||
if (
|
||||
cls._enum is PlatformEnum.OOT
|
||||
and qualname == Platform.get_device_communicator_cls()
|
||||
):
|
||||
raise NotImplementedError(
|
||||
"Out-of-tree diffusion platforms must implement "
|
||||
"get_all_to_all_communicator_cls()"
|
||||
)
|
||||
return qualname
|
||||
|
||||
@classmethod
|
||||
def get_cpu_architecture(cls) -> CpuArchEnum:
|
||||
"""Get the CPU architecture of the current platform."""
|
||||
|
||||
@@ -121,6 +121,13 @@ class MpsPlatform(Platform):
|
||||
# Use base communicator for MPS
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase"
|
||||
|
||||
@classmethod
|
||||
def get_all_to_all_communicator_cls(cls) -> str:
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.distributed.device_communicators."
|
||||
"cpu_communicator.CpuCommunicator"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def seed_everything(cls, seed: int | None = None) -> None:
|
||||
"""Set the seed for MPS device."""
|
||||
|
||||
@@ -197,6 +197,13 @@ class NPUPlatformBase(Platform):
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
|
||||
|
||||
@classmethod
|
||||
def get_all_to_all_communicator_cls(cls) -> str:
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.distributed.device_communicators."
|
||||
"cpu_communicator.CpuCommunicator"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def enable_dit_layerwise_offload_by_default(cls) -> bool:
|
||||
"""Whether automatic DiT layerwise offload is enabled on this platform."""
|
||||
|
||||
@@ -0,0 +1,274 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from collections import Counter, defaultdict
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from importlib.metadata import EntryPoint, entry_points
|
||||
from typing import Any
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import (
|
||||
PLATFORM_PLUGINS_GROUP,
|
||||
get_selected_platform_dist,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.plugins.hook_registry import HookRegistry as _HookRegistry
|
||||
from sglang.srt.plugins.hook_registry import (
|
||||
HookSource,
|
||||
HookType,
|
||||
_current_plugin_source,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GENERAL_PLUGINS_GROUP = "sglang.multimodal_gen.plugins"
|
||||
|
||||
|
||||
class _OnceState(Enum):
|
||||
NOT_STARTED = "not_started"
|
||||
RUNNING = "running"
|
||||
COMPLETE = "complete"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class _Once:
|
||||
"""A process-local initialization gate.
|
||||
|
||||
The lock spans the action, so RUNNING means the caller is the thread already
|
||||
inside: a callback that hands activation to another thread and joins it
|
||||
deadlocks rather than racing. A failure is terminal.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
self.state = _OnceState.NOT_STARTED
|
||||
self.error: BaseException | None = None
|
||||
self._lock = threading.RLock()
|
||||
|
||||
def run(self, action: Callable[[], None]) -> bool:
|
||||
"""Run *action* once; return False only for a re-entrant call."""
|
||||
with self._lock:
|
||||
if self.state is _OnceState.COMPLETE:
|
||||
return True
|
||||
if self.state is _OnceState.FAILED:
|
||||
raise RuntimeError(
|
||||
f"{self.name} previously failed: {self.error}"
|
||||
) from self.error
|
||||
if self.state is _OnceState.RUNNING:
|
||||
return False
|
||||
|
||||
self.state = _OnceState.RUNNING
|
||||
try:
|
||||
action()
|
||||
except BaseException as exc:
|
||||
self.error = exc
|
||||
self.state = _OnceState.FAILED
|
||||
raise
|
||||
self.state = _OnceState.COMPLETE
|
||||
return True
|
||||
|
||||
def reset(self) -> None:
|
||||
with self._lock:
|
||||
self.state = _OnceState.NOT_STARTED
|
||||
self.error = None
|
||||
|
||||
|
||||
_plugin_registration = _Once("Diffusion plugin registration")
|
||||
_hook_application = _Once("Diffusion hook application")
|
||||
_required_dist: str | None = None
|
||||
|
||||
|
||||
class HookRegistry(_HookRegistry):
|
||||
# Rebound so diffusion hooks do not land in SRT's registry.
|
||||
_hooks = defaultdict(list)
|
||||
_patched = set()
|
||||
|
||||
|
||||
def plugin_hook(target: str, type: HookType = HookType.AFTER) -> Callable:
|
||||
def decorator(hook: Callable) -> Callable:
|
||||
HookRegistry.register(target, hook, type)
|
||||
return hook
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def _discard_hooks_from_source(source: HookSource) -> None:
|
||||
"""Discard hooks registered by a plugin that failed to load or execute."""
|
||||
for target in tuple(HookRegistry._hooks):
|
||||
remaining = [hook for hook in HookRegistry._hooks[target] if hook[2] != source]
|
||||
if remaining:
|
||||
HookRegistry._hooks[target] = remaining
|
||||
else:
|
||||
del HookRegistry._hooks[target]
|
||||
|
||||
|
||||
def _get_excluded_dists() -> set[str]:
|
||||
selected_dist = get_selected_platform_dist()
|
||||
platform_dists = {
|
||||
entry_point.dist.name
|
||||
for entry_point in entry_points(group=PLATFORM_PLUGINS_GROUP)
|
||||
if entry_point.dist
|
||||
}
|
||||
if selected_dist is None:
|
||||
return platform_dists
|
||||
return platform_dists - {selected_dist}
|
||||
|
||||
|
||||
def _discover() -> dict[str, tuple[Callable[[], Any], str | None]]:
|
||||
allowed: set[str] | None = None
|
||||
allowed_str = envs.SGLANG_PLUGINS.get()
|
||||
if allowed_str:
|
||||
allowed = {name.strip() for name in allowed_str.split(",") if name.strip()}
|
||||
|
||||
discovered = tuple(entry_points(group=GENERAL_PLUGINS_GROUP))
|
||||
if not discovered:
|
||||
logger.debug("No diffusion plugins found for group %s.", GENERAL_PLUGINS_GROUP)
|
||||
return {}
|
||||
|
||||
excluded_dists = _get_excluded_dists()
|
||||
required_dist = get_selected_platform_dist()
|
||||
candidates: list[EntryPoint] = []
|
||||
for entry_point in discovered:
|
||||
dist_name = entry_point.dist.name if entry_point.dist else None
|
||||
if allowed is not None and entry_point.name not in allowed:
|
||||
logger.info(
|
||||
"Skipping diffusion plugin %s (not in SGLANG_PLUGINS)",
|
||||
entry_point.name,
|
||||
)
|
||||
continue
|
||||
if dist_name in excluded_dists:
|
||||
logger.info(
|
||||
"Skipping diffusion plugin %s (dist %s is not the selected platform)",
|
||||
entry_point.name,
|
||||
dist_name,
|
||||
)
|
||||
continue
|
||||
candidates.append(entry_point)
|
||||
|
||||
counts = Counter(entry_point.name for entry_point in candidates)
|
||||
duplicates = sorted(name for name, count in counts.items() if count > 1)
|
||||
if duplicates:
|
||||
raise RuntimeError(
|
||||
"Diffusion plugin entry-point names must be unique: "
|
||||
+ ", ".join(repr(name) for name in duplicates)
|
||||
)
|
||||
|
||||
plugins: dict[str, tuple[Callable[[], Any], str | None]] = {}
|
||||
for entry_point in candidates:
|
||||
dist_name = entry_point.dist.name if entry_point.dist else None
|
||||
source = HookSource(plugin_name=entry_point.name, dist_name=dist_name)
|
||||
token = _current_plugin_source.set(source)
|
||||
try:
|
||||
callback = entry_point.load()
|
||||
if not callable(callback):
|
||||
raise TypeError(
|
||||
f"Diffusion plugin {entry_point.name!r} must resolve to a callable"
|
||||
)
|
||||
plugins[entry_point.name] = (callback, dist_name)
|
||||
logger.info("Loaded diffusion plugin %s", entry_point.name)
|
||||
except Exception:
|
||||
_discard_hooks_from_source(source)
|
||||
if required_dist and dist_name == required_dist:
|
||||
raise
|
||||
logger.exception("Failed to load diffusion plugin %s", entry_point.name)
|
||||
finally:
|
||||
_current_plugin_source.reset(token)
|
||||
|
||||
return plugins
|
||||
|
||||
|
||||
def _require_hooks_applied(required_dist: str) -> None:
|
||||
unapplied = sorted(
|
||||
target
|
||||
for target, hooks in HookRegistry._hooks.items()
|
||||
if target not in HookRegistry._patched
|
||||
and any(source and source.dist_name == required_dist for _, _, source in hooks)
|
||||
)
|
||||
if unapplied:
|
||||
raise RuntimeError(
|
||||
f"Selected platform package {required_dist!r} could not apply hooks on: "
|
||||
+ ", ".join(unapplied)
|
||||
)
|
||||
|
||||
|
||||
def _register_plugins_once() -> str | None:
|
||||
plugins = _discover()
|
||||
# The selected platform's own plugins carry its hardware contract, so their
|
||||
# failures abort startup; third-party ones stay best-effort.
|
||||
required_dist = get_selected_platform_dist() if plugins else None
|
||||
|
||||
for name, (func, dist_name) in plugins.items():
|
||||
source = HookSource(plugin_name=name, dist_name=dist_name)
|
||||
token = _current_plugin_source.set(source)
|
||||
try:
|
||||
func()
|
||||
logger.info("Executed diffusion plugin: %s", name)
|
||||
except Exception:
|
||||
_discard_hooks_from_source(source)
|
||||
if required_dist and dist_name == required_dist:
|
||||
raise
|
||||
logger.exception("Failed to execute diffusion plugin: %s", name)
|
||||
finally:
|
||||
_current_plugin_source.reset(token)
|
||||
|
||||
return required_dist
|
||||
|
||||
|
||||
def load_plugins() -> None:
|
||||
"""Discover and execute diffusion plugin callbacks once per process.
|
||||
|
||||
This phase only registers hooks. It deliberately does not resolve hook
|
||||
targets: resolving a dotted target imports its module, which makes a
|
||||
seemingly harmless plugin-discovery call capable of importing the entire
|
||||
worker runtime.
|
||||
|
||||
Re-entrant calls from a plugin callback return immediately; a caller on
|
||||
another thread waits for the in-flight registration. A failed load is
|
||||
terminal for the process because arbitrary callback side effects cannot be
|
||||
rolled back safely.
|
||||
"""
|
||||
_ensure_plugins_loaded()
|
||||
|
||||
|
||||
def _ensure_plugins_loaded() -> bool:
|
||||
def register() -> None:
|
||||
global _required_dist
|
||||
_required_dist = _register_plugins_once()
|
||||
|
||||
return _plugin_registration.run(register)
|
||||
|
||||
|
||||
def apply_plugin_hooks() -> None:
|
||||
"""Apply registered hooks once, at an explicit runtime-safe boundary.
|
||||
|
||||
Hook target resolution is allowed to import target modules. Callers that
|
||||
require import ordering, notably spawned accelerator workers, must finish
|
||||
platform initialization before entering this phase.
|
||||
"""
|
||||
if not _ensure_plugins_loaded():
|
||||
# The outer activation applies the complete registry after registration.
|
||||
return
|
||||
|
||||
def apply() -> None:
|
||||
HookRegistry.apply_hooks()
|
||||
if _required_dist:
|
||||
_require_hooks_applied(_required_dist)
|
||||
|
||||
_hook_application.run(apply)
|
||||
|
||||
|
||||
def _reset_lifecycle_for_tests() -> None:
|
||||
global _required_dist
|
||||
_plugin_registration.reset()
|
||||
_hook_application.reset()
|
||||
_required_dist = None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"HookRegistry",
|
||||
"HookType",
|
||||
"apply_plugin_hooks",
|
||||
"load_plugins",
|
||||
"plugin_hook",
|
||||
]
|
||||
@@ -191,3 +191,10 @@ class XpuPlatform(Platform):
|
||||
"""Get device communicator class for Intel XPU distributed communication."""
|
||||
# Use base communicator for now; can be updated to use oneCCL-based communicator
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator.DeviceCommunicatorBase"
|
||||
|
||||
@classmethod
|
||||
def get_all_to_all_communicator_cls(cls) -> str:
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.distributed.device_communicators."
|
||||
"cpu_communicator.CpuCommunicator"
|
||||
)
|
||||
|
||||
@@ -1897,6 +1897,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
self
|
||||
)
|
||||
|
||||
current_platform.apply_server_args_defaults(self)
|
||||
# configure logger before use
|
||||
configure_logger(server_args=self)
|
||||
|
||||
|
||||
@@ -20,7 +20,10 @@ def maybe_enable_inductor_compute_comm_overlap() -> None:
|
||||
|
||||
|
||||
def build_torch_compile_kwargs(
|
||||
*, mode: str | None, module: nn.Module | None = None
|
||||
*,
|
||||
mode: str | None,
|
||||
module: nn.Module | None = None,
|
||||
enable_inductor_compute_comm_overlap: bool = False,
|
||||
) -> dict[str, object]:
|
||||
compile_kwargs: dict[str, object] = {"fullgraph": False, "dynamic": None}
|
||||
if current_platform.is_out_of_tree():
|
||||
@@ -43,6 +46,11 @@ def build_torch_compile_kwargs(
|
||||
compile_kwargs["dynamic"] = False
|
||||
elif mode is not None:
|
||||
compile_kwargs["mode"] = mode
|
||||
if (
|
||||
enable_inductor_compute_comm_overlap
|
||||
and compile_kwargs.get("backend", "inductor") == "inductor"
|
||||
):
|
||||
maybe_enable_inductor_compute_comm_overlap()
|
||||
return compile_kwargs
|
||||
|
||||
|
||||
@@ -61,6 +69,24 @@ def resolve_torch_compile_mode(
|
||||
return default
|
||||
|
||||
|
||||
def resolve_torch_compile_kwargs(
|
||||
*env_names: str,
|
||||
config: object | None = None,
|
||||
default: str,
|
||||
module: nn.Module | None = None,
|
||||
enable_inductor_compute_comm_overlap: bool = False,
|
||||
) -> tuple[dict[str, object], str | None]:
|
||||
mode = None
|
||||
if not current_platform.is_npu():
|
||||
mode = resolve_torch_compile_mode(*env_names, config=config, default=default)
|
||||
compile_kwargs = build_torch_compile_kwargs(
|
||||
mode=mode,
|
||||
module=module,
|
||||
enable_inductor_compute_comm_overlap=enable_inductor_compute_comm_overlap,
|
||||
)
|
||||
return compile_kwargs, mode
|
||||
|
||||
|
||||
def compile_matching_submodules(
|
||||
module: nn.Module,
|
||||
*,
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
"""The supported offline-script layout: only the facade at module scope.
|
||||
|
||||
``spawn`` re-executes this in every child, so binding ``DiffGenerator`` here
|
||||
must not import the diffusion runtime.
|
||||
"""
|
||||
|
||||
from offline_script_runner import main
|
||||
|
||||
from sglang.multimodal_gen import DiffGenerator # noqa: F401
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
"""Shared body of the offline-script fixtures in ``test_worker_bootstrap``.
|
||||
|
||||
Held apart so the scripts differ only in the module-scope import each is named
|
||||
for; nothing here may import what the child's bootstrap has to precede.
|
||||
"""
|
||||
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import sys
|
||||
|
||||
from sglang.multimodal_gen.runtime.managers import worker_bootstrap
|
||||
|
||||
_CHILD_REPLY_TIMEOUT_S = 120
|
||||
_CHILD_JOIN_TIMEOUT_S = 10
|
||||
|
||||
|
||||
def run_offline_script(result_path: str) -> None:
|
||||
"""Spawn one scheduler child and write back what it observed."""
|
||||
# Inside the guarded call on purpose: this is the import whose absence from
|
||||
# the child's module scope the test is measuring.
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
spec = worker_bootstrap.SchedulerProcessSpec(
|
||||
local_rank=0,
|
||||
rank=0,
|
||||
server_args=worker_bootstrap.ServerArgsPayload.capture(
|
||||
ServerArgs.__new__(ServerArgs)
|
||||
),
|
||||
pipe_writer=writer,
|
||||
)
|
||||
|
||||
process = mp.get_context("spawn").Process(
|
||||
target=worker_bootstrap.bootstrap_scheduler_process,
|
||||
args=(spec,),
|
||||
)
|
||||
process.start()
|
||||
writer.close()
|
||||
|
||||
observed = None
|
||||
if reader.poll(_CHILD_REPLY_TIMEOUT_S):
|
||||
try:
|
||||
observed = reader.recv()
|
||||
except EOFError:
|
||||
pass
|
||||
process.join(_CHILD_JOIN_TIMEOUT_S)
|
||||
if process.is_alive():
|
||||
process.kill()
|
||||
process.join(_CHILD_JOIN_TIMEOUT_S)
|
||||
|
||||
with open(result_path, "w") as result_file:
|
||||
json.dump({"observed": observed, "exitcode": process.exitcode}, result_file)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
run_offline_script(sys.argv[1])
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
"""An offline script that reaches past the facade at module scope.
|
||||
|
||||
The child therefore imports the runtime before it can initialize its platform.
|
||||
Nothing in-tree can reorder that; the child is expected to say so.
|
||||
"""
|
||||
|
||||
from offline_script_runner import main
|
||||
|
||||
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import ( # noqa: F401
|
||||
DiffGenerator,
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,86 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
"""Fake out-of-tree platform package used by ``test_worker_bootstrap``.
|
||||
|
||||
The test copies this file into a temporary directory alongside a generated
|
||||
``sgl_fake_plugin-0.1.dist-info`` and puts that directory on ``sys.path``, so a
|
||||
spawned child discovers it through real entry-point metadata.
|
||||
|
||||
It records which diffusion modules were already imported at each bootstrap
|
||||
boundary, which is why it must stay free of diffusion runtime imports beyond
|
||||
``Platform``: importing an observed module here would corrupt the measurement.
|
||||
"""
|
||||
|
||||
import sys
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import Platform, PlatformEnum
|
||||
|
||||
WORKER_MODULE = "sglang.multimodal_gen.runtime.managers.gpu_worker"
|
||||
GENERATOR_MODULE = "sglang.multimodal_gen.runtime.entrypoints.diffusion_generator"
|
||||
SERVER_ARGS_MODULE = "sglang.multimodal_gen.runtime.server_args.server_args"
|
||||
|
||||
worker_imported_when_plugin_ran = None
|
||||
generator_imported_when_plugin_ran = None
|
||||
server_args_imported_when_plugin_ran = None
|
||||
backend_initialized_when_plugin_ran = None
|
||||
worker_imported_when_backend_initialized = None
|
||||
server_args_imported_when_backend_initialized = None
|
||||
backend_initialized = False
|
||||
|
||||
|
||||
class FakePlatform(Platform):
|
||||
_enum = PlatformEnum.OOT
|
||||
device_name = "fake"
|
||||
device_type = "fake"
|
||||
dispatch_key = "PrivateUse1"
|
||||
|
||||
def init_backend(self):
|
||||
global backend_initialized, server_args_imported_when_backend_initialized
|
||||
global worker_imported_when_backend_initialized
|
||||
worker_imported_when_backend_initialized = WORKER_MODULE in sys.modules
|
||||
server_args_imported_when_backend_initialized = (
|
||||
SERVER_ARGS_MODULE in sys.modules
|
||||
)
|
||||
backend_initialized = True
|
||||
|
||||
|
||||
def activate():
|
||||
return "sgl_fake_plugin.FakePlatform"
|
||||
|
||||
|
||||
def replacement(pipe_writer, *args, **kwargs):
|
||||
pipe_writer.send(
|
||||
{
|
||||
"override_ran": True,
|
||||
"worker_imported_when_plugin_ran": worker_imported_when_plugin_ran,
|
||||
"generator_imported_when_plugin_ran": generator_imported_when_plugin_ran,
|
||||
"server_args_imported_when_plugin_ran": (
|
||||
server_args_imported_when_plugin_ran
|
||||
),
|
||||
"backend_initialized_when_plugin_ran": backend_initialized_when_plugin_ran,
|
||||
"worker_imported_when_backend_initialized": (
|
||||
worker_imported_when_backend_initialized
|
||||
),
|
||||
"server_args_imported_when_backend_initialized": (
|
||||
server_args_imported_when_backend_initialized
|
||||
),
|
||||
"backend_initialized": backend_initialized,
|
||||
}
|
||||
)
|
||||
pipe_writer.close()
|
||||
|
||||
|
||||
def register():
|
||||
global backend_initialized_when_plugin_ran
|
||||
global generator_imported_when_plugin_ran, server_args_imported_when_plugin_ran
|
||||
global worker_imported_when_plugin_ran
|
||||
backend_initialized_when_plugin_ran = backend_initialized
|
||||
worker_imported_when_plugin_ran = WORKER_MODULE in sys.modules
|
||||
generator_imported_when_plugin_ran = GENERATOR_MODULE in sys.modules
|
||||
server_args_imported_when_plugin_ran = SERVER_ARGS_MODULE in sys.modules
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms.plugins import HookRegistry, HookType
|
||||
|
||||
HookRegistry.register(
|
||||
WORKER_MODULE + ".run_scheduler_process", replacement, HookType.REPLACE
|
||||
)
|
||||
@@ -0,0 +1,505 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import threading
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import plugins
|
||||
from sglang.srt.plugins.hook_registry import HookRegistry as SrtHookRegistry
|
||||
from sglang.srt.plugins.hook_registry import (
|
||||
HookSource,
|
||||
HookType,
|
||||
)
|
||||
|
||||
_THREAD_TIMEOUT_S = 10
|
||||
|
||||
|
||||
class _Caller(threading.Thread):
|
||||
"""Runs one activation call on its own thread, keeping what it raised."""
|
||||
|
||||
def __init__(self, call):
|
||||
super().__init__(daemon=True)
|
||||
self._call = call
|
||||
self.error = None
|
||||
|
||||
def run(self):
|
||||
try:
|
||||
self._call()
|
||||
except BaseException as exc:
|
||||
self.error = exc
|
||||
|
||||
|
||||
class _GateProbe:
|
||||
"""A gate lock that reports when a thread genuinely has to wait on it.
|
||||
|
||||
The non-blocking attempt fails only for a non-owner, so the arrival is
|
||||
observable instead of guessed at with a sleep that can silently miss.
|
||||
"""
|
||||
|
||||
def __init__(self, lock):
|
||||
self._lock = lock
|
||||
self.blocked = threading.Event()
|
||||
|
||||
def __enter__(self):
|
||||
if not self._lock.acquire(blocking=False):
|
||||
self.blocked.set()
|
||||
self._lock.acquire()
|
||||
|
||||
def __exit__(self, *exc_info):
|
||||
self._lock.release()
|
||||
|
||||
|
||||
def _entry_point(name, distribution):
|
||||
entry_point = MagicMock(name=f"entry_point_{name}")
|
||||
entry_point.name = name
|
||||
entry_point.value = f"test_plugin:{name}"
|
||||
entry_point.dist = SimpleNamespace(name=distribution)
|
||||
entry_point.load.return_value = MagicMock(name=f"plugin_{name}")
|
||||
return entry_point
|
||||
|
||||
|
||||
class _ThreadedTestCase(unittest.TestCase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self._callers = []
|
||||
# A failed assertion can leave a caller unjoined; it must not run on
|
||||
# into the next test.
|
||||
self.addCleanup(self._join_started_callers)
|
||||
|
||||
def _join_started_callers(self):
|
||||
for caller in self._callers:
|
||||
caller.join(_THREAD_TIMEOUT_S)
|
||||
|
||||
def _finish(self):
|
||||
for caller in self._callers:
|
||||
caller.join(_THREAD_TIMEOUT_S)
|
||||
self.assertFalse(caller.is_alive(), "activation thread never finished")
|
||||
|
||||
def _start_caller(self, call):
|
||||
caller = _Caller(call)
|
||||
self._callers.append(caller)
|
||||
caller.start()
|
||||
return caller
|
||||
|
||||
def _start_second_caller(self, call, probe):
|
||||
"""Start *call* elsewhere and wait until it is provably blocked at the gate."""
|
||||
caller = self._start_caller(call)
|
||||
self.assertTrue(
|
||||
probe.blocked.wait(_THREAD_TIMEOUT_S),
|
||||
"second thread never reached the gate",
|
||||
)
|
||||
return caller
|
||||
|
||||
|
||||
class TestOnceGate(_ThreadedTestCase):
|
||||
"""Gate semantics on a fresh instance, clear of the module-global phases."""
|
||||
|
||||
def test_a_waiting_thread_inherits_the_failure(self):
|
||||
"""A failure must reach a caller that arrived while the phase ran."""
|
||||
once = plugins._Once("test gate")
|
||||
probe = _GateProbe(once._lock)
|
||||
once._lock = probe
|
||||
inside = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def boom():
|
||||
inside.set()
|
||||
self.assertTrue(
|
||||
release.wait(_THREAD_TIMEOUT_S), "release was never signalled"
|
||||
)
|
||||
raise RuntimeError("vendor plugin exploded")
|
||||
|
||||
first = self._start_caller(lambda: once.run(boom))
|
||||
try:
|
||||
self.assertTrue(inside.wait(_THREAD_TIMEOUT_S))
|
||||
second = self._start_second_caller(lambda: once.run(lambda: None), probe)
|
||||
finally:
|
||||
release.set()
|
||||
self._finish()
|
||||
|
||||
self.assertIsInstance(first.error, RuntimeError)
|
||||
self.assertIn("exploded", str(first.error))
|
||||
self.assertIsInstance(second.error, RuntimeError)
|
||||
self.assertIn("previously failed", str(second.error))
|
||||
|
||||
|
||||
class TestDiffusionPluginBarrier(_ThreadedTestCase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
state = (
|
||||
plugins._plugin_registration.state,
|
||||
plugins._plugin_registration.error,
|
||||
plugins._hook_application.state,
|
||||
plugins._hook_application.error,
|
||||
plugins._required_dist,
|
||||
)
|
||||
self.addCleanup(self._restore_lifecycle, state)
|
||||
plugins._reset_lifecycle_for_tests()
|
||||
self._callers = []
|
||||
# A failed assertion can leave a caller unjoined; it must not run on
|
||||
# into the next test.
|
||||
self.addCleanup(self._join_started_callers)
|
||||
|
||||
@staticmethod
|
||||
def _restore_lifecycle(state):
|
||||
(
|
||||
plugins._plugin_registration.state,
|
||||
plugins._plugin_registration.error,
|
||||
plugins._hook_application.state,
|
||||
plugins._hook_application.error,
|
||||
plugins._required_dist,
|
||||
) = state
|
||||
|
||||
def test_body_runs_once_across_repeated_calls(self):
|
||||
with patch.object(plugins, "_register_plugins_once") as load_once:
|
||||
plugins.load_plugins()
|
||||
plugins.load_plugins()
|
||||
|
||||
load_once.assert_called_once_with()
|
||||
|
||||
def test_failure_is_terminal_instead_of_replaying_partial_side_effects(self):
|
||||
with patch.object(
|
||||
plugins,
|
||||
"_register_plugins_once",
|
||||
side_effect=RuntimeError("plugin init exploded"),
|
||||
) as load_once:
|
||||
with self.assertRaisesRegex(RuntimeError, "exploded"):
|
||||
plugins.load_plugins()
|
||||
with self.assertRaisesRegex(RuntimeError, "previously failed"):
|
||||
plugins.load_plugins()
|
||||
|
||||
load_once.assert_called_once_with()
|
||||
|
||||
def test_reentrant_call_returns_instead_of_recursing(self):
|
||||
calls = []
|
||||
|
||||
def reentrant():
|
||||
calls.append(1)
|
||||
plugins.load_plugins()
|
||||
|
||||
with patch.object(plugins, "_register_plugins_once", reentrant):
|
||||
plugins.load_plugins()
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
|
||||
def test_hook_application_is_a_separate_once_only_phase(self):
|
||||
with (
|
||||
patch.object(plugins, "_register_plugins_once", return_value="vendor-pkg"),
|
||||
patch.object(plugins.HookRegistry, "apply_hooks") as apply_hooks,
|
||||
patch.object(plugins, "_require_hooks_applied") as require_hooks,
|
||||
):
|
||||
plugins.load_plugins()
|
||||
apply_hooks.assert_not_called()
|
||||
|
||||
plugins.apply_plugin_hooks()
|
||||
plugins.apply_plugin_hooks()
|
||||
|
||||
apply_hooks.assert_called_once_with()
|
||||
require_hooks.assert_called_once_with("vendor-pkg")
|
||||
|
||||
def test_hook_application_failure_is_terminal(self):
|
||||
with (
|
||||
patch.object(plugins, "_register_plugins_once", return_value=None),
|
||||
patch.object(
|
||||
plugins.HookRegistry,
|
||||
"apply_hooks",
|
||||
side_effect=RuntimeError("hook application exploded"),
|
||||
) as apply_hooks,
|
||||
):
|
||||
with self.assertRaisesRegex(RuntimeError, "application exploded"):
|
||||
plugins.apply_plugin_hooks()
|
||||
with self.assertRaisesRegex(RuntimeError, "previously failed"):
|
||||
plugins.apply_plugin_hooks()
|
||||
|
||||
apply_hooks.assert_called_once_with()
|
||||
|
||||
def test_hooks_from_a_reentrant_callback_are_applied_by_the_outer_call(self):
|
||||
def reentrant():
|
||||
plugins.apply_plugin_hooks()
|
||||
|
||||
with (
|
||||
patch.object(plugins, "_register_plugins_once", reentrant),
|
||||
patch.object(plugins.HookRegistry, "apply_hooks") as apply_hooks,
|
||||
):
|
||||
plugins.apply_plugin_hooks()
|
||||
|
||||
apply_hooks.assert_called_once_with()
|
||||
|
||||
def _probe_registration_gate(self):
|
||||
probe = _GateProbe(plugins._plugin_registration._lock)
|
||||
patcher = patch.object(plugins._plugin_registration, "_lock", probe)
|
||||
patcher.start()
|
||||
self.addCleanup(patcher.stop)
|
||||
return probe
|
||||
|
||||
def test_a_second_thread_does_not_skip_hook_application(self):
|
||||
"""A caller that arrives during registration must not return before the
|
||||
registry is applied."""
|
||||
inside = threading.Event()
|
||||
release = threading.Event()
|
||||
applied_on_return = []
|
||||
probe = self._probe_registration_gate()
|
||||
|
||||
def slow_register():
|
||||
inside.set()
|
||||
self.assertTrue(
|
||||
release.wait(_THREAD_TIMEOUT_S), "release was never signalled"
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(plugins, "_register_plugins_once", slow_register),
|
||||
patch.object(plugins.HookRegistry, "apply_hooks") as apply_hooks,
|
||||
):
|
||||
|
||||
def apply_and_report():
|
||||
plugins.apply_plugin_hooks()
|
||||
applied_on_return.append(apply_hooks.call_count)
|
||||
|
||||
first = self._start_caller(plugins.load_plugins)
|
||||
try:
|
||||
self.assertTrue(inside.wait(_THREAD_TIMEOUT_S))
|
||||
second = self._start_second_caller(apply_and_report, probe)
|
||||
finally:
|
||||
release.set()
|
||||
self._finish()
|
||||
|
||||
apply_hooks.assert_called_once_with()
|
||||
|
||||
self.assertIsNone(first.error)
|
||||
self.assertIsNone(second.error)
|
||||
self.assertEqual(
|
||||
applied_on_return,
|
||||
[1],
|
||||
"second thread returned before the registry was applied",
|
||||
)
|
||||
|
||||
|
||||
class TestDiffusionPlugins(unittest.TestCase):
|
||||
def setUp(self):
|
||||
# _discover() reads this live, so an allowlist set in the environment
|
||||
# would filter the mocked entry points out from under these tests.
|
||||
patcher = patch.dict(os.environ, {"SGLANG_PLUGINS": ""})
|
||||
patcher.start()
|
||||
self.addCleanup(patcher.stop)
|
||||
|
||||
def test_load_executes_callbacks_without_resolving_hook_targets(self):
|
||||
register = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
plugins, "_discover", return_value={"test": (register, "test-package")}
|
||||
),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="test-package"
|
||||
),
|
||||
patch.object(plugins.HookRegistry, "apply_hooks") as apply_hooks,
|
||||
):
|
||||
required_dist = plugins._register_plugins_once()
|
||||
|
||||
register.assert_called_once_with()
|
||||
apply_hooks.assert_not_called()
|
||||
self.assertEqual(required_dist, "test-package")
|
||||
|
||||
def test_a_failing_callback_does_not_stop_the_others(self):
|
||||
healthy = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
plugins,
|
||||
"_discover",
|
||||
return_value={
|
||||
"broken": (MagicMock(side_effect=RuntimeError("boom")), "a"),
|
||||
"healthy": (healthy, "b"),
|
||||
},
|
||||
),
|
||||
# Unpatched, this runs real platform detection.
|
||||
patch.object(plugins, "get_selected_platform_dist", return_value=None),
|
||||
):
|
||||
plugins._register_plugins_once()
|
||||
|
||||
healthy.assert_called_once_with()
|
||||
|
||||
def test_hooks_registered_while_importing_a_plugin_keep_their_source(self):
|
||||
target = "test_diffusion_plugins.import_time_target"
|
||||
hook = MagicMock()
|
||||
entry_point = _entry_point("vendor", "vendor-pkg")
|
||||
|
||||
def load():
|
||||
plugins.plugin_hook(target)(hook)
|
||||
return MagicMock()
|
||||
|
||||
entry_point.load.side_effect = load
|
||||
self.addCleanup(plugins.HookRegistry._hooks.pop, target, None)
|
||||
|
||||
with (
|
||||
patch.object(plugins, "entry_points", return_value=[entry_point]),
|
||||
patch.object(plugins, "_get_excluded_dists", return_value=set()),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
):
|
||||
plugins._discover()
|
||||
|
||||
self.assertEqual(
|
||||
plugins.HookRegistry._hooks[target][0][2],
|
||||
HookSource("vendor", "vendor-pkg"),
|
||||
)
|
||||
|
||||
def test_failing_optional_callback_discards_its_registered_hooks(self):
|
||||
target = "test_diffusion_plugins.partial_callback_target"
|
||||
|
||||
def register_then_fail():
|
||||
plugins.plugin_hook(target)(lambda result: result)
|
||||
raise RuntimeError("boom")
|
||||
|
||||
self.addCleanup(plugins.HookRegistry._hooks.pop, target, None)
|
||||
with (
|
||||
patch.object(
|
||||
plugins,
|
||||
"_discover",
|
||||
return_value={"broken": (register_then_fail, "optional-pkg")},
|
||||
),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
):
|
||||
plugins._register_plugins_once()
|
||||
|
||||
self.assertNotIn(target, plugins.HookRegistry._hooks)
|
||||
|
||||
def test_failing_optional_import_discards_its_registered_hooks(self):
|
||||
target = "test_diffusion_plugins.partial_import_target"
|
||||
entry_point = _entry_point("broken", "optional-pkg")
|
||||
|
||||
def load_then_fail():
|
||||
plugins.plugin_hook(target)(lambda result: result)
|
||||
raise RuntimeError("boom")
|
||||
|
||||
entry_point.load.side_effect = load_then_fail
|
||||
self.addCleanup(plugins.HookRegistry._hooks.pop, target, None)
|
||||
with (
|
||||
patch.object(plugins, "entry_points", return_value=[entry_point]),
|
||||
patch.object(plugins, "_get_excluded_dists", return_value=set()),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
):
|
||||
self.assertEqual(plugins._discover(), {})
|
||||
|
||||
self.assertNotIn(target, plugins.HookRegistry._hooks)
|
||||
|
||||
def test_duplicate_plugin_names_fail_before_import(self):
|
||||
entries = [_entry_point("duplicate", "one"), _entry_point("duplicate", "two")]
|
||||
with (
|
||||
patch.object(plugins, "entry_points", return_value=entries),
|
||||
patch.object(plugins, "_get_excluded_dists", return_value=set()),
|
||||
patch.object(plugins, "get_selected_platform_dist", return_value=None),
|
||||
self.assertRaisesRegex(RuntimeError, "must be unique"),
|
||||
):
|
||||
plugins._discover()
|
||||
|
||||
for entry_point in entries:
|
||||
entry_point.load.assert_not_called()
|
||||
|
||||
def test_a_failing_load_from_the_selected_platform_aborts_startup(self):
|
||||
broken = _entry_point("vendor", "vendor-pkg")
|
||||
broken.load.side_effect = RuntimeError("vendor wheel is broken")
|
||||
unrelated = _entry_point("other", "other-pkg")
|
||||
unrelated.load.side_effect = RuntimeError("third party is broken")
|
||||
|
||||
with (
|
||||
patch.object(plugins, "entry_points", return_value=[unrelated]),
|
||||
patch.object(plugins, "_get_excluded_dists", return_value=set()),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
):
|
||||
self.assertEqual(plugins._discover(), {})
|
||||
|
||||
with (
|
||||
patch.object(plugins, "entry_points", return_value=[broken]),
|
||||
patch.object(plugins, "_get_excluded_dists", return_value=set()),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
self.assertRaisesRegex(RuntimeError, "vendor wheel is broken"),
|
||||
):
|
||||
plugins._discover()
|
||||
|
||||
def test_a_failing_callback_from_the_selected_platform_aborts_startup(self):
|
||||
boom = MagicMock(side_effect=RuntimeError("vendor hook is broken"))
|
||||
|
||||
with (
|
||||
patch.object(plugins, "_discover", return_value={"x": (boom, "other-pkg")}),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
):
|
||||
plugins._register_plugins_once()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
plugins, "_discover", return_value={"v": (boom, "vendor-pkg")}
|
||||
),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value="vendor-pkg"
|
||||
),
|
||||
self.assertRaisesRegex(RuntimeError, "vendor hook is broken"),
|
||||
):
|
||||
plugins._register_plugins_once()
|
||||
|
||||
def test_an_unapplied_hook_from_the_selected_platform_aborts_startup(self):
|
||||
# apply_hooks() logs and moves on, so a required target can go unpatched.
|
||||
target = "test_diffusion_plugins.unappliable"
|
||||
self.addCleanup(plugins.HookRegistry._hooks.pop, target, None)
|
||||
self.addCleanup(plugins.HookRegistry._patched.discard, target)
|
||||
plugins.HookRegistry._hooks[target] = [
|
||||
(HookType.AFTER, lambda r: r, HookSource("v", "vendor-pkg"))
|
||||
]
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "could not apply hooks"):
|
||||
plugins._require_hooks_applied("vendor-pkg")
|
||||
|
||||
plugins._require_hooks_applied("other-pkg")
|
||||
|
||||
plugins.HookRegistry._patched.add(target)
|
||||
plugins._require_hooks_applied("vendor-pkg")
|
||||
|
||||
def test_plugin_hook_uses_the_diffusion_registry(self):
|
||||
target = "test_diffusion_plugins.target"
|
||||
|
||||
def hook():
|
||||
pass
|
||||
|
||||
self.addCleanup(plugins.HookRegistry._hooks.pop, target, None)
|
||||
plugins.plugin_hook(target)(hook)
|
||||
|
||||
self.assertIs(plugins.HookRegistry._hooks[target][0][1], hook)
|
||||
self.assertNotIn(target, SrtHookRegistry._hooks)
|
||||
|
||||
def test_excludes_every_unselected_platform_distribution(self):
|
||||
entries = [
|
||||
_entry_point("selected", "selected-package"),
|
||||
_entry_point("selected_extra", "selected-package"),
|
||||
_entry_point("other", "other-package"),
|
||||
]
|
||||
# None is a built-in platform: nothing installed is in use.
|
||||
cases = (
|
||||
("selected-package", {"other-package"}),
|
||||
(None, {"selected-package", "other-package"}),
|
||||
)
|
||||
for selected_dist, expected in cases:
|
||||
with (
|
||||
self.subTest(selected_dist=selected_dist),
|
||||
patch.object(
|
||||
plugins, "get_selected_platform_dist", return_value=selected_dist
|
||||
),
|
||||
patch.object(plugins, "entry_points", return_value=entries),
|
||||
):
|
||||
self.assertEqual(plugins._get_excluded_dists(), expected)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -550,7 +550,7 @@ def test_validate_server_args_requires_packed_varlen_backend():
|
||||
resolve_component_attention_backend=lambda *_names: (None, None),
|
||||
)
|
||||
with patch(
|
||||
"sglang.multimodal_gen.configs.pipeline_configs.minimax_h3.get_attn_backend"
|
||||
"sglang.multimodal_gen.runtime.layers.attention.selector.get_attn_backend"
|
||||
) as get_attn_backend:
|
||||
MiniMaxH3PipelineConfig.validate_server_args(config, server_args)
|
||||
get_attn_backend.assert_called_once_with(
|
||||
@@ -560,7 +560,7 @@ def test_validate_server_args_requires_packed_varlen_backend():
|
||||
attention_requirements=AttentionRequirements(packed_varlen=True),
|
||||
)
|
||||
with patch(
|
||||
"sglang.multimodal_gen.configs.pipeline_configs.minimax_h3.get_attn_backend",
|
||||
"sglang.multimodal_gen.runtime.layers.attention.selector.get_attn_backend",
|
||||
side_effect=ValueError("does not implement packed varlen attention"),
|
||||
):
|
||||
with pytest.raises(ValueError, match="does not implement packed varlen"):
|
||||
@@ -590,7 +590,7 @@ def test_validate_server_args_accepts_transformer_backend_override():
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.multimodal_gen.configs.pipeline_configs.minimax_h3.get_attn_backend"
|
||||
"sglang.multimodal_gen.runtime.layers.attention.selector.get_attn_backend"
|
||||
) as get_attn_backend:
|
||||
MiniMaxH3PipelineConfig.validate_server_args(config, server_args)
|
||||
get_attn_backend.assert_called_once_with(
|
||||
@@ -621,7 +621,7 @@ def test_resolve_transformer_attention_backend_uses_selector_precedence():
|
||||
),
|
||||
)
|
||||
with patch(
|
||||
"sglang.multimodal_gen.configs.pipeline_configs.minimax_h3."
|
||||
"sglang.multimodal_gen.runtime.layers.attention.selector."
|
||||
"get_global_forced_attn_backend",
|
||||
return_value=forced_backend,
|
||||
):
|
||||
|
||||
@@ -0,0 +1,345 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.multimodal_gen.runtime.platforms as runtime_platforms
|
||||
from sglang.multimodal_gen.runtime.distributed import group_coordinator, parallel_state
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator import (
|
||||
CpuCommunicator,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers import custom_op
|
||||
from sglang.multimodal_gen.runtime.layers.activation import SiluAndMul
|
||||
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
||||
from sglang.multimodal_gen.runtime.managers import gpu_worker
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
|
||||
|
||||
class _OotPlatform(Platform):
|
||||
_enum = PlatformEnum.OOT
|
||||
device_name = "test"
|
||||
device_type = "test"
|
||||
|
||||
def get_dispatch_key_name(self) -> str:
|
||||
return "test"
|
||||
|
||||
|
||||
class _DispatchKeyOotPlatform(_OotPlatform):
|
||||
dispatch_key = "PrivateUse1"
|
||||
|
||||
|
||||
class _ExistingCommunicatorPlatform(_OotPlatform):
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "test.LegacyCommunicator"
|
||||
|
||||
|
||||
class _TestOp(CustomOp):
|
||||
def forward_native(self, value):
|
||||
return ("native", value)
|
||||
|
||||
|
||||
class _CudaCompatibleTestOp(_TestOp):
|
||||
def forward_cuda(self, value):
|
||||
return ("cuda", value)
|
||||
|
||||
|
||||
class _TestCommunicator(DeviceCommunicatorBase):
|
||||
pass
|
||||
|
||||
|
||||
class TestOotCustomOpDispatch(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
CustomOp._oot_forward_registry.pop("test", None)
|
||||
CustomOp._oot_forward_registry.pop("other", None)
|
||||
|
||||
def test_registered_forward_is_used(self):
|
||||
forward = MagicMock(return_value=("registered", 7))
|
||||
CustomOp.register_oot_forward(_TestOp, fn=forward, platform_key="test")
|
||||
with patch.object(custom_op.platforms, "_current_platform", _OotPlatform()):
|
||||
op = _TestOp()
|
||||
self.assertEqual(op(7), ("registered", 7))
|
||||
forward.assert_called_once_with(op, 7)
|
||||
|
||||
def test_missing_registration_uses_native_fallback(self):
|
||||
CustomOp.register_oot_forward(_TestOp, fn=MagicMock(), platform_key="other")
|
||||
with patch.object(custom_op.platforms, "_current_platform", _OotPlatform()):
|
||||
op = _TestOp()
|
||||
self.assertEqual(op(7), ("native", 7))
|
||||
|
||||
def test_compiled_forward_does_not_recompile_after_dispatch(self):
|
||||
def oot_forward(op, value):
|
||||
return op.forward_native(value)
|
||||
|
||||
CustomOp.register_oot_forward(SiluAndMul, fn=oot_forward, platform_key="test")
|
||||
cases = (
|
||||
(SiluAndMul, torch.randn(2, 8)),
|
||||
(lambda: RMSNorm(4), torch.randn(2, 4)),
|
||||
)
|
||||
|
||||
for factory, value in cases:
|
||||
with self.subTest(op=factory):
|
||||
compile_count = 0
|
||||
|
||||
def counting_backend(graph_module, _example_inputs):
|
||||
nonlocal compile_count
|
||||
compile_count += 1
|
||||
return graph_module.forward
|
||||
|
||||
with patch.object(
|
||||
custom_op.platforms, "_current_platform", _OotPlatform()
|
||||
):
|
||||
op = factory()
|
||||
selected_forward = op._forward_method
|
||||
compiled = torch.compile(op, backend=counting_backend, fullgraph=True)
|
||||
expected = op.forward_native(value)
|
||||
torch.testing.assert_close(compiled(value), expected)
|
||||
torch.testing.assert_close(compiled(value), expected)
|
||||
|
||||
self.assertEqual(compile_count, 1)
|
||||
self.assertIs(op._forward_method, selected_forward)
|
||||
|
||||
def test_platform_dispatch_key_can_reuse_an_existing_forward(self):
|
||||
platform = _OotPlatform()
|
||||
platform.get_dispatch_key_name = lambda: "cuda"
|
||||
with patch.object(custom_op.platforms, "_current_platform", platform):
|
||||
self.assertEqual(_CudaCompatibleTestOp()(7), ("cuda", 7))
|
||||
|
||||
def test_platform_dispatch_key_must_be_nonempty(self):
|
||||
platform = _OotPlatform()
|
||||
platform.get_dispatch_key_name = lambda: " "
|
||||
with (
|
||||
patch.object(custom_op.platforms, "_current_platform", platform),
|
||||
self.assertRaisesRegex(ValueError, "non-empty"),
|
||||
):
|
||||
_TestOp()(7)
|
||||
|
||||
|
||||
class TestOotBackendInit(unittest.TestCase):
|
||||
def setUp(self):
|
||||
state = (
|
||||
runtime_platforms._backend_init_done,
|
||||
runtime_platforms._backend_init_error,
|
||||
)
|
||||
self.addCleanup(self._restore_backend_state, state)
|
||||
runtime_platforms._backend_init_done = False
|
||||
runtime_platforms._backend_init_error = None
|
||||
|
||||
@staticmethod
|
||||
def _restore_backend_state(state):
|
||||
(
|
||||
runtime_platforms._backend_init_done,
|
||||
runtime_platforms._backend_init_error,
|
||||
) = state
|
||||
|
||||
def test_backend_initialization_runs_once(self):
|
||||
platform = _OotPlatform()
|
||||
with (
|
||||
patch.object(runtime_platforms, "_current_platform", platform),
|
||||
patch.object(platform, "init_backend") as init_backend,
|
||||
):
|
||||
runtime_platforms.initialize_current_platform()
|
||||
runtime_platforms.initialize_current_platform()
|
||||
|
||||
init_backend.assert_called_once_with()
|
||||
|
||||
def test_backend_initialization_failure_is_not_retried(self):
|
||||
platform = _OotPlatform()
|
||||
error = RuntimeError("backend unavailable")
|
||||
with (
|
||||
patch.object(runtime_platforms, "_current_platform", platform),
|
||||
patch.object(platform, "init_backend", side_effect=error) as init_backend,
|
||||
):
|
||||
with self.assertRaisesRegex(RuntimeError, "backend unavailable"):
|
||||
runtime_platforms.initialize_current_platform()
|
||||
with self.assertRaisesRegex(RuntimeError, "previously failed"):
|
||||
runtime_platforms.initialize_current_platform()
|
||||
|
||||
init_backend.assert_called_once_with()
|
||||
|
||||
def test_interrupted_backend_initialization_is_not_reported_as_success(self):
|
||||
"""An interrupt must leave the process failed, not silently initialized."""
|
||||
platform = _OotPlatform()
|
||||
with (
|
||||
patch.object(runtime_platforms, "_current_platform", platform),
|
||||
patch.object(
|
||||
platform, "init_backend", side_effect=KeyboardInterrupt
|
||||
) as init_backend,
|
||||
):
|
||||
with self.assertRaises(KeyboardInterrupt):
|
||||
runtime_platforms.initialize_current_platform()
|
||||
with self.assertRaisesRegex(RuntimeError, "previously failed"):
|
||||
runtime_platforms.initialize_current_platform()
|
||||
|
||||
init_backend.assert_called_once_with()
|
||||
|
||||
def test_worker_runs_init_backend_before_building_the_scheduler(self):
|
||||
order = []
|
||||
platform = MagicMock()
|
||||
platform.is_cuda.return_value = False
|
||||
platform.is_musa.return_value = False
|
||||
|
||||
with (
|
||||
patch.object(gpu_worker, "current_platform", platform),
|
||||
patch.object(
|
||||
gpu_worker,
|
||||
"initialize_current_platform",
|
||||
side_effect=lambda: order.append("init_backend"),
|
||||
),
|
||||
patch.object(gpu_worker, "kill_itself_when_parent_died"),
|
||||
patch.object(gpu_worker, "configure_logger"),
|
||||
patch.object(gpu_worker, "globally_suppress_loggers"),
|
||||
patch.object(
|
||||
gpu_worker,
|
||||
"init_diffusion_tracing",
|
||||
side_effect=lambda *a, **k: order.append("tracing"),
|
||||
),
|
||||
patch.object(
|
||||
gpu_worker.PortArgs,
|
||||
"from_server_args",
|
||||
side_effect=RuntimeError("stop before Scheduler"),
|
||||
),
|
||||
self.assertRaisesRegex(RuntimeError, "stop before Scheduler"),
|
||||
):
|
||||
gpu_worker.run_scheduler_process(
|
||||
local_rank=0,
|
||||
rank=0,
|
||||
server_args=MagicMock(),
|
||||
pipe_writer=None,
|
||||
)
|
||||
|
||||
self.assertEqual(order, ["init_backend", "tracing"])
|
||||
|
||||
|
||||
class TestOotRequiredConfiguration(unittest.TestCase):
|
||||
def test_device_and_dispatch_defaults_fail_loudly(self):
|
||||
platform = type("Oot", (Platform,), {"_enum": PlatformEnum.OOT})()
|
||||
|
||||
with self.assertRaisesRegex(NotImplementedError, "implement get_device"):
|
||||
platform.get_device(0)
|
||||
with self.assertRaisesRegex(NotImplementedError, "define dispatch_key"):
|
||||
platform.get_torch_library_dispatch_key()
|
||||
with self.assertRaisesRegex(
|
||||
NotImplementedError, "implement get_all_to_all_communicator_cls"
|
||||
):
|
||||
platform.get_all_to_all_communicator_cls()
|
||||
|
||||
self.assertEqual(
|
||||
_DispatchKeyOotPlatform().get_torch_library_dispatch_key(),
|
||||
"PrivateUse1",
|
||||
)
|
||||
|
||||
def test_builtin_torch_library_dispatch_is_preserved(self):
|
||||
platform = Platform()
|
||||
for is_npu, expected in ((False, "CUDA"), (True, "PrivateUse1")):
|
||||
with (
|
||||
self.subTest(is_npu=is_npu),
|
||||
patch.object(platform, "is_out_of_tree", return_value=False),
|
||||
patch.object(platform, "is_npu", return_value=is_npu),
|
||||
):
|
||||
self.assertEqual(platform.get_torch_library_dispatch_key(), expected)
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms.xpu import XpuPlatform
|
||||
|
||||
self.assertEqual(XpuPlatform().get_torch_library_dispatch_key(), "CUDA")
|
||||
|
||||
def test_existing_communicator_override_remains_the_fallback(self):
|
||||
self.assertEqual(
|
||||
_ExistingCommunicatorPlatform.get_all_to_all_communicator_cls(),
|
||||
"test.LegacyCommunicator",
|
||||
)
|
||||
|
||||
|
||||
class TestOotRuntimeHooks(unittest.TestCase):
|
||||
def test_builtin_overrides_keep_non_cuda_all_to_all_on_cpu(self):
|
||||
from sglang.multimodal_gen.runtime.platforms.mps import MpsPlatform
|
||||
from sglang.multimodal_gen.runtime.platforms.npu import NPUPlatformBase
|
||||
from sglang.multimodal_gen.runtime.platforms.xpu import XpuPlatform
|
||||
|
||||
for platform_cls in (MpsPlatform, NPUPlatformBase, XpuPlatform):
|
||||
with self.subTest(platform=platform_cls.__name__):
|
||||
with patch.object(
|
||||
group_coordinator, "current_platform", platform_cls()
|
||||
):
|
||||
self.assertIs(
|
||||
group_coordinator._resolve_all_to_all_communicator_cls(),
|
||||
CpuCommunicator,
|
||||
)
|
||||
|
||||
def test_platform_selects_all_to_all_communicator(self):
|
||||
platform = MagicMock()
|
||||
platform.get_all_to_all_communicator_cls.return_value = "vendor.Communicator"
|
||||
|
||||
with (
|
||||
patch.object(group_coordinator, "current_platform", platform),
|
||||
patch.object(
|
||||
group_coordinator,
|
||||
"resolve_name",
|
||||
return_value=_TestCommunicator,
|
||||
) as resolve_name,
|
||||
):
|
||||
self.assertIs(
|
||||
group_coordinator._resolve_all_to_all_communicator_cls(),
|
||||
_TestCommunicator,
|
||||
)
|
||||
|
||||
resolve_name.assert_called_once_with("vendor.Communicator")
|
||||
|
||||
def test_rejects_invalid_all_to_all_communicator(self):
|
||||
platform = MagicMock()
|
||||
platform.get_all_to_all_communicator_cls.return_value = "vendor.Communicator"
|
||||
|
||||
with (
|
||||
patch.object(group_coordinator, "current_platform", platform),
|
||||
patch.object(group_coordinator, "resolve_name", return_value=object),
|
||||
self.assertRaisesRegex(TypeError, "DeviceCommunicatorBase subclass"),
|
||||
):
|
||||
group_coordinator._resolve_all_to_all_communicator_cls()
|
||||
|
||||
def test_platform_controls_distributed_device_id(self):
|
||||
device_id = object()
|
||||
for supported in (False, True):
|
||||
platform = MagicMock(device_name="test")
|
||||
platform.get_torch_distributed_backend_str.return_value = "gloo"
|
||||
platform.supports_distributed_device_id.return_value = supported
|
||||
|
||||
with (
|
||||
self.subTest(supported=supported),
|
||||
patch.object(runtime_platforms, "_current_platform", platform),
|
||||
patch.object(parallel_state, "_WORLD", SimpleNamespace(world_size=1)),
|
||||
patch.object(
|
||||
parallel_state.torch.distributed,
|
||||
"is_initialized",
|
||||
return_value=False,
|
||||
),
|
||||
patch.object(
|
||||
parallel_state.torch.distributed, "init_process_group"
|
||||
) as init_process_group,
|
||||
patch.object(
|
||||
parallel_state.torch.distributed,
|
||||
"get_world_size",
|
||||
return_value=1,
|
||||
),
|
||||
patch.object(parallel_state, "_sync_srt_world_group"),
|
||||
):
|
||||
parallel_state.init_distributed_environment(device_id=device_id)
|
||||
|
||||
kwargs = init_process_group.call_args.kwargs
|
||||
if supported:
|
||||
self.assertIs(kwargs["device_id"], device_id)
|
||||
else:
|
||||
self.assertNotIn("device_id", kwargs)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,7 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
@@ -12,6 +14,15 @@ class NVMLUnavailableError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _entry_point(name: str, result: str | None, dist: str | None = None):
|
||||
entry_point = MagicMock(name=f"entry_point_{name}")
|
||||
entry_point.name = name
|
||||
entry_point.value = f"test_plugin:{name}"
|
||||
entry_point.dist = SimpleNamespace(name=dist) if dist else None
|
||||
entry_point.load.return_value = MagicMock(return_value=result)
|
||||
return entry_point
|
||||
|
||||
|
||||
class TestCudaPlatformDetection(unittest.TestCase):
|
||||
def test_torch_fallback_excludes_hip(self):
|
||||
cases = (
|
||||
@@ -29,8 +40,8 @@ class TestCudaPlatformDetection(unittest.TestCase):
|
||||
"sglang.multimodal_gen.runtime.platforms.pynvml.nvmlInit",
|
||||
side_effect=NVMLUnavailableError,
|
||||
),
|
||||
patch.object(platforms.os.path, "isfile", return_value=False),
|
||||
patch.object(platforms.os.path, "exists", return_value=False),
|
||||
patch.object(os.path, "isfile", return_value=False),
|
||||
patch.object(os.path, "exists", return_value=False),
|
||||
patch.object(torch.version, "hip", hip_version, create=True),
|
||||
patch.object(torch.cuda, "is_available", return_value=True),
|
||||
patch.object(torch.cuda, "device_count", return_value=1),
|
||||
@@ -38,5 +49,233 @@ class TestCudaPlatformDetection(unittest.TestCase):
|
||||
self.assertEqual(platforms.cuda_platform_plugin(), expected)
|
||||
|
||||
|
||||
class TestDiffusionPlatformPlugins(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.env = patch.dict(
|
||||
os.environ,
|
||||
{"SGLANG_DIFFUSION_PLATFORM_OVERRIDE": ""},
|
||||
)
|
||||
self.env.start()
|
||||
self.addCleanup(self.env.stop)
|
||||
|
||||
current_platform = platforms._current_platform
|
||||
current_selection = platforms._current_platform_selection
|
||||
self.addCleanup(setattr, platforms, "_current_platform", current_platform)
|
||||
self.addCleanup(
|
||||
setattr, platforms, "_current_platform_selection", current_selection
|
||||
)
|
||||
self._reset_current_platform()
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_selected_platform_records_its_distribution(self, entry_points):
|
||||
class _FakeOot(platforms.Platform):
|
||||
_enum = platforms.PlatformEnum.OOT
|
||||
device_name = "fake"
|
||||
device_type = "fake"
|
||||
|
||||
cases = (("selected", "explicit override"), ("", "automatic discovery"))
|
||||
for override, description in cases:
|
||||
with (
|
||||
self.subTest(description=description),
|
||||
patch.object(platforms, "resolve_name", return_value=_FakeOot),
|
||||
):
|
||||
self._reset_current_platform()
|
||||
entry_points.return_value = [
|
||||
_entry_point("selected", "vendor.platform.Platform", "vendor-pkg"),
|
||||
_entry_point("inactive", None, "other-pkg"),
|
||||
]
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = override
|
||||
|
||||
self.assertEqual(platforms.get_selected_platform_dist(), "vendor-pkg")
|
||||
self.assertIsInstance(platforms._current_platform, _FakeOot)
|
||||
|
||||
def test_accessor_reports_no_distribution_for_a_builtin(self):
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = "cpu"
|
||||
self._reset_current_platform()
|
||||
|
||||
self.assertIsNone(platforms.get_selected_platform_dist())
|
||||
|
||||
def _reset_current_platform(self):
|
||||
platforms._current_platform = None
|
||||
platforms._current_platform_selection = None
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_explicit_selection_loads_only_selected_plugin(self, entry_points):
|
||||
selected = _entry_point("selected", "vendor.platform.Platform")
|
||||
ignored = [
|
||||
_entry_point("duplicate", None),
|
||||
_entry_point("duplicate", None),
|
||||
_entry_point("cuda", "squatter.Platform"),
|
||||
]
|
||||
entry_points.return_value = [selected, *ignored]
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = "selected"
|
||||
|
||||
self.assertEqual(
|
||||
platforms.resolve_current_platform_cls_qualname(),
|
||||
"vendor.platform.Platform",
|
||||
)
|
||||
selected.load.assert_called_once_with()
|
||||
for entry_point in ignored:
|
||||
entry_point.load.assert_not_called()
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_auto_detection_requires_one_active_plugin(self, entry_points):
|
||||
entry_points.return_value = [
|
||||
_entry_point("inactive", None),
|
||||
_entry_point("active", "vendor.platform.Platform"),
|
||||
]
|
||||
self.assertEqual(
|
||||
platforms.resolve_current_platform_cls_qualname(),
|
||||
"vendor.platform.Platform",
|
||||
)
|
||||
|
||||
entry_points.return_value = [
|
||||
_entry_point("first", "first.Platform"),
|
||||
_entry_point("second", "second.Platform"),
|
||||
]
|
||||
with self.assertRaisesRegex(RuntimeError, "Multiple platform plugins"):
|
||||
platforms.resolve_current_platform_cls_qualname()
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_invalid_entry_point_names_fail_before_import(self, entry_points):
|
||||
cases = (
|
||||
([_entry_point("same", None), _entry_point("same", None)], "", "unique"),
|
||||
(
|
||||
[_entry_point("same", None), _entry_point("same", None)],
|
||||
"same",
|
||||
"unique",
|
||||
),
|
||||
([_entry_point("XPU", None)], "", "built-in"),
|
||||
)
|
||||
for entries, selected, message in cases:
|
||||
with self.subTest(message=message):
|
||||
entry_points.return_value = entries
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = selected
|
||||
with self.assertRaisesRegex(RuntimeError, message):
|
||||
platforms.resolve_current_platform_cls_qualname()
|
||||
for entry_point in entries:
|
||||
entry_point.load.assert_not_called()
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_explicit_selection_requires_active_match(self, entry_points):
|
||||
cases = (
|
||||
([], ValueError, "not found"),
|
||||
([_entry_point("selected", None)], RuntimeError, "returned None"),
|
||||
)
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = "selected"
|
||||
for entries, error_type, message in cases:
|
||||
with self.subTest(message=message):
|
||||
entry_points.return_value = entries
|
||||
with self.assertRaisesRegex(error_type, message):
|
||||
platforms.resolve_current_platform_cls_qualname()
|
||||
|
||||
def test_builtin_override_bypasses_plugin_selection(self):
|
||||
expected = {
|
||||
"cpu": "sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform",
|
||||
"cuda": "sglang.multimodal_gen.runtime.platforms.cuda.CudaPlatform",
|
||||
"rocm": "sglang.multimodal_gen.runtime.platforms.rocm.RocmPlatform",
|
||||
"mps": "sglang.multimodal_gen.runtime.platforms.mps.MpsPlatform",
|
||||
"npu": "sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase",
|
||||
"musa": "sglang.multimodal_gen.runtime.platforms.musa.MusaPlatform",
|
||||
"xpu": "sglang.multimodal_gen.runtime.platforms.xpu.XpuPlatform",
|
||||
}
|
||||
|
||||
for name, qualname in expected.items():
|
||||
with (
|
||||
self.subTest(name=name),
|
||||
patch.object(platforms, "entry_points") as entry_points,
|
||||
):
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = name
|
||||
self.assertEqual(
|
||||
platforms.resolve_current_platform_cls_qualname(),
|
||||
qualname,
|
||||
)
|
||||
entry_points.assert_not_called()
|
||||
|
||||
@patch.object(platforms, "entry_points", return_value=[])
|
||||
def test_xpu_keeps_automatic_detection_priority(self, _entry_points):
|
||||
xpu_qualname = "sglang.multimodal_gen.runtime.platforms.xpu.XpuPlatform"
|
||||
detectors = {
|
||||
"mps": MagicMock(return_value=None),
|
||||
"xpu": MagicMock(return_value=xpu_qualname),
|
||||
"rocm": MagicMock(return_value=None),
|
||||
"cuda": MagicMock(return_value=None),
|
||||
"npu": MagicMock(return_value=None),
|
||||
"musa": MagicMock(return_value=None),
|
||||
"cpu": MagicMock(return_value=None),
|
||||
}
|
||||
with patch.object(platforms, "builtin_platform_plugins", detectors):
|
||||
self.assertEqual(
|
||||
platforms.resolve_current_platform_cls_qualname(), xpu_qualname
|
||||
)
|
||||
detectors["mps"].assert_called_once_with()
|
||||
detectors["xpu"].assert_called_once_with()
|
||||
for name in ("rocm", "cuda", "npu", "musa", "cpu"):
|
||||
detectors[name].assert_not_called()
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_external_plugin_cannot_return_a_builtin_platform(self, entry_points):
|
||||
entry_points.return_value = [
|
||||
_entry_point(
|
||||
"selected",
|
||||
"sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform",
|
||||
"vendor-pkg",
|
||||
)
|
||||
]
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = "selected"
|
||||
|
||||
with self.assertRaisesRegex(TypeError, "PlatformEnum.OOT"):
|
||||
platforms.get_selected_platform_dist()
|
||||
|
||||
self.assertIsNone(platforms._current_platform)
|
||||
self.assertIsNone(platforms._current_platform_selection)
|
||||
|
||||
def test_external_platform_identity_is_validated_before_publication(self):
|
||||
class BadPlatform(platforms.Platform):
|
||||
_enum = platforms.PlatformEnum.OOT
|
||||
device_name = "fake"
|
||||
device_type = "fake"
|
||||
|
||||
selection = platforms._PlatformSelection(
|
||||
"vendor.BadPlatform", "selected", "vendor-pkg"
|
||||
)
|
||||
for attribute in ("device_name", "device_type"):
|
||||
with (
|
||||
self.subTest(attribute=attribute),
|
||||
patch.object(BadPlatform, attribute, " "),
|
||||
patch.object(
|
||||
platforms, "_select_current_platform", return_value=selection
|
||||
),
|
||||
patch.object(
|
||||
platforms, "_load_platform_class", return_value=BadPlatform
|
||||
),
|
||||
self.assertRaisesRegex(TypeError, attribute),
|
||||
):
|
||||
platforms._resolve_current_platform()
|
||||
|
||||
self.assertIsNone(platforms._current_platform)
|
||||
self.assertIsNone(platforms._current_platform_selection)
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_explicit_plugin_must_return_a_class_qualname(self, entry_points):
|
||||
entry_points.return_value = [_entry_point("selected", 42, "vendor-pkg")]
|
||||
os.environ["SGLANG_DIFFUSION_PLATFORM_OVERRIDE"] = "selected"
|
||||
|
||||
with self.assertRaisesRegex(TypeError, "non-empty class qualname"):
|
||||
platforms.resolve_current_platform_cls_qualname()
|
||||
|
||||
@patch.object(platforms, "entry_points")
|
||||
def test_a_failing_activation_is_not_downgraded_to_a_builtin(self, entry_points):
|
||||
# Skipping the plugin here would run the whole job on the wrong hardware.
|
||||
broken = _entry_point("broken", None)
|
||||
broken.load.return_value = MagicMock(
|
||||
side_effect=RuntimeError("vendor runtime is broken")
|
||||
)
|
||||
entry_points.return_value = [broken]
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "vendor runtime is broken"):
|
||||
platforms.resolve_current_platform_cls_qualname()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import torch.nn as nn
|
||||
@@ -7,9 +8,13 @@ import torch.nn as nn
|
||||
from sglang.multimodal_gen.runtime.models.dits.ltx_2 import (
|
||||
LTX2VideoTransformer3DModel,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import DecodingStage
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
||||
DenoisingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.mova import (
|
||||
MOVADenoisingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.torch_compile import (
|
||||
CompiledModuleRegistry,
|
||||
build_torch_compile_kwargs,
|
||||
@@ -99,6 +104,84 @@ def test_out_of_tree_platform_controls_compile_kwargs(backend, options, expected
|
||||
get_compile_options.assert_called_once_with(module)
|
||||
|
||||
|
||||
def test_mova_uses_platform_compile_kwargs():
|
||||
stage = MOVADenoisingStage.__new__(MOVADenoisingStage)
|
||||
module = _CompilableModule()
|
||||
server_args = SimpleNamespace(enable_torch_compile=True)
|
||||
module_path = (
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.mova"
|
||||
)
|
||||
compile_path = "sglang.multimodal_gen.runtime.utils.torch_compile"
|
||||
|
||||
with (
|
||||
# This test exercises the real mode resolver, which reads the env first.
|
||||
patch.dict(os.environ, {}, clear=False),
|
||||
patch(f"{module_path}.current_platform.is_hip", return_value=False),
|
||||
patch(f"{module_path}.current_platform.is_npu", return_value=False),
|
||||
patch(
|
||||
f"{compile_path}.current_platform.is_npu",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
f"{compile_path}.current_platform.is_out_of_tree",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
f"{compile_path}.current_platform.get_compile_backend",
|
||||
return_value="custom_backend",
|
||||
) as get_compile_backend,
|
||||
patch(
|
||||
f"{compile_path}.current_platform.get_compile_options",
|
||||
return_value={"max_autotune": True},
|
||||
) as get_compile_options,
|
||||
):
|
||||
os.environ.pop("SGLANG_TORCH_COMPILE_MODE", None)
|
||||
stage._maybe_enable_torch_compile(module, server_args)
|
||||
|
||||
get_compile_backend.assert_called_once_with("max-autotune-no-cudagraphs")
|
||||
get_compile_options.assert_called_once_with(module)
|
||||
assert module.compile_calls == [
|
||||
{
|
||||
"backend": "custom_backend",
|
||||
"dynamic": None,
|
||||
"fullgraph": False,
|
||||
"options": {"max_autotune": True},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_vae_compile_options_receive_the_target_module():
|
||||
stage = DecodingStage.__new__(DecodingStage)
|
||||
vae = _CompilableModule()
|
||||
decode = MagicMock()
|
||||
compiled_callable = MagicMock(target_id=None, compiled_module=None)
|
||||
compiled_callable.get_or_compile.return_value = decode
|
||||
server_args = SimpleNamespace(enable_torch_compile=True)
|
||||
module_path = "sglang.multimodal_gen.runtime.pipelines_core.stages.decoding"
|
||||
|
||||
with (
|
||||
patch(f"{module_path}.current_platform.is_npu", return_value=False),
|
||||
patch(
|
||||
f"{module_path}.resolve_torch_compile_kwargs",
|
||||
return_value=({"backend": "custom_backend"}, "default"),
|
||||
) as resolve_compile_kwargs,
|
||||
):
|
||||
result = stage._get_vae_decode_fn(
|
||||
vae,
|
||||
server_args,
|
||||
decode_fn=decode,
|
||||
compiled_callable=compiled_callable,
|
||||
)
|
||||
|
||||
assert result is decode
|
||||
resolve_compile_kwargs.assert_called_once_with(
|
||||
"SGLANG_VAE_TORCH_COMPILE_MODE",
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
default="default",
|
||||
module=vae,
|
||||
)
|
||||
|
||||
|
||||
def test_ltx2_compile_conditions_match_only_direct_blocks():
|
||||
conditions = LTX2VideoTransformer3DModel._compile_conditions
|
||||
|
||||
@@ -179,17 +262,19 @@ def test_denoising_stage_selects_regional_compile():
|
||||
),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.stages.denoising."
|
||||
"maybe_enable_inductor_compute_comm_overlap"
|
||||
),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.stages.denoising."
|
||||
"build_torch_compile_kwargs",
|
||||
return_value=compile_kwargs,
|
||||
) as build_compile_kwargs,
|
||||
"resolve_torch_compile_kwargs",
|
||||
return_value=(compile_kwargs, "default"),
|
||||
) as resolve_compile_kwargs,
|
||||
):
|
||||
stage._maybe_torch_compile(model)
|
||||
|
||||
build_compile_kwargs.assert_called_once_with(mode="default", module=model)
|
||||
resolve_compile_kwargs.assert_called_once_with(
|
||||
"SGLANG_TORCH_COMPILE_MODE",
|
||||
config=stage.server_args.pipeline_config.dit_config,
|
||||
default="max-autotune-no-cudagraphs",
|
||||
module=model,
|
||||
enable_inductor_compute_comm_overlap=True,
|
||||
)
|
||||
assert [len(block.compile_calls) for block in model.transformer_blocks] == [1, 1]
|
||||
assert [block.compile_calls for block in model.transformer_blocks] == [
|
||||
[compile_kwargs],
|
||||
|
||||
@@ -146,6 +146,18 @@ def _from_dict_without_model_resolution(
|
||||
return ServerArgs.from_dict(kwargs)
|
||||
|
||||
|
||||
class TestPlatformLifecycleHooks(unittest.TestCase):
|
||||
def test_server_args_applies_platform_defaults(self):
|
||||
with patch.object(
|
||||
current_platform, "apply_server_args_defaults"
|
||||
) as apply_defaults:
|
||||
server_args = _from_dict_without_model_resolution(
|
||||
{"model_path": "test/model"}
|
||||
)
|
||||
|
||||
apply_defaults.assert_called_once_with(server_args)
|
||||
|
||||
|
||||
class TestServerArgsPathExpansion(unittest.TestCase):
|
||||
def _from_dict_without_model_resolution(self, kwargs):
|
||||
return _from_dict_without_model_resolution(kwargs)
|
||||
|
||||
@@ -0,0 +1,337 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import pathlib
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.multimodal_gen.runtime.managers import worker_bootstrap
|
||||
|
||||
WORKER_MODULE = "sglang.multimodal_gen.runtime.managers.gpu_worker"
|
||||
GENERATOR_MODULE = "sglang.multimodal_gen.runtime.entrypoints.diffusion_generator"
|
||||
SERVER_ARGS_MODULE = "sglang.multimodal_gen.runtime.server_args.server_args"
|
||||
HTTP_SERVER_MODULE = "sglang.multimodal_gen.runtime.launch_server"
|
||||
|
||||
METADATA = "Metadata-Version: 2.1\nName: sgl-fake-plugin\nVersion: 0.1\n"
|
||||
ENTRY_POINTS = """\
|
||||
[sglang.multimodal_gen.platforms]
|
||||
fake = sgl_fake_plugin:activate
|
||||
[sglang.multimodal_gen.plugins]
|
||||
fake = sgl_fake_plugin:register
|
||||
"""
|
||||
|
||||
|
||||
# Real modules, not embedded source strings, so they are linted like any other
|
||||
# file. The plugin imports no diffusion module of its own, which is what keeps
|
||||
# the import-order measurement honest.
|
||||
FIXTURES_DIR = pathlib.Path(__file__).parent / "fixtures"
|
||||
FAKE_PLUGIN_FIXTURE = FIXTURES_DIR / "sgl_fake_plugin.py"
|
||||
FACADE_IMPORT_SCRIPT = FIXTURES_DIR / "offline_script_facade_import.py"
|
||||
RUNTIME_IMPORT_SCRIPT = FIXTURES_DIR / "offline_script_runtime_import.py"
|
||||
|
||||
PYTHON_ROOT = pathlib.Path(__file__).parents[4]
|
||||
EARLY_IMPORT_WARNING = "imported before this worker initialized its platform"
|
||||
SCRIPT_TIMEOUT_S = 300
|
||||
|
||||
|
||||
def _install_fake_plugin_dist(root: pathlib.Path) -> None:
|
||||
shutil.copy(FAKE_PLUGIN_FIXTURE, root / "sgl_fake_plugin.py")
|
||||
dist_info = root / "sgl_fake_plugin-0.1.dist-info"
|
||||
dist_info.mkdir()
|
||||
(dist_info / "METADATA").write_text(METADATA)
|
||||
(dist_info / "entry_points.txt").write_text(ENTRY_POINTS)
|
||||
|
||||
|
||||
def _check_cli_import_order(pipe_writer) -> None:
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli import main as cli_main
|
||||
|
||||
imported_before_activation = GENERATOR_MODULE in sys.modules
|
||||
|
||||
class StopAtPluginBoundary(Exception):
|
||||
pass
|
||||
|
||||
def stop_before_command_imports():
|
||||
raise StopAtPluginBoundary
|
||||
|
||||
cli_main.apply_plugin_hooks = stop_before_command_imports
|
||||
try:
|
||||
cli_main.generate_cmd_init()
|
||||
except StopAtPluginBoundary:
|
||||
pass
|
||||
|
||||
pipe_writer.send(
|
||||
{
|
||||
"imported_before_activation": imported_before_activation,
|
||||
"imported_after_failed_activation": GENERATOR_MODULE in sys.modules,
|
||||
}
|
||||
)
|
||||
pipe_writer.close()
|
||||
|
||||
|
||||
def _check_http_server_import_order(pipe_writer) -> None:
|
||||
from sglang.multimodal_gen.runtime.platforms import plugins
|
||||
|
||||
class StopAtPluginBoundary(Exception):
|
||||
pass
|
||||
|
||||
observed = {
|
||||
"http_server_imported_before_bootstrap": HTTP_SERVER_MODULE in sys.modules,
|
||||
"server_args_imported_before_bootstrap": SERVER_ARGS_MODULE in sys.modules,
|
||||
}
|
||||
|
||||
def stop_before_runtime_imports():
|
||||
observed.update(
|
||||
http_server_imported_when_hooks_applied=(HTTP_SERVER_MODULE in sys.modules),
|
||||
server_args_imported_when_hooks_applied=(SERVER_ARGS_MODULE in sys.modules),
|
||||
)
|
||||
raise StopAtPluginBoundary
|
||||
|
||||
plugins.load_plugins = lambda: None
|
||||
plugins.apply_plugin_hooks = stop_before_runtime_imports
|
||||
try:
|
||||
worker_bootstrap.bootstrap_http_server_process(None)
|
||||
except StopAtPluginBoundary:
|
||||
pass
|
||||
|
||||
pipe_writer.send(observed)
|
||||
pipe_writer.close()
|
||||
|
||||
|
||||
class TestBootstrapImportBoundary(unittest.TestCase):
|
||||
def test_manager_namespace_does_not_hide_early_worker_imports(self):
|
||||
for module, warned in (
|
||||
("sglang.multimodal_gen.runtime.managers", False),
|
||||
("sglang.multimodal_gen.runtime.platforms.plugins", False),
|
||||
(worker_bootstrap.__name__, False),
|
||||
(WORKER_MODULE, True),
|
||||
):
|
||||
modules = {
|
||||
"__main__": SimpleNamespace(__file__="offline.py"),
|
||||
module: None,
|
||||
}
|
||||
with (
|
||||
self.subTest(module=module),
|
||||
patch.object(worker_bootstrap, "sys", SimpleNamespace(modules=modules)),
|
||||
patch.object(worker_bootstrap.logging, "getLogger") as get_logger,
|
||||
):
|
||||
worker_bootstrap._warn_if_runtime_imported_early()
|
||||
if warned:
|
||||
warning = get_logger.return_value.warning
|
||||
warning.assert_called_once()
|
||||
self.assertEqual(warning.call_args.args[1], module)
|
||||
else:
|
||||
get_logger.assert_not_called()
|
||||
|
||||
|
||||
class TestSpawnedWorkerReceivesPluginOverride(unittest.TestCase):
|
||||
"""End-to-end over a real spawn, with a real entry-point distribution.
|
||||
|
||||
A spawned child re-imports from a blank interpreter, so nothing the parent
|
||||
patched survives. The child must initialize its backend, register its own
|
||||
hooks, and apply them before invoking the worker.
|
||||
"""
|
||||
|
||||
def test_lifecycle_precedes_worker_import_and_argument_materialization(self):
|
||||
# The parent has a real ServerArgs object; the process boundary must
|
||||
# keep its class opaque until bootstrap chooses to materialize it.
|
||||
# Imported before the fake dist reaches sys.path: this resolves the
|
||||
# platform, which must not land on a class living in a temp directory.
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = pathlib.Path(tmp)
|
||||
_install_fake_plugin_dist(root)
|
||||
|
||||
# Spawn ships sys.path to the child, so the dist is discoverable there.
|
||||
sys.path.insert(0, str(root))
|
||||
self.addCleanup(sys.path.remove, str(root))
|
||||
importlib.invalidate_caches()
|
||||
|
||||
server_args = ServerArgs.__new__(ServerArgs)
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
spec = worker_bootstrap.SchedulerProcessSpec(
|
||||
local_rank=0,
|
||||
rank=0,
|
||||
server_args=worker_bootstrap.ServerArgsPayload.capture(server_args),
|
||||
pipe_writer=writer,
|
||||
)
|
||||
|
||||
process = mp.get_context("spawn").Process(
|
||||
target=worker_bootstrap.bootstrap_scheduler_process,
|
||||
args=(spec,),
|
||||
)
|
||||
process.start()
|
||||
writer.close()
|
||||
self.addCleanup(process.join, 10)
|
||||
self.addCleanup(process.kill)
|
||||
|
||||
result = None
|
||||
if reader.poll(120):
|
||||
try:
|
||||
result = reader.recv()
|
||||
except EOFError:
|
||||
pass
|
||||
if result is None:
|
||||
process.join(10)
|
||||
self.fail(
|
||||
"child sent nothing back, so the override never ran "
|
||||
f"(exit code {process.exitcode})"
|
||||
)
|
||||
|
||||
self.assertTrue(result["override_ran"], "plugin override did not run")
|
||||
self.assertIs(
|
||||
result["worker_imported_when_plugin_ran"],
|
||||
False,
|
||||
"plugins loaded after the worker module was already imported",
|
||||
)
|
||||
self.assertIs(
|
||||
result["generator_imported_when_plugin_ran"],
|
||||
False,
|
||||
"the package facade imported the diffusion runtime before plugins loaded",
|
||||
)
|
||||
self.assertIs(
|
||||
result["worker_imported_when_backend_initialized"],
|
||||
False,
|
||||
"worker imports preceded platform backend initialization",
|
||||
)
|
||||
self.assertIs(
|
||||
result["server_args_imported_when_backend_initialized"],
|
||||
False,
|
||||
"spawn unpickled ServerArgs before platform backend initialization",
|
||||
)
|
||||
self.assertIs(
|
||||
result["server_args_imported_when_plugin_ran"],
|
||||
False,
|
||||
"spawn materialized ServerArgs before plugin registration",
|
||||
)
|
||||
self.assertTrue(
|
||||
result["backend_initialized_when_plugin_ran"],
|
||||
"plugin registration ran before platform backend initialization",
|
||||
)
|
||||
self.assertTrue(
|
||||
result["backend_initialized"],
|
||||
"platform backend initialized after the worker override ran",
|
||||
)
|
||||
|
||||
def test_cli_activates_plugins_before_importing_commands(self):
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
process = mp.get_context("spawn").Process(
|
||||
target=_check_cli_import_order,
|
||||
args=(writer,),
|
||||
)
|
||||
process.start()
|
||||
writer.close()
|
||||
self.addCleanup(process.join, 10)
|
||||
self.addCleanup(process.kill)
|
||||
|
||||
self.assertTrue(reader.poll(30), "child did not report CLI import state")
|
||||
result = reader.recv()
|
||||
self.assertFalse(result["imported_before_activation"])
|
||||
self.assertFalse(result["imported_after_failed_activation"])
|
||||
|
||||
def test_http_hooks_apply_before_runtime_imports(self):
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
process = mp.get_context("spawn").Process(
|
||||
target=_check_http_server_import_order,
|
||||
args=(writer,),
|
||||
)
|
||||
process.start()
|
||||
writer.close()
|
||||
self.addCleanup(process.join, 10)
|
||||
self.addCleanup(process.kill)
|
||||
|
||||
self.assertTrue(reader.poll(30), "child did not report HTTP import state")
|
||||
result = reader.recv()
|
||||
self.assertEqual(
|
||||
result,
|
||||
{
|
||||
"http_server_imported_before_bootstrap": False,
|
||||
"server_args_imported_before_bootstrap": False,
|
||||
"http_server_imported_when_hooks_applied": False,
|
||||
"server_args_imported_when_hooks_applied": False,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class TestOfflineScriptImportContract(unittest.TestCase):
|
||||
"""Real scripts in a real interpreter, because spawn re-executes the
|
||||
launching script's module scope before it unpickles anything."""
|
||||
|
||||
def _run_offline_script(self, script: pathlib.Path):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = pathlib.Path(tmp)
|
||||
dist_root = root / "site"
|
||||
dist_root.mkdir()
|
||||
_install_fake_plugin_dist(dist_root)
|
||||
result_path = root / "observed.json"
|
||||
|
||||
env = dict(os.environ)
|
||||
env["PYTHONPATH"] = os.pathsep.join(
|
||||
path
|
||||
for path in (
|
||||
str(dist_root),
|
||||
str(FIXTURES_DIR),
|
||||
str(PYTHON_ROOT),
|
||||
env.get("PYTHONPATH", ""),
|
||||
)
|
||||
if path
|
||||
)
|
||||
completed = subprocess.run(
|
||||
[sys.executable, str(script), str(result_path)],
|
||||
env=env,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=SCRIPT_TIMEOUT_S,
|
||||
check=False,
|
||||
)
|
||||
self.assertEqual(completed.returncode, 0, completed.stderr)
|
||||
payload = json.loads(result_path.read_text())
|
||||
|
||||
self.assertIsNotNone(
|
||||
payload["observed"],
|
||||
f"child sent nothing back (exit code {payload['exitcode']})",
|
||||
)
|
||||
return payload["observed"], completed.stderr
|
||||
|
||||
def test_a_module_scope_facade_import_leaves_the_child_lifecycle_intact(self):
|
||||
observed, stderr = self._run_offline_script(FACADE_IMPORT_SCRIPT)
|
||||
|
||||
self.assertTrue(observed["override_ran"], "plugin override did not run")
|
||||
self.assertIs(
|
||||
observed["generator_imported_when_plugin_ran"],
|
||||
False,
|
||||
"re-executing the script imported the generator before plugins loaded",
|
||||
)
|
||||
self.assertNotIn(EARLY_IMPORT_WARNING, stderr)
|
||||
|
||||
def test_a_module_scope_runtime_import_is_reported_by_the_child(self):
|
||||
observed, stderr = self._run_offline_script(RUNTIME_IMPORT_SCRIPT)
|
||||
|
||||
self.assertTrue(observed["override_ran"], "plugin override did not run")
|
||||
self.assertIs(
|
||||
observed["generator_imported_when_plugin_ran"],
|
||||
True,
|
||||
"the script layout under test no longer imports the runtime early",
|
||||
)
|
||||
self.assertIn(
|
||||
EARLY_IMPORT_WARNING,
|
||||
stderr,
|
||||
"the child accepted a mis-ordered import without reporting it",
|
||||
)
|
||||
self.assertIn(
|
||||
RUNTIME_IMPORT_SCRIPT.name,
|
||||
stderr,
|
||||
"the report did not name the script whose imports have to move",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user