[diffusion] refactor: refactor pipeline folders (#13253)

This commit is contained in:
Mick
2025-11-20 12:56:51 +08:00
committed by GitHub
parent 2e3a69ae05
commit bc42c8c415
60 changed files with 303 additions and 300 deletions
+1 -2
View File
@@ -1,6 +1,5 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # 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.pipelines import PipelineConfig
from sglang.multimodal_gen.configs.sample import SamplingParams from sglang.multimodal_gen.configs.sample import SamplingParams
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator
@@ -1,16 +1,16 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
from sglang.multimodal_gen.configs.pipelines.base import ( from sglang.multimodal_gen.configs.pipeline_configs.base import (
PipelineConfig, PipelineConfig,
SlidingTileAttnConfig, SlidingTileAttnConfig,
) )
from sglang.multimodal_gen.configs.pipelines.flux import FluxPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.flux import FluxPipelineConfig
from sglang.multimodal_gen.configs.pipelines.hunyuan import ( from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import (
FastHunyuanConfig, FastHunyuanConfig,
HunyuanConfig, HunyuanConfig,
) )
from sglang.multimodal_gen.configs.pipelines.stepvideo import StepVideoT2VConfig from sglang.multimodal_gen.configs.pipeline_configs.stepvideo import StepVideoT2VConfig
from sglang.multimodal_gen.configs.pipelines.wan import ( from sglang.multimodal_gen.configs.pipeline_configs.wan import (
SelfForcingWanT2V480PConfig, SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V480PConfig,
WanI2V720PConfig, WanI2V720PConfig,
@@ -13,12 +13,12 @@ from sglang.multimodal_gen.configs.models.encoders import (
T5Config, T5Config,
) )
from sglang.multimodal_gen.configs.models.vaes.flux import FluxVAEConfig from sglang.multimodal_gen.configs.models.vaes.flux import FluxVAEConfig
from sglang.multimodal_gen.configs.pipelines.base import ( from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType, ModelTaskType,
PipelineConfig, PipelineConfig,
preprocess_text, preprocess_text,
) )
from sglang.multimodal_gen.configs.pipelines.hunyuan import ( from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import (
clip_postprocess_text, clip_postprocess_text,
clip_preprocess_text, clip_preprocess_text,
) )
@@ -15,7 +15,7 @@ from sglang.multimodal_gen.configs.models.encoders import (
LlamaConfig, LlamaConfig,
) )
from sglang.multimodal_gen.configs.models.vaes import HunyuanVAEConfig from sglang.multimodal_gen.configs.models.vaes import HunyuanVAEConfig
from sglang.multimodal_gen.configs.pipelines.base import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs import PipelineConfig
PROMPT_TEMPLATE_ENCODE_VIDEO = ( PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: " "<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
@@ -9,7 +9,10 @@ from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAECo
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig
from sglang.multimodal_gen.configs.pipelines.base import ModelTaskType, PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType,
PipelineConfig,
)
from sglang.multimodal_gen.utils import calculate_dimensions from sglang.multimodal_gen.utils import calculate_dimensions
@@ -6,7 +6,7 @@ from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig
from sglang.multimodal_gen.configs.models.dits import StepVideoConfig from sglang.multimodal_gen.configs.models.dits import StepVideoConfig
from sglang.multimodal_gen.configs.models.vaes import StepVideoVAEConfig from sglang.multimodal_gen.configs.models.vaes import StepVideoVAEConfig
from sglang.multimodal_gen.configs.pipelines.base import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
@dataclass @dataclass
@@ -14,7 +14,10 @@ from sglang.multimodal_gen.configs.models.encoders import (
T5Config, T5Config,
) )
from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig
from sglang.multimodal_gen.configs.pipelines.base import ModelTaskType, PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType,
PipelineConfig,
)
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__)
+27 -39
View File
@@ -15,7 +15,7 @@ import re
from functools import lru_cache from functools import lru_cache
from typing import Any, Callable, Dict, List, Optional, Tuple, Type from typing import Any, Callable, Dict, List, Optional, Tuple, Type
from sglang.multimodal_gen.configs.pipelines import ( from sglang.multimodal_gen.configs.pipeline_configs import (
FastHunyuanConfig, FastHunyuanConfig,
FluxPipelineConfig, FluxPipelineConfig,
HunyuanConfig, HunyuanConfig,
@@ -25,12 +25,12 @@ from sglang.multimodal_gen.configs.pipelines import (
WanT2V480PConfig, WanT2V480PConfig,
WanT2V720PConfig, WanT2V720PConfig,
) )
from sglang.multimodal_gen.configs.pipelines.base import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
from sglang.multimodal_gen.configs.pipelines.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
from sglang.multimodal_gen.configs.pipelines.wan import ( from sglang.multimodal_gen.configs.pipeline_configs.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_I2V_A14B_Config,
@@ -55,7 +55,7 @@ from sglang.multimodal_gen.configs.sample.wan import (
WanT2V_1_3B_SamplingParams, WanT2V_1_3B_SamplingParams,
WanT2V_14B_SamplingParams, WanT2V_14B_SamplingParams,
) )
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
@@ -74,49 +74,37 @@ _PIPELINE_REGISTRY: Dict[str, Type[ComposedPipelineBase]] = {}
def _discover_and_register_pipelines(): def _discover_and_register_pipelines():
""" """
Automatically discover and register all ComposedPipelineBase subclasses. Automatically discover and register all ComposedPipelineBase subclasses.
This function scans the 'sglang.multimodal_gen.runtime.architectures' package, This function scans the 'sglang.multimodal_gen.runtime.pipelines' package,
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: # E-run only once if _PIPELINE_REGISTRY: # run only once
return return
package_name = "sglang.multimodal_gen.runtime.architectures" package_name = "sglang.multimodal_gen.runtime.pipelines"
package = importlib.import_module(package_name) package = importlib.import_module(package_name)
for _, pipeline_type_str, ispkg in pkgutil.iter_modules(package.__path__): for _, module_name, ispkg in pkgutil.walk_packages(
package.__path__, package.__name__ + "."
):
if not ispkg: if not ispkg:
continue pipeline_module = importlib.import_module(module_name)
pipeline_type_package_name = f"{package_name}.{pipeline_type_str}" if hasattr(pipeline_module, "EntryClass"):
pipeline_type_package = importlib.import_module(pipeline_type_package_name) entry_cls = pipeline_module.EntryClass
for _, arch, ispkg_arch in pkgutil.iter_modules(pipeline_type_package.__path__): entry_cls_list = (
if not ispkg_arch: [entry_cls] if not isinstance(entry_cls, list) else entry_cls
continue )
arch_package_name = f"{pipeline_type_package_name}.{arch}"
arch_package = importlib.import_module(arch_package_name)
for _, module_name, ispkg_module in pkgutil.walk_packages(
arch_package.__path__, arch_package.__name__ + "."
):
if not ispkg_module:
pipeline_module = importlib.import_module(module_name)
if hasattr(pipeline_module, "EntryClass"):
entry_cls = pipeline_module.EntryClass
if not isinstance(entry_cls, list):
entry_cls_list = [entry_cls]
else:
entry_cls_list = entry_cls
for cls in entry_cls_list: for cls in entry_cls_list:
if hasattr(cls, "pipeline_name"): if hasattr(cls, "pipeline_name"):
if cls.pipeline_name in _PIPELINE_REGISTRY: if cls.pipeline_name in _PIPELINE_REGISTRY:
logger.warning( logger.warning(
f"Duplicate pipeline name '{cls.pipeline_name}' found. Overwriting." f"Duplicate pipeline name '{cls.pipeline_name}' found. Overwriting."
) )
_PIPELINE_REGISTRY[cls.pipeline_name] = cls _PIPELINE_REGISTRY[cls.pipeline_name] = cls
# else: logger.debug(
# logger.warning( f"Registering pipelines complete, {len(_PIPELINE_REGISTRY)} pipelines registered"
# f"Pipeline class {cls.__name__} does not have a 'pipeline_name' attribute." )
# )
# --- Part 2: Config Registration --- # --- Part 2: Config Registration ---
@@ -1,8 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Basic inference pipelines for sglang.multimodal_gen.
This package contains basic pipelines for video and image generation.
"""
@@ -1 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
@@ -1 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
@@ -1 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
@@ -1 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
@@ -1 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
@@ -21,8 +21,8 @@ import torch
import torchvision import torchvision
from einops import rearrange from einops import rearrange
from sglang.multimodal_gen.runtime.pipelines import Req from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
# Suppress verbose logging from imageio, which is triggered when saving images. # Suppress verbose logging from imageio, which is triggered when saving images.
logging.getLogger("imageio").setLevel(logging.WARNING) logging.getLogger("imageio").setLevel(logging.WARNING)
@@ -24,7 +24,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
post_process_sample, post_process_sample,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.scheduler_client import scheduler_client from sglang.multimodal_gen.runtime.scheduler_client import scheduler_client
from sglang.multimodal_gen.runtime.server_args import get_global_server_args from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -34,7 +34,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
post_process_sample, post_process_sample,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import get_global_server_args from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -16,7 +16,7 @@ logging.getLogger("imageio").setLevel(logging.WARNING)
logging.getLogger("imageio_ffmpeg").setLevel(logging.WARNING) logging.getLogger("imageio_ffmpeg").setLevel(logging.WARNING)
from sglang.multimodal_gen.configs.sample.base import DataType, SamplingParams from sglang.multimodal_gen.configs.sample.base import DataType, SamplingParams
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import shallow_asdict from sglang.multimodal_gen.utils import shallow_asdict
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.multimodal_gen.runtime.layers.attention import AttentionMetadata from sglang.multimodal_gen.runtime.layers.attention import AttentionMetadata
from sglang.multimodal_gen.runtime.pipelines import Req from sglang.multimodal_gen.runtime.pipelines_core import Req
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -16,8 +16,8 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_cfg_group, get_cfg_group,
get_tp_group, get_tp_group,
) )
from sglang.multimodal_gen.runtime.pipelines import Req, build_pipeline from sglang.multimodal_gen.runtime.pipelines_core import Req, build_pipeline
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs
from sglang.multimodal_gen.runtime.utils.common import set_cuda_arch from sglang.multimodal_gen.runtime.utils.common import set_cuda_arch
from sglang.multimodal_gen.runtime.utils.logging_utils import ( from sglang.multimodal_gen.runtime.utils.logging_utils import (
@@ -6,7 +6,7 @@ from typing import Any
import zmq import zmq
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.server_args import ( from sglang.multimodal_gen.runtime.server_args import (
PortArgs, PortArgs,
ServerArgs, ServerArgs,
@@ -6,8 +6,8 @@ from typing import TypeVar
import zmq import zmq
from sglang.multimodal_gen.runtime.pipelines import Req from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.utils import init_logger from sglang.multimodal_gen.utils import init_logger
@@ -1,64 +1 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Diffusion pipelines for sglang.multimodal_gen.
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.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class PipelineWithLoRA(LoRAPipeline, ComposedPipelineBase):
"""Type for a pipeline that has both ComposedPipelineBase and LoRAPipeline functionality."""
pass
def build_pipeline(
server_args: ServerArgs,
) -> PipelineWithLoRA:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
1. download the model from the hub if it's not already downloaded
2. verify the model config and directory
3. based on the config, determine the pipeline class
"""
model_path = server_args.model_path
model_info = get_model_info(model_path)
if model_info is None:
raise ValueError(f"Unsupported model: {model_path}")
pipeline_cls = model_info.pipeline_cls
# instantiate the pipelines
pipeline = pipeline_cls(model_path, server_args)
logger.info("Pipelines instantiated")
return cast(PipelineWithLoRA, pipeline)
__all__ = [
"build_pipeline",
"ComposedPipelineBase",
"Req",
"LoRAPipeline",
]
@@ -1,11 +1,11 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
DenoisingStage, DenoisingStage,
@@ -9,10 +9,10 @@ using the modular pipeline architecture.
""" """
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
DenoisingStage, DenoisingStage,
@@ -3,11 +3,11 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from diffusers.image_processor import VaeImageProcessor from diffusers.image_processor import VaeImageProcessor
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
DecodingStage, DecodingStage,
DenoisingStage, DenoisingStage,
ImageEncodingStage, ImageEncodingStage,
@@ -17,7 +17,7 @@ from sglang.multimodal_gen.runtime.pipelines.stages import (
TextEncodingStage, TextEncodingStage,
TimestepPreparationStage, TimestepPreparationStage,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.conditioning import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.conditioning import (
ConditioningStage, ConditioningStage,
) )
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
@@ -1,59 +0,0 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Pipeline stages for diffusion models.
This package contains the various stages that can be composed to create
complete diffusion pipelines.
"""
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.causal_denoising import (
CausalDMDDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.conditioning import (
ConditioningStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.decoding import DecodingStage
from sglang.multimodal_gen.runtime.pipelines.stages.denoising import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines.stages.denoising_dmd import (
DmdDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.encoding import EncodingStage
from sglang.multimodal_gen.runtime.pipelines.stages.image_encoding import (
ImageEncodingStage,
ImageVAEEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.input_validation import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.latent_preparation import (
LatentPreparationStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.text_encoding import (
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines.stages.timestep_preparation import (
TimestepPreparationStage,
)
__all__ = [
"PipelineStage",
"InputValidationStage",
"TimestepPreparationStage",
"LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"ImageVAEEncodingStage",
"TextEncodingStage",
"StepvideoPromptEncodingStage",
]
@@ -24,11 +24,11 @@ from sglang.multimodal_gen.runtime.models.encoders.bert import (
HunyuanClip, # type: ignore HunyuanClip, # type: ignore
) )
from sglang.multimodal_gen.runtime.models.encoders.stepllm import STEP1TextEncoder from sglang.multimodal_gen.runtime.models.encoders.stepllm import STEP1TextEncoder
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
DecodingStage, DecodingStage,
DenoisingStage, DenoisingStage,
InputValidationStage, InputValidationStage,
@@ -7,13 +7,13 @@ Wan causal DMD pipeline implementation.
This module wires the causal DMD denoising stage into the modular pipeline. This module wires the causal DMD denoising stage into the modular pipeline.
""" """
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
# isort: off # isort: off
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
CausalDMDDenoisingStage, CausalDMDDenoisingStage,
@@ -11,15 +11,15 @@ using the modular pipeline architecture.
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import ( from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, FlowMatchEulerDiscreteScheduler,
) )
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# isort: off # isort: off
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
DmdDenoisingStage, DmdDenoisingStage,
@@ -8,15 +8,15 @@ This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture. using the modular pipeline architecture.
""" """
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# isort: off # isort: off
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ImageEncodingStage, ImageEncodingStage,
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
@@ -8,15 +8,15 @@ This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture. using the modular pipeline architecture.
""" """
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# isort: off # isort: off
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ImageEncodingStage, ImageEncodingStage,
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
@@ -11,11 +11,11 @@ using the modular pipeline architecture.
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import ( from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler, FlowUniPCMultistepScheduler,
) )
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines.lora_pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ConditioningStage, ConditioningStage,
DecodingStage, DecodingStage,
DenoisingStage, DenoisingStage,
@@ -0,0 +1,64 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Diffusion pipelines for sglang.multimodal_gen.
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,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class PipelineWithLoRA(LoRAPipeline, ComposedPipelineBase):
"""Type for a pipeline that has both ComposedPipelineBase and LoRAPipeline functionality."""
pass
def build_pipeline(
server_args: ServerArgs,
) -> PipelineWithLoRA:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
1. download the model from the hub if it's not already downloaded
2. verify the model config and directory
3. based on the config, determine the pipeline class
"""
model_path = server_args.model_path
model_info = get_model_info(model_path)
if model_info is None:
raise ValueError(f"Unsupported model: {model_path}")
pipeline_cls = model_info.pipeline_cls
# instantiate the pipelines
pipeline = pipeline_cls(model_path, server_args)
logger.info("Pipelines instantiated")
return cast(PipelineWithLoRA, pipeline)
__all__ = [
"build_pipeline",
"ComposedPipelineBase",
"Req",
"LoRAPipeline",
]
@@ -15,15 +15,15 @@ from typing import Any, cast
import torch import torch
from tqdm import tqdm from tqdm import tqdm
from sglang.multimodal_gen.configs.pipelines import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs import PipelineConfig
from sglang.multimodal_gen.runtime.loader.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loader import (
PipelineComponentLoader, PipelineComponentLoader,
) )
from sglang.multimodal_gen.runtime.pipelines.executors.pipeline_executor import ( from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
PipelineExecutor, PipelineExecutor,
) )
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model, maybe_download_model,
@@ -90,7 +90,7 @@ class ComposedPipelineBase(ABC):
def build_executor(self, server_args: ServerArgs): def build_executor(self, server_args: ServerArgs):
# TODO # TODO
from sglang.multimodal_gen.runtime.pipelines.executors.parallel_executor import ( from sglang.multimodal_gen.runtime.pipelines_core.executors.parallel_executor import (
ParallelExecutor, ParallelExecutor,
) )
@@ -9,12 +9,12 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_cfg_group, get_cfg_group,
get_classifier_free_guidance_rank, get_classifier_free_guidance_rank,
) )
from sglang.multimodal_gen.runtime.pipelines import Req from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.pipelines.executors.pipeline_executor import ( from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
PipelineExecutor, PipelineExecutor,
Timer, Timer,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
PipelineStage, PipelineStage,
StageParallelismType, StageParallelismType,
) )
@@ -8,8 +8,8 @@ import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import List from typing import List
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -6,13 +6,13 @@ Synchronous pipeline executor implementation.
""" """
from typing import List from typing import List
from sglang.multimodal_gen.runtime.pipelines.executors.pipeline_executor import ( from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
PipelineExecutor, PipelineExecutor,
Timer, Timer,
logger, logger,
) )
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
@@ -16,7 +16,7 @@ from sglang.multimodal_gen.runtime.layers.lora.linear import (
replace_submodule, replace_submodule,
) )
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
from sglang.multimodal_gen.runtime.pipelines.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
@@ -0,0 +1,59 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Pipeline stages for diffusion models.
This package contains the various stages that can be composed to create
complete diffusion pipelines.
"""
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import (
CausalDMDDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.conditioning import (
ConditioningStage,
)
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.denoising_dmd import (
DmdDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.encoding import EncodingStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
ImageEncodingStage,
ImageVAEEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import (
LatentPreparationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.timestep_preparation import (
TimestepPreparationStage,
)
__all__ = [
"PipelineStage",
"InputValidationStage",
"TimestepPreparationStage",
"LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"ImageVAEEncodingStage",
"TextEncodingStage",
"StepvideoPromptEncodingStage",
]
@@ -16,8 +16,10 @@ from enum import Enum, auto
import torch import torch
import sglang.multimodal_gen.envs as envs import sglang.multimodal_gen.envs as envs
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -5,12 +5,14 @@ import torch # type: ignore
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.denoising import DenoisingStage from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -7,12 +7,14 @@ Conditioning stage for diffusion pipelines.
import torch import torch
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -10,19 +10,21 @@ import weakref
import torch import torch
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig
from sglang.multimodal_gen.configs.pipelines.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loader import VAELoader from sglang.multimodal_gen.runtime.loader.component_loader import VAELoader
from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import OutputBatch, Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
PipelineStage, PipelineStage,
StageParallelismType, StageParallelismType,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
@@ -19,7 +19,7 @@ import torch.profiler
from einops import rearrange from einops import rearrange
from tqdm.auto import tqdm from tqdm.auto import tqdm
from sglang.multimodal_gen.configs.pipelines.base import ModelTaskType, STA_Mode from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
cfg_model_parallel_all_reduce, cfg_model_parallel_all_reduce,
get_local_torch_device, get_local_torch_device,
@@ -44,15 +44,17 @@ from sglang.multimodal_gen.runtime.layers.attention.STA_configuration import (
) )
from sglang.multimodal_gen.runtime.loader.component_loader import TransformerLoader from sglang.multimodal_gen.runtime.loader.component_loader import TransformerLoader
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
PipelineStage, PipelineStage,
StageParallelismType, StageParallelismType,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -22,9 +22,9 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler
FlowMatchEulerDiscreteScheduler, FlowMatchEulerDiscreteScheduler,
) )
from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages import DenoisingStage from sglang.multimodal_gen.runtime.pipelines_core.stages import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines.stages.denoising import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
st_attn_available, st_attn_available,
vsa_available, vsa_available,
) )
@@ -9,12 +9,14 @@ import torch
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
V, # Import validators V, # Import validators
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
@@ -10,7 +10,7 @@ This module contains implementations of image encoding stages for diffusion pipe
import PIL import PIL
import torch import torch
from sglang.multimodal_gen.configs.pipelines.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
_pack_latents, _pack_latents,
@@ -25,12 +25,14 @@ from sglang.multimodal_gen.runtime.models.vision_utils import (
pil_to_numpy, pil_to_numpy,
resize, resize,
) )
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ExecutionMode, ServerArgs from sglang.multimodal_gen.runtime.server_args import ExecutionMode, ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
@@ -9,15 +9,15 @@ import torch
import torchvision.transforms.functional as TF import torchvision.transforms.functional as TF
from PIL import Image from PIL import Image
from sglang.multimodal_gen.configs.pipelines import WanI2V480PConfig from sglang.multimodal_gen.configs.pipeline_configs import WanI2V480PConfig
from sglang.multimodal_gen.configs.pipelines.base import ModelTaskType from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
from sglang.multimodal_gen.configs.pipelines.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
) )
from sglang.multimodal_gen.runtime.models.vision_utils import load_image, load_video from sglang.multimodal_gen.runtime.models.vision_utils import load_image, load_video
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators, StageValidators,
VerificationResult, VerificationResult,
) )
@@ -7,12 +7,14 @@ Latent preparation stage for diffusion pipelines.
from diffusers.utils.torch_utils import randn_tensor from diffusers.utils.torch_utils import randn_tensor
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -5,12 +5,14 @@
import torch import torch
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -10,15 +10,17 @@ This module contains implementations of prompt encoding stages for diffusion pip
import torch import torch
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
from sglang.multimodal_gen.configs.pipelines import FluxPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs import FluxPipelineConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -12,21 +12,23 @@ from typing import Any, Callable, Tuple
import numpy as np import numpy as np
from sglang.multimodal_gen.configs.pipelines import FluxPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs import FluxPipelineConfig
from sglang.multimodal_gen.configs.pipelines.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
PipelineStage, PipelineStage,
StageParallelismType, StageParallelismType,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
from sglang.multimodal_gen.runtime.pipelines.stages.validators import VerificationResult from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -5,7 +5,7 @@ import asyncio
import zmq import zmq
import zmq.asyncio import zmq.asyncio
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -15,9 +15,9 @@ from dataclasses import field
from enum import Enum from enum import Enum
from typing import Any, Optional from typing import Any, Optional
from sglang.multimodal_gen.configs.pipelines import FluxPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs import FluxPipelineConfig
from sglang.multimodal_gen.configs.pipelines.base import PipelineConfig, STA_Mode from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig, STA_Mode
from sglang.multimodal_gen.configs.pipelines.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
@@ -2,7 +2,7 @@
import zmq import zmq
from sglang.multimodal_gen.runtime.pipelines.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger