[diffusion] chore: clean ComposedPipelineBase (#15937)
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user