[diffusion] feat: support out-of-tree models and pipelines (#35713)

This commit is contained in:
Mick
2026-08-21 00:33:34 +08:00
committed by GitHub
parent 7f8f030000
commit be373395b4
8 changed files with 297 additions and 16 deletions
@@ -18,6 +18,11 @@ description: "Configure SGLang diffusion behavior with environment variables."
</tr> </tr>
</thead> </thead>
<tbody> <tbody>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_EXTERNAL_MODEL_PACKAGE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>not set</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Installed package that registers out-of-tree diffusion pipelines and component models. The package is imported once in every process.</td>
</tr>
<tr> <tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DIFFUSION_TARGET_DEVICE</code></td> <td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DIFFUSION_TARGET_DEVICE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>cuda</code></td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>cuda</code></td>
@@ -33,6 +33,51 @@ utilities, and common action-policy helpers. Model packages may call these
helpers. Keep ownership in shared runtime folders unless the code is truly helpers. Keep ownership in shared runtime folders unless the code is truly
architecture-specific. architecture-specific.
## Out-of-Tree Models and Pipelines
An installed package can register native component models and a pipeline
without modifying SGLang-Diffusion. Register them in the package's
`__init__.py`:
```python
from sglang.multimodal_gen.registry import register_pipeline
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from .configs import CustomPipelineConfig, CustomSamplingParams
from .pipeline import CustomPipeline
ModelRegistry.register_model(
"CustomTransformer2DModel",
"custom_diffusion.models:CustomTransformer2DModel",
)
register_pipeline(
CustomPipeline,
sampling_param_cls=CustomSamplingParams,
pipeline_config_cls=CustomPipelineConfig,
hf_model_paths=["my-org/custom-diffusion-model"],
model_detectors=[lambda value: "custom-diffusion" in value.lower()],
)
```
Install the package in the server environment and set the same environment
variable used by SRT plugins:
```bash
pip install -e /path/to/custom-diffusion
SGLANG_EXTERNAL_MODEL_PACKAGE=custom_diffusion \
sglang serve --model-path my-org/custom-diffusion-model
```
Notes:
- The string form of `register_model` keeps component imports lazy.
- `hf_model_paths` also supports checkpoints without `model_index.json`. Other
Diffusers checkpoints can select the pipeline through `_class_name`.
- For a standalone safetensors file, pass `--pipeline CustomPipeline`.
- Set the environment variable before startup. Each process imports the package
once. Use `overwrite=True` only to intentionally replace a built-in pipeline.
## Start With the Smallest Change ## Start With the Smallest Change
Before adding files, decide which path fits the model. Before adding files, decide which path fits the model.
+4
View File
@@ -24,6 +24,7 @@ if TYPE_CHECKING:
SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork" SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork"
SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda" SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda"
SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = "" SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = ""
SGLANG_EXTERNAL_MODEL_PACKAGE: str = ""
MAX_JOBS: str | None = None MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None CMAKE_BUILD_TYPE: str | None = None
@@ -209,6 +210,9 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_DIFFUSION_PLATFORM_OVERRIDE": _lazy_str( "SGLANG_DIFFUSION_PLATFORM_OVERRIDE": _lazy_str(
"SGLANG_DIFFUSION_PLATFORM_OVERRIDE", "" "SGLANG_DIFFUSION_PLATFORM_OVERRIDE", ""
), ),
# Import an installed package that registers out-of-tree diffusion models
# and pipelines. This is shared with the SRT model plugin mechanism.
"SGLANG_EXTERNAL_MODEL_PACKAGE": _lazy_str("SGLANG_EXTERNAL_MODEL_PACKAGE", ""),
# Enables torch profiler if set. Path to the directory where torch profiler # Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path. # traces are saved. Note that it must be an absolute path.
"SGLANG_DIFFUSION_TORCH_PROFILER_DIR": _lazy_path( "SGLANG_DIFFUSION_TORCH_PROFILER_DIR": _lazy_path(
+91 -4
View File
@@ -196,6 +196,9 @@ from sglang.multimodal_gen.configs.sample.zimage import (
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.utils.external_model_package import (
load_external_model_package,
)
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model_index, maybe_download_model_index,
) )
@@ -206,6 +209,7 @@ logger = init_logger(__name__)
# --- Part 1: Pipeline Discovery --- # --- Part 1: Pipeline Discovery ---
_PIPELINE_REGISTRY: Dict[str, Type[ComposedPipelineBase]] = {} _PIPELINE_REGISTRY: Dict[str, Type[ComposedPipelineBase]] = {}
_BUILTIN_PIPELINES_DISCOVERED = False
# Registry for pipeline configuration classes (for safetensors files without model_index.json) # Registry for pipeline configuration classes (for safetensors files without model_index.json)
# Maps pipeline_class_name -> (PipelineConfig class, SamplingParams class) # Maps pipeline_class_name -> (PipelineConfig class, SamplingParams class)
@@ -219,11 +223,13 @@ def _discover_and_register_pipelines():
finds modules with an 'EntryClass' attribute, and maps the class's 'pipeline_name' finds modules with an 'EntryClass' attribute, and maps the class's 'pipeline_name'
to the class itself in a global registry. to the class itself in a global registry.
""" """
if _PIPELINE_REGISTRY: # run only once global _BUILTIN_PIPELINES_DISCOVERED
if _BUILTIN_PIPELINES_DISCOVERED:
return return
package_name = "sglang.multimodal_gen.runtime.pipelines" package_name = "sglang.multimodal_gen.runtime.pipelines"
package = importlib.import_module(package_name) package = importlib.import_module(package_name)
_BUILTIN_PIPELINES_DISCOVERED = True
for _, module_name, ispkg in pkgutil.walk_packages( for _, module_name, ispkg in pkgutil.walk_packages(
package.__path__, package.__name__ + "." package.__path__, package.__name__ + "."
@@ -276,6 +282,11 @@ def _discover_and_register_pipelines():
) )
def _ensure_registry_initialized() -> None:
_discover_and_register_pipelines()
load_external_model_package()
def get_pipeline_config_classes( def get_pipeline_config_classes(
pipeline_class_name: str, pipeline_class_name: str,
) -> Tuple[Type[PipelineConfig], Type[Any]] | None: ) -> Tuple[Type[PipelineConfig], Type[Any]] | None:
@@ -283,10 +294,23 @@ def get_pipeline_config_classes(
Get the configuration classes for a pipeline. Get the configuration classes for a pipeline.
""" """
# Ensure pipelines are discovered first # Ensure pipelines are discovered first
_discover_and_register_pipelines() _ensure_registry_initialized()
return _PIPELINE_CONFIG_REGISTRY.get(pipeline_class_name) return _PIPELINE_CONFIG_REGISTRY.get(pipeline_class_name)
def get_pipeline_class(
pipeline_class_name: str,
) -> Type[ComposedPipelineBase] | None:
"""Get a registered pipeline class by name."""
_ensure_registry_initialized()
return _PIPELINE_REGISTRY.get(pipeline_class_name)
def get_registered_pipeline_names() -> List[str]:
_ensure_registry_initialized()
return list(_PIPELINE_REGISTRY)
# --- Part 2: Config Registration --- # --- Part 2: Config Registration ---
@dataclasses.dataclass @dataclasses.dataclass
class ConfigInfo: class ConfigInfo:
@@ -328,7 +352,7 @@ def register_configs(
pipeline_config_cls: Type[PipelineConfig], pipeline_config_cls: Type[PipelineConfig],
hf_model_paths: Optional[List[str]] = None, hf_model_paths: Optional[List[str]] = None,
model_detectors: Optional[List[Callable[[str], bool]]] = None, model_detectors: Optional[List[Callable[[str], bool]]] = None,
): ) -> str:
""" """
Registers configuration classes for a new model family. Registers configuration classes for a new model family.
""" """
@@ -349,6 +373,67 @@ def register_configs(
if model_detectors: if model_detectors:
for detector in model_detectors: for detector in model_detectors:
_MODEL_NAME_DETECTORS.append((model_id, detector)) _MODEL_NAME_DETECTORS.append((model_id, detector))
return model_id
def register_pipeline(
pipeline_cls: Type[ComposedPipelineBase],
*,
sampling_param_cls: Any,
pipeline_config_cls: Type[PipelineConfig],
hf_model_paths: Optional[List[str]] = None,
model_detectors: Optional[List[Callable[[str], bool]]] = None,
overwrite: bool = False,
) -> None:
"""Register an out-of-tree native diffusion pipeline and its configs."""
_discover_and_register_pipelines()
if not issubclass(pipeline_cls, ComposedPipelineBase):
raise TypeError("pipeline_cls must inherit from ComposedPipelineBase")
if not issubclass(pipeline_config_cls, PipelineConfig):
raise TypeError("pipeline_config_cls must inherit from PipelineConfig")
pipeline_name = pipeline_cls.pipeline_name
existing_pipeline = _PIPELINE_REGISTRY.get(pipeline_name)
if existing_pipeline is not None and not overwrite:
raise ValueError(
f"Pipeline '{pipeline_name}' is already registered; pass overwrite=True to replace it"
)
for model_path in hf_model_paths or []:
if model_path in _MODEL_HF_PATH_TO_NAME and not overwrite:
raise ValueError(f"Model path '{model_path}' is already registered")
registered_pipeline = KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS.get(
model_path.lower()
)
if registered_pipeline is not None and not overwrite:
raise ValueError(
f"Model path '{model_path}' is already registered for pipeline "
f"'{registered_pipeline}'"
)
_PIPELINE_REGISTRY[pipeline_name] = pipeline_cls
_PIPELINE_CONFIG_REGISTRY[pipeline_name] = (
pipeline_config_cls,
sampling_param_cls,
)
config_id = register_configs(
sampling_param_cls=sampling_param_cls,
pipeline_config_cls=pipeline_config_cls,
hf_model_paths=hf_model_paths,
model_detectors=None if overwrite else model_detectors,
)
if overwrite and model_detectors:
_MODEL_NAME_DETECTORS[:0] = [
(config_id, detector) for detector in model_detectors
]
for model_path in hf_model_paths or []:
KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS[model_path.lower()] = pipeline_name
_get_config_info.cache_clear()
get_model_info.cache_clear()
logger.info(
"Registered external diffusion pipeline '%s' from %s",
pipeline_name,
pipeline_cls.__module__,
)
def get_model_short_name(model_id: str) -> str: def get_model_short_name(model_id: str) -> str:
@@ -367,6 +452,7 @@ def _normalize_hf_cache_path(path: str) -> str:
def has_registered_diffusion_model_path(model_path: str) -> bool: def has_registered_diffusion_model_path(model_path: str) -> bool:
_ensure_registry_initialized()
all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(), key=len, reverse=True) all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(), key=len, reverse=True)
if model_path in _MODEL_HF_PATH_TO_NAME: if model_path in _MODEL_HF_PATH_TO_NAME:
@@ -396,6 +482,7 @@ def _get_config_info(
""" """
Gets the ConfigInfo for a given model path using mappings and detectors. Gets the ConfigInfo for a given model path using mappings and detectors.
""" """
_ensure_registry_initialized()
all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(), key=len, reverse=True) all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(), key=len, reverse=True)
# 0. Explicit model_id override: match by short name # 0. Explicit model_id override: match by short name
@@ -580,7 +667,7 @@ def get_model_info(
# For AUTO or SGLANG backend, try native implementation first # For AUTO or SGLANG backend, try native implementation first
# 1. Discover all available pipeline classes and cache them # 1. Discover all available pipeline classes and cache them
_discover_and_register_pipelines() _ensure_registry_initialized()
# Detect quantized models and fallback to diffusers # Detect quantized models and fallback to diffusers
is_quantized = any(q in model_path.lower() for q in ["-4bit", "-awq", "-gptq"]) is_quantized = any(q in model_path.lower() for q in ["-4bit", "-awq", "-gptq"])
@@ -19,6 +19,9 @@ from typing import NoReturn, TypeVar, cast
import cloudpickle import cloudpickle
from torch import nn from torch import nn
from sglang.multimodal_gen.runtime.utils.external_model_package import (
load_external_model_package,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -242,7 +245,6 @@ class _LazyRegisteredModel(_BaseRegisteredModel):
""" """
module_name: str module_name: str
component_name: str
class_name: str class_name: str
# Performed in another process to avoid initializing CUDA # Performed in another process to avoid initializing CUDA
@@ -289,10 +291,12 @@ class _ModelRegistry:
registered_models: dict[str, _BaseRegisteredModel] = field(default_factory=dict) registered_models: dict[str, _BaseRegisteredModel] = field(default_factory=dict)
def get_supported_archs(self) -> Set[str]: def get_supported_archs(self) -> Set[str]:
load_external_model_package()
return self.registered_models.keys() return self.registered_models.keys()
def resolve_by_alias(self, alias: str) -> type[nn.Module] | None: def resolve_by_alias(self, alias: str) -> type[nn.Module] | None:
"""Resolve a model class by its alias (external module path).""" """Resolve a model class by its alias (external module path)."""
load_external_model_package()
if alias in _ALIAS_TO_MODEL: if alias in _ALIAS_TO_MODEL:
canonical_name = _ALIAS_TO_MODEL[alias] canonical_name = _ALIAS_TO_MODEL[alias]
return self._try_load_model_cls(canonical_name) return self._try_load_model_cls(canonical_name)
@@ -304,7 +308,7 @@ class _ModelRegistry:
model_cls: type[nn.Module] | str, model_cls: type[nn.Module] | str,
) -> None: ) -> None:
""" """
Register an external model to be used in vLLM. Register an external model to be used in SGLang-Diffusion.
:code:`model_cls` can be either: :code:`model_cls` can be either:
@@ -328,7 +332,9 @@ class _ModelRegistry:
msg = "Expected a string in the format `<module>:<class>`" msg = "Expected a string in the format `<module>:<class>`"
raise ValueError(msg) raise ValueError(msg)
model = _LazyRegisteredModel(*split_str) model = _LazyRegisteredModel(
module_name=split_str[0], class_name=split_str[1]
)
else: else:
model = _RegisteredModel.from_model_cls(model_cls) model = _RegisteredModel.from_model_cls(model_cls)
@@ -364,6 +370,7 @@ class _ModelRegistry:
self, self,
architectures: str | list[str], architectures: str | list[str],
) -> list[str]: ) -> list[str]:
load_external_model_package()
if isinstance(architectures, str): if isinstance(architectures, str):
architectures = [architectures] architectures = [architectures]
if not architectures: if not architectures:
@@ -417,7 +424,6 @@ ModelRegistry = _ModelRegistry(
{ {
model_arch: _LazyRegisteredModel( model_arch: _LazyRegisteredModel(
module_name=f"sglang.multimodal_gen.runtime.models.{component_name}.{mod_relname}", module_name=f"sglang.multimodal_gen.runtime.models.{component_name}.{mod_relname}",
component_name=component_name,
class_name=cls_name, class_name=cls_name,
) )
for model_arch, ( for model_arch, (
@@ -48,20 +48,18 @@ def build_pipeline(
# Check if pipeline class is explicitly specified # Check if pipeline class is explicitly specified
if server_args.pipeline_class_name: if server_args.pipeline_class_name:
from sglang.multimodal_gen.registry import ( from sglang.multimodal_gen.registry import (
_PIPELINE_REGISTRY, get_pipeline_class,
_discover_and_register_pipelines, get_registered_pipeline_names,
) )
_discover_and_register_pipelines() available_pipelines = get_registered_pipeline_names()
logger.info(f"Requested pipeline_class_name: {server_args.pipeline_class_name}") logger.info(f"Requested pipeline_class_name: {server_args.pipeline_class_name}")
logger.info( logger.info(f"Available pipelines in registry: {available_pipelines}")
f"Available pipelines in registry: {list(_PIPELINE_REGISTRY.keys())}" pipeline_cls = get_pipeline_class(server_args.pipeline_class_name)
)
pipeline_cls = _PIPELINE_REGISTRY.get(server_args.pipeline_class_name)
if pipeline_cls is None: if pipeline_cls is None:
raise ValueError( raise ValueError(
f"Pipeline class '{server_args.pipeline_class_name}' not found in registry. " f"Pipeline class '{server_args.pipeline_class_name}' not found in registry. "
f"Available pipelines: {list(_PIPELINE_REGISTRY.keys())}" f"Available pipelines: {available_pipelines}"
) )
logger.info( logger.info(
f"✓ Using explicitly specified pipeline: {server_args.pipeline_class_name} (class: {pipeline_cls.__name__})" f"✓ Using explicitly specified pipeline: {server_args.pipeline_class_name} (class: {pipeline_cls.__name__})"
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
import importlib
from functools import lru_cache
from sglang.multimodal_gen import envs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@lru_cache(maxsize=1)
def load_external_model_package() -> None:
"""Import the configured out-of-tree model package once per process."""
package_name = envs.SGLANG_EXTERNAL_MODEL_PACKAGE
if not package_name:
return
logger.info("Loading external model package: %s", package_name)
importlib.import_module(package_name)
@@ -0,0 +1,116 @@
# SPDX-License-Identifier: Apache-2.0
import os
import subprocess
import sys
import textwrap
def test_external_diffusion_package_registration(tmp_path):
package_dir = tmp_path / "external_diffusion_plugin"
package_dir.mkdir()
(package_dir / "__init__.py").write_text(
textwrap.dedent("""
from sglang.multimodal_gen.registry import register_pipeline
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from .plugin import (
ExternalPipeline,
ExternalPipelineConfig,
ExternalSamplingParams,
)
ModelRegistry.register_model(
"ExternalTransformer",
"external_diffusion_plugin.plugin:ExternalTransformer",
)
register_pipeline(
ExternalPipeline,
sampling_param_cls=ExternalSamplingParams,
pipeline_config_cls=ExternalPipelineConfig,
hf_model_paths=["external-org/external-checkpoint"],
model_detectors=[lambda value: "external-checkpoint" in value],
)
"""),
encoding="utf-8",
)
(package_dir / "plugin.py").write_text(
textwrap.dedent("""
from torch import nn
from sglang.multimodal_gen.configs.pipeline_configs.base import (
PipelineConfig,
)
from sglang.multimodal_gen.configs.sample.sampling_params import (
SamplingParams,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
class ExternalTransformer(nn.Module):
pass
class ExternalPipelineConfig(PipelineConfig):
pass
class ExternalSamplingParams(SamplingParams):
pass
class ExternalPipeline(ComposedPipelineBase):
pipeline_name = "ExternalPipeline"
def create_pipeline_stages(self, server_args):
pass
"""),
encoding="utf-8",
)
model_dir = tmp_path / "external-checkpoint"
model_dir.mkdir()
script = textwrap.dedent("""
import sys
from sglang.cli.utils import get_is_diffusion_model
from sglang.multimodal_gen.registry import (
get_model_info,
get_pipeline_config_classes,
)
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
assert get_is_diffusion_model(sys.argv[1])
model_cls, architecture = ModelRegistry.resolve_model_cls(
"ExternalTransformer"
)
assert model_cls.__name__ == "ExternalTransformer"
assert architecture == "ExternalTransformer"
model_info = get_model_info(sys.argv[1], backend="sglang")
assert model_info.pipeline_cls.__name__ == "ExternalPipeline"
assert model_info.pipeline_config_cls.__name__ == "ExternalPipelineConfig"
assert model_info.sampling_param_cls.__name__ == "ExternalSamplingParams"
pipeline_config_cls, sampling_param_cls = get_pipeline_config_classes(
"ExternalPipeline"
)
assert pipeline_config_cls is model_info.pipeline_config_cls
assert sampling_param_cls is model_info.sampling_param_cls
""")
env = os.environ.copy()
env["SGLANG_EXTERNAL_MODEL_PACKAGE"] = "external_diffusion_plugin"
env["PYTHONPATH"] = os.pathsep.join(
path for path in (str(tmp_path), env.get("PYTHONPATH")) if path
)
result = subprocess.run(
[sys.executable, "-c", script, str(model_dir)],
check=False,
env=env,
capture_output=True,
text=True,
)
assert result.returncode == 0, result.stderr