[diffusion] chore: clean ComposedPipelineBase (#15937)

This commit is contained in:
Mick
2025-12-28 11:43:25 +08:00
committed by GitHub
parent f55d608c89
commit b4a00ed2d9
@@ -7,7 +7,6 @@ Base class for composed pipelines.
This module defines the base class for pipelines that are composed of multiple stages. This module defines the base class for pipelines that are composed of multiple stages.
""" """
import argparse
import os import os
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, cast from typing import Any, cast
@@ -15,7 +14,6 @@ from typing import Any, cast
import torch import torch
from tqdm import tqdm from tqdm import tqdm
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,
) )
@@ -49,7 +47,6 @@ class ComposedPipelineBase(ABC):
_extra_config_module_map: dict[str, str] = {} _extra_config_module_map: dict[str, str] = {}
server_args: ServerArgs | None = None server_args: ServerArgs | None = None
modules: dict[str, Any] = {} modules: dict[str, Any] = {}
post_init_called: bool = False
executor: PipelineExecutor | None = None executor: PipelineExecutor | None = None
# the name of the pipeline it associated with, in diffusers # the name of the pipeline it associated with, in diffusers
@@ -92,6 +89,8 @@ class ComposedPipelineBase(ABC):
logger.info("Loading pipeline modules...") logger.info("Loading pipeline modules...")
self.modules = self.load_modules(server_args, loaded_modules) self.modules = self.load_modules(server_args, loaded_modules)
self.__post_init__()
def build_executor(self, server_args: ServerArgs): def build_executor(self, server_args: ServerArgs):
# TODO # TODO
from sglang.multimodal_gen.runtime.pipelines_core.executors.parallel_executor import ( from sglang.multimodal_gen.runtime.pipelines_core.executors.parallel_executor import (
@@ -101,48 +100,13 @@ class ComposedPipelineBase(ABC):
# return SyncExecutor(server_args=server_args) # return SyncExecutor(server_args=server_args)
return ParallelExecutor(server_args=server_args) return ParallelExecutor(server_args=server_args)
def post_init(self) -> None: def __post_init__(self) -> None:
assert self.server_args is not None, "server_args must be set" assert self.server_args is not None, "server_args must be set"
if self.post_init_called:
return
self.post_init_called = True
self.initialize_pipeline(self.server_args) self.initialize_pipeline(self.server_args)
logger.info("Creating pipeline stages...") logger.info("Creating pipeline stages...")
self.create_pipeline_stages(self.server_args) self.create_pipeline_stages(self.server_args)
@classmethod
def from_pretrained(
cls,
model_path: str,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: str | PipelineConfig | None = None,
args: argparse.Namespace | None = None,
required_config_modules: list[str] | None = None,
loaded_modules: dict[str, torch.nn.Module] | None = None,
**kwargs,
) -> "ComposedPipelineBase":
"""
Load a pipeline from a pretrained model.
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
kwargs["model_path"] = model_path
server_args = ServerArgs.from_kwargs(**kwargs)
logger.info("server_args in from_pretrained: %s", server_args)
pipe = cls(
model_path,
server_args,
required_config_modules=required_config_modules,
loaded_modules=loaded_modules,
)
pipe.post_init()
return pipe
def get_module(self, module_name: str, default_value: Any = None) -> Any: def get_module(self, module_name: str, default_value: Any = None) -> Any:
if module_name not in self.modules: if module_name not in self.modules:
return default_value return default_value
@@ -357,8 +321,6 @@ class ComposedPipelineBase(ABC):
Returns: Returns:
Req: The batch with the generated video or image. Req: The batch with the generated video or image.
""" """
if not self.post_init_called:
self.post_init()
if self.is_lora_set() and not self.is_lora_effective(): if self.is_lora_set() and not self.is_lora_effective():
logger.warning( logger.warning(