sglang/srt/platforms/device_mixin.py |
- PlatformEnum + DeviceMixin base class |
+ DeviceMixin base class and SRT platform identity types |
sglang/srt/platforms/interface.py |
@@ -847,6 +1056,22 @@ Target paths use fully-qualified dotted notation. Both formats are supported:
sglang/srt/platforms/__init__.py |
current_platform lazy singleton + discovery logic |
+
+ sglang/multimodal_gen/runtime/platforms/interface.py |
+ Diffusion Platform base class |
+
+
+ sglang/multimodal_gen/runtime/platforms/__init__.py |
+ Diffusion current_platform lazy singleton and built-in fallback order |
+
+
+ sglang/multimodal_gen/runtime/platforms/plugins.py |
+ Diffusion plugin registration, explicit hook-application phase, and hook registry |
+
+
+ sglang/multimodal_gen/runtime/managers/worker_bootstrap.py |
+ Import-neutral process specifications and spawn targets that initialize the backend before resolving runtime hooks |
+
sglang/srt/plugins/__init__.py |
load_plugins() + load_plugins_by_group() |
diff --git a/docs/docs/sglang-diffusion/contributing.mdx b/docs/docs/sglang-diffusion/contributing.mdx
index 1b2130a06..f219b6889 100644
--- a/docs/docs/sglang-diffusion/contributing.mdx
+++ b/docs/docs/sglang-diffusion/contributing.mdx
@@ -11,7 +11,7 @@ This guide outlines the requirements for contributing to the SGLang Diffusion mo
- [Support New Models](/docs/sglang-diffusion/support_new_models): implementation guide for adding new diffusion pipelines
- [Helper ownership](/docs/sglang-diffusion/support_new_models#place-helpers-with-their-owners): where to put shared and model-specific utilities
- [CI Performance](./ci_perf): update and regenerate perf baselines
-
+- [SGLang Plugin System](/docs/hardware-platforms/plugin): package a hardware platform outside the SGLang repository
## On AI-Assisted ("Vibe Coding") PRs
@@ -76,3 +76,13 @@ Consider adding tests to the `pr-test` or `nightly-test` suites to safeguard you
Please run the according testcase, then update/add the baseline to `perf_baselines.json` by following the instruction in console if applicable.
See [test](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen/test) for examples
+
+## Out-of-tree platform changes
+
+Use the `sglang.multimodal_gen.platforms` entry-point group when a hardware integration can live in a separately installed package. The [plugin guide](/docs/hardware-platforms/plugin#add-diffusion-support) describes the activation callback, diffusion platform contract, and selection behavior.
+
+When you change the platform interface or discovery implementation in SGLang:
+
+- Add focused CPU-only tests with a fake OOT platform.
+- Run the relevant SRT plugin tests when changing shared hook-registry behavior.
+- Document workload-specific kernels or optimized paths that still require vendor integration.
diff --git a/docs/docs/sglang-diffusion/disaggregation.mdx b/docs/docs/sglang-diffusion/disaggregation.mdx
index 593b0df38..e9b721ffb 100644
--- a/docs/docs/sglang-diffusion/disaggregation.mdx
+++ b/docs/docs/sglang-diffusion/disaggregation.mdx
@@ -424,24 +424,25 @@ If not specified, parallelism is auto-derived from `--num-gpus`.
## Python API
-For programmatic single-machine deployment, `launch_pool_disagg_server()` is available:
+For programmatic single-machine deployment, `launch_pool_disagg_server()` is available. It spawns its workers, and each child re-executes this script's module scope before initializing its platform, so keep the diffusion imports inside the guard:
```python
-from sglang.multimodal_gen.runtime.server_args import ServerArgs
-from sglang.multimodal_gen.runtime.launch_server import launch_pool_disagg_server
+if __name__ == "__main__":
+ from sglang.multimodal_gen.runtime.launch_server import launch_pool_disagg_server
+ from sglang.multimodal_gen.runtime.server_args import ServerArgs
-server_args = ServerArgs.from_kwargs(
- model_path="Wan-AI/Wan2.1-T2V-14B-Diffusers",
- denoiser_sp=4, denoiser_ulysses=2, denoiser_ring=2,
- disagg_ib_device="mlx5_0",
-)
+ server_args = ServerArgs.from_kwargs(
+ model_path="Wan-AI/Wan2.1-T2V-14B-Diffusers",
+ denoiser_sp=4, denoiser_ulysses=2, denoiser_ring=2,
+ disagg_ib_device="mlx5_0",
+ )
-launch_pool_disagg_server(
- server_args,
- encoder_gpus=[[0]],
- denoiser_gpus=[[1, 2, 3, 4], [5, 6, 7, 8]],
- decoder_gpus=[[0]],
-)
+ launch_pool_disagg_server(
+ server_args,
+ encoder_gpus=[[0]],
+ denoiser_gpus=[[1, 2, 3, 4], [5, 6, 7, 8]],
+ decoder_gpus=[[0]],
+ )
```
## Architecture
diff --git a/docs/docs/sglang-diffusion/environment_variables.mdx b/docs/docs/sglang-diffusion/environment_variables.mdx
index 7ef3aa19d..3f837fbe8 100644
--- a/docs/docs/sglang-diffusion/environment_variables.mdx
+++ b/docs/docs/sglang-diffusion/environment_variables.mdx
@@ -23,6 +23,16 @@ description: "Configure SGLang diffusion behavior with environment variables."
not set |
Installed package that registers out-of-tree diffusion pipelines and component models. The package is imported once in every process. |
+
+ SGLANG_DIFFUSION_PLATFORM_OVERRIDE |
+ not set |
+ Select cpu, cuda, rocm, mps, npu, musa, or an installed sglang.multimodal_gen.platforms entry-point name. XPU remains automatic-only. When unset, SGLang Diffusion detects installed platform plugins before trying built-in platforms. |
+
+
+ SGLANG_PLUGINS |
+ not set |
+ Comma-separated allowlist for installed sglang.multimodal_gen.plugins hooks. When unset, all discovered diffusion hooks load. |
+
SGLANG_DIFFUSION_TARGET_DEVICE |
cuda |
@@ -68,11 +78,6 @@ description: "Configure SGLang diffusion behavior with environment variables."
INFO |
Default logging level |
-
- SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD |
- fork |
- Multiprocess context for workers (fork or spawn) |
-
SGLANG_DIFFUSION_IPC_A2A |
true |
diff --git a/python/sglang/cli/serve.py b/python/sglang/cli/serve.py
index 6e913622d..516e7a693 100644
--- a/python/sglang/cli/serve.py
+++ b/python/sglang/cli/serve.py
@@ -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,
diff --git a/python/sglang/multimodal_gen/__init__.py b/python/sglang/multimodal_gen/__init__.py
index 3d7545560..4bf4e0c9e 100644
--- a/python/sglang/multimodal_gen/__init__.py
+++ b/python/sglang/multimodal_gen/__init__.py
@@ -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
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py b/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py
index 791bb0f4f..db7557bd6 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py
@@ -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,
diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py
index b13616dbf..a06a57857 100644
--- a/python/sglang/multimodal_gen/envs.py
+++ b/python/sglang/multimodal_gen/envs.py
@@ -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", ""
),
diff --git a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py
index 3e253da66..a6fc8d961 100644
--- a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py
+++ b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py
@@ -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
diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py
index e0143a0d8..5123d0dbb 100644
--- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py
+++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py
@@ -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:
diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py
index af5be8722..2d824f9ef 100644
--- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py
+++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py
@@ -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")
diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py
index d55bd454f..dc6efe5d0 100644
--- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py
+++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py
@@ -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)
diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py
index 15cd99608..5d06929b9 100644
--- a/python/sglang/multimodal_gen/runtime/launch_server.py
+++ b/python/sglang/multimodal_gen/runtime/launch_server.py
@@ -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:
diff --git a/python/sglang/multimodal_gen/runtime/layers/custom_op.py b/python/sglang/multimodal_gen/runtime/layers/custom_op.py
index 4151cc2b0..0fa1c3fb1 100644
--- a/python/sglang/multimodal_gen/runtime/layers/custom_op.py
+++ b/python/sglang/multimodal_gen/runtime/layers/custom_op.py
@@ -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
diff --git a/python/sglang/multimodal_gen/runtime/layers/utils.py b/python/sglang/multimodal_gen/runtime/layers/utils.py
index 9cc70491a..58e2416ca 100644
--- a/python/sglang/multimodal_gen/runtime/layers/utils.py
+++ b/python/sglang/multimodal_gen/runtime/layers/utils.py
@@ -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)
diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py
index 5f1826594..da4a412c8 100644
--- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py
+++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py
@@ -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():
diff --git a/python/sglang/multimodal_gen/runtime/managers/worker_bootstrap.py b/python/sglang/multimodal_gen/runtime/managers/worker_bootstrap.py
new file mode 100644
index 000000000..0c59a60ac
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/managers/worker_bootstrap.py
@@ -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())
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/__init__.py b/python/sglang/multimodal_gen/runtime/pipelines_core/__init__.py
index 2f39d2c8c..ea5e80a56 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/__init__.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/__init__.py
@@ -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
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
index ef91dfbe9..b0e5df668 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
@@ -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)
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
index 4fa862cc9..b7b6adaf2 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
@@ -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):
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py
index df50ae9ff..5eace7fd2 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py
@@ -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
diff --git a/python/sglang/multimodal_gen/runtime/platforms/__init__.py b/python/sglang/multimodal_gen/runtime/platforms/__init__.py
index ce46524ab..12fb2309a 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/__init__.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/__init__.py
@@ -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",
+]
diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py
index 129e92c5d..7d971204f 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/interface.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py
@@ -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_`` 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."""
diff --git a/python/sglang/multimodal_gen/runtime/platforms/mps.py b/python/sglang/multimodal_gen/runtime/platforms/mps.py
index d4547407d..e4b25a4aa 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/mps.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/mps.py
@@ -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."""
diff --git a/python/sglang/multimodal_gen/runtime/platforms/npu.py b/python/sglang/multimodal_gen/runtime/platforms/npu.py
index 478a0f2db..032cf71f4 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/npu.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/npu.py
@@ -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."""
diff --git a/python/sglang/multimodal_gen/runtime/platforms/plugins.py b/python/sglang/multimodal_gen/runtime/platforms/plugins.py
new file mode 100644
index 000000000..186f7b45f
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/platforms/plugins.py
@@ -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",
+]
diff --git a/python/sglang/multimodal_gen/runtime/platforms/xpu.py b/python/sglang/multimodal_gen/runtime/platforms/xpu.py
index f6064455f..5fa2d52e5 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/xpu.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/xpu.py
@@ -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"
+ )
diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py
index c0ea92da0..cf0198f99 100644
--- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py
+++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py
@@ -1897,6 +1897,7 @@ class ServerArgs(DisaggServerArgsMixin):
self
)
+ current_platform.apply_server_args_defaults(self)
# configure logger before use
configure_logger(server_args=self)
diff --git a/python/sglang/multimodal_gen/runtime/utils/torch_compile.py b/python/sglang/multimodal_gen/runtime/utils/torch_compile.py
index 677ba2871..d67f35076 100644
--- a/python/sglang/multimodal_gen/runtime/utils/torch_compile.py
+++ b/python/sglang/multimodal_gen/runtime/utils/torch_compile.py
@@ -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,
*,
diff --git a/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_facade_import.py b/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_facade_import.py
new file mode 100644
index 000000000..ae2143e81
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_facade_import.py
@@ -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()
diff --git a/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_runner.py b/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_runner.py
new file mode 100644
index 000000000..4ee88957c
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_runner.py
@@ -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])
diff --git a/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_runtime_import.py b/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_runtime_import.py
new file mode 100644
index 000000000..9b8b02274
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/fixtures/offline_script_runtime_import.py
@@ -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()
diff --git a/python/sglang/multimodal_gen/test/unit/fixtures/sgl_fake_plugin.py b/python/sglang/multimodal_gen/test/unit/fixtures/sgl_fake_plugin.py
new file mode 100644
index 000000000..598147767
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/fixtures/sgl_fake_plugin.py
@@ -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
+ )
diff --git a/python/sglang/multimodal_gen/test/unit/test_diffusion_plugins.py b/python/sglang/multimodal_gen/test/unit/test_diffusion_plugins.py
new file mode 100644
index 000000000..75d2dce0d
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_diffusion_plugins.py
@@ -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()
diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py
index def45664d..1c53dce4f 100644
--- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py
+++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py
@@ -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,
):
diff --git a/python/sglang/multimodal_gen/test/unit/test_oot_platform_runtime.py b/python/sglang/multimodal_gen/test/unit/test_oot_platform_runtime.py
new file mode 100644
index 000000000..fec96a2a1
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_oot_platform_runtime.py
@@ -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()
diff --git a/python/sglang/multimodal_gen/test/unit/test_platform_detection.py b/python/sglang/multimodal_gen/test/unit/test_platform_detection.py
index 23cf65c51..7dcf7250f 100644
--- a/python/sglang/multimodal_gen/test/unit/test_platform_detection.py
+++ b/python/sglang/multimodal_gen/test/unit/test_platform_detection.py
@@ -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()
diff --git a/python/sglang/multimodal_gen/test/unit/test_regional_torch_compile.py b/python/sglang/multimodal_gen/test/unit/test_regional_torch_compile.py
index 80eeff9de..9ce2e22b6 100644
--- a/python/sglang/multimodal_gen/test/unit/test_regional_torch_compile.py
+++ b/python/sglang/multimodal_gen/test/unit/test_regional_torch_compile.py
@@ -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],
diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py
index ca942cca0..fa2efb18a 100644
--- a/python/sglang/multimodal_gen/test/unit/test_server_args.py
+++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py
@@ -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)
diff --git a/python/sglang/multimodal_gen/test/unit/test_worker_bootstrap.py b/python/sglang/multimodal_gen/test/unit/test_worker_bootstrap.py
new file mode 100644
index 000000000..0ba34bee8
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_worker_bootstrap.py
@@ -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()