[diffusion] feat: support diffusers backend - run any model supported by diffusers (#14112)

This commit is contained in:
Adarsh Shirawalmath
2026-01-06 12:30:57 +08:00
committed by GitHub
parent 84d13c54bb
commit 7be1a8c70c
17 changed files with 1120 additions and 57 deletions
+100 -19
View File
@@ -12,7 +12,20 @@ import importlib
import os
import pkgutil
from functools import lru_cache
from typing import Any, Callable, Dict, List, Optional, Tuple, Type
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
List,
Optional,
Tuple,
Type,
Union,
)
if TYPE_CHECKING:
from sglang.multimodal_gen.runtime.server_args import Backend
from sglang.multimodal_gen.configs.pipeline_configs import (
FastHunyuanConfig,
@@ -236,8 +249,34 @@ class ModelInfo:
pipeline_config_cls: Type[PipelineConfig]
def _get_diffusers_model_info(model_path: str) -> ModelInfo:
"""
Get model info for diffusers backend.
Returns a ModelInfo with DiffusersPipeline and generic configs.
"""
from sglang.multimodal_gen.configs.pipeline_configs.diffusers_generic import (
DiffusersGenericPipelineConfig,
)
from sglang.multimodal_gen.configs.sample.diffusers_generic import (
DiffusersGenericSamplingParams,
)
from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import (
DiffusersPipeline,
)
return ModelInfo(
pipeline_cls=DiffusersPipeline,
sampling_param_cls=DiffusersGenericSamplingParams,
pipeline_config_cls=DiffusersGenericPipelineConfig,
)
@lru_cache(maxsize=1)
def get_model_info(model_path: str) -> Optional[ModelInfo]:
def get_model_info(
model_path: str,
backend: Optional[Union[str, "Backend"]] = None,
) -> Optional[ModelInfo]:
"""
Resolves all necessary classes (pipeline, sampling, config) for a given model path.
@@ -246,7 +285,31 @@ def get_model_info(model_path: str) -> Optional[ModelInfo]:
'_class_name' against an auto-discovered registry of pipeline implementations.
2. Resolves the associated configuration classes (for sampling and pipeline) using a
manually registered mapping based on the model path.
Args:
model_path: Path to the model or HuggingFace model ID
backend: Backend to use ('auto', 'sglang', 'diffusers'). If None, uses 'auto'.
Returns:
ModelInfo with the resolved pipeline class and config classes, or None if not found.
"""
# import Backend enum here to avoid circular imports
from sglang.multimodal_gen.runtime.server_args import Backend
# Normalize backend
if backend is None:
backend = Backend.AUTO
elif isinstance(backend, str):
backend = Backend.from_string(backend)
# Handle explicit diffusers backend
if backend == Backend.DIFFUSERS:
logger.info(
"Using diffusers backend for model '%s' (explicitly requested)", model_path
)
return _get_diffusers_model_info(model_path)
# For AUTO or SGLANG backend, try native implementation first
# 1. Discover all available pipeline classes and cache them
_discover_and_register_pipelines()
@@ -258,32 +321,55 @@ def get_model_info(model_path: str) -> Optional[ModelInfo]:
config = maybe_download_model_index(model_path)
except Exception as e:
logger.error(f"Could not read model config for '{model_path}': {e}")
if backend == Backend.AUTO:
logger.info("Falling back to diffusers backend")
return _get_diffusers_model_info(model_path)
return None
pipeline_class_name = config.get("_class_name")
if not pipeline_class_name:
logger.error(f"'_class_name' not found in model_index.json for '{model_path}'")
if backend == Backend.AUTO:
logger.info("Falling back to diffusers backend")
return _get_diffusers_model_info(model_path)
return None
pipeline_cls = _PIPELINE_REGISTRY.get(pipeline_class_name)
if not pipeline_cls:
logger.error(
f"Pipeline class '{pipeline_class_name}' specified in '{model_path}' is not a registered EntryClass in the framework. "
f"Available pipelines: {list(_PIPELINE_REGISTRY.keys())}"
)
return None
if backend == Backend.AUTO:
logger.warning(
f"Pipeline class '{pipeline_class_name}' specified in '{model_path}' has no native sglang support. "
f"Falling back to diffusers backend."
)
return _get_diffusers_model_info(model_path)
else:
logger.error(
f"Pipeline class '{pipeline_class_name}' specified in '{model_path}' is not a registered EntryClass in the framework. "
f"Available pipelines: {list(_PIPELINE_REGISTRY.keys())}. "
f"Consider using --backend diffusers to use vanilla diffusers pipeline."
)
return None
# 3. Get configuration classes (sampling, pipeline config)
config_info = _get_config_info(model_path)
if not config_info:
logger.error(
f"Could not resolve configuration for model '{model_path}'. "
"It is not a registered model path or detected by any registered model family detectors. "
f"Known model paths: {list(_MODEL_HF_PATH_TO_NAME.keys())}"
)
return None
if backend == Backend.AUTO:
logger.warning(
f"Could not resolve native configuration for model '{model_path}'. "
f"Falling back to diffusers backend."
)
return _get_diffusers_model_info(model_path)
else:
logger.error(
f"Could not resolve configuration for model '{model_path}'. "
"It is not a registered model path or detected by any registered model family detectors. "
f"Known model paths: {list(_MODEL_HF_PATH_TO_NAME.keys())}. "
f"Consider using --backend diffusers to use vanilla diffusers pipeline."
)
return None
# 4. Combine the complete model info
# 4. Combine and return the complete model info
logger.info("Using native sglang backend for model '%s'", model_path)
model_info = ModelInfo(
pipeline_cls=pipeline_cls,
sampling_param_cls=config_info.sampling_param_cls,
@@ -312,7 +398,6 @@ def _register_configs():
"FastVideo/FastHunyuan-diffusers",
],
)
# Wan
register_configs(
sampling_param_cls=WanT2V_1_3B_SamplingParams,
@@ -372,7 +457,6 @@ def _register_configs():
"Wan-AI/Wan2.2-TI2V-5B-Diffusers",
],
)
register_configs(
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
@@ -381,7 +465,6 @@ def _register_configs():
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
],
)
register_configs(
sampling_param_cls=Wan2_2_T2V_A14B_SamplingParam,
pipeline_config_cls=Wan2_2_T2V_A14B_Config,
@@ -399,7 +482,6 @@ def _register_configs():
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
],
)
# FLUX
register_configs(
sampling_param_cls=FluxSamplingParams,
@@ -425,7 +507,6 @@ def _register_configs():
],
model_detectors=[lambda hf_id: "z-image" in hf_id.lower()],
)
# Qwen-Image
register_configs(
sampling_param_cls=QwenImageSamplingParams,