[diffusion] feat: support role-based component loading and stage affinity (#25168)
This commit is contained in:
@@ -25,7 +25,7 @@ class RoleType(str, Enum):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def choices(cls) -> list[str]:
|
def choices(cls) -> list[str]:
|
||||||
return [role.value for role in cls]
|
return [role.value for role in cls] + sorted(_ROLE_ALIASES)
|
||||||
|
|
||||||
|
|
||||||
def get_module_role(module_name: str) -> "RoleType | None":
|
def get_module_role(module_name: str) -> "RoleType | None":
|
||||||
@@ -37,32 +37,53 @@ def get_module_role(module_name: str) -> "RoleType | None":
|
|||||||
"image_processor",
|
"image_processor",
|
||||||
"processor",
|
"processor",
|
||||||
"connectors",
|
"connectors",
|
||||||
|
"vision_language_encoder",
|
||||||
)
|
)
|
||||||
if any(
|
if any(
|
||||||
module_name == p or module_name.startswith(p + "_") for p in encoder_prefixes
|
module_name == p or module_name.startswith(p + "_") for p in encoder_prefixes
|
||||||
):
|
):
|
||||||
return RoleType.ENCODER
|
return RoleType.ENCODER
|
||||||
|
|
||||||
denoising_prefixes = ("transformer",)
|
if module_name in {"hy3dshape_conditioner", "hy3dshape_image_processor"}:
|
||||||
|
return RoleType.ENCODER
|
||||||
|
|
||||||
|
denoising_prefixes = (
|
||||||
|
"transformer",
|
||||||
|
"video_dit",
|
||||||
|
"audio_dit",
|
||||||
|
"dual_tower_bridge",
|
||||||
|
)
|
||||||
if any(
|
if any(
|
||||||
module_name == p or module_name.startswith(p + "_") for p in denoising_prefixes
|
module_name == p or module_name.startswith(p + "_") for p in denoising_prefixes
|
||||||
):
|
):
|
||||||
return RoleType.DENOISER
|
return RoleType.DENOISER
|
||||||
|
|
||||||
|
if module_name == "hy3dshape_model":
|
||||||
|
return RoleType.DENOISER
|
||||||
|
|
||||||
decoder_prefixes = ("vae", "audio_vae", "video_vae", "vocoder")
|
decoder_prefixes = ("vae", "audio_vae", "video_vae", "vocoder")
|
||||||
if any(
|
if any(
|
||||||
module_name == p or module_name.startswith(p + "_") for p in decoder_prefixes
|
module_name == p or module_name.startswith(p + "_") for p in decoder_prefixes
|
||||||
):
|
):
|
||||||
return RoleType.DECODER
|
return RoleType.DECODER
|
||||||
|
|
||||||
|
if module_name == "hy3dshape_vae":
|
||||||
|
return RoleType.DECODER
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def filter_modules_for_role(module_names: list[str], role: "RoleType") -> list[str]:
|
def filter_modules_for_role(
|
||||||
|
module_names: list[str],
|
||||||
|
role: "RoleType",
|
||||||
|
*,
|
||||||
|
extra_allowed_modules: set[str] | None = None,
|
||||||
|
) -> list[str]:
|
||||||
"""Filter module names to only those needed by the given role."""
|
"""Filter module names to only those needed by the given role."""
|
||||||
if role in (RoleType.MONOLITHIC, RoleType.SERVER):
|
if role in (RoleType.MONOLITHIC, RoleType.SERVER):
|
||||||
return module_names
|
return module_names
|
||||||
|
|
||||||
|
extra_allowed_modules = extra_allowed_modules or set()
|
||||||
filtered = []
|
filtered = []
|
||||||
for name in module_names:
|
for name in module_names:
|
||||||
module_role = get_module_role(name)
|
module_role = get_module_role(name)
|
||||||
@@ -71,8 +92,7 @@ def filter_modules_for_role(module_names: list[str], role: "RoleType") -> list[s
|
|||||||
filtered.append(name)
|
filtered.append(name)
|
||||||
elif module_role == role:
|
elif module_role == role:
|
||||||
filtered.append(name)
|
filtered.append(name)
|
||||||
elif role == RoleType.ENCODER and module_role == RoleType.DECODER:
|
elif name in extra_allowed_modules:
|
||||||
# Encoder also needs VAE for ImageVAEEncoding stages
|
|
||||||
filtered.append(name)
|
filtered.append(name)
|
||||||
|
|
||||||
return filtered
|
return filtered
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ import torch.nn as nn
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
|
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
|
||||||
Hunyuan3D2PipelineConfig,
|
Hunyuan3D2PipelineConfig,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
|
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
|
||||||
load_model_from_full_model_state_dict,
|
load_model_from_full_model_state_dict,
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
@@ -58,6 +59,21 @@ class Hunyuan3D2Pipeline(ComposedPipelineBase):
|
|||||||
"hy3dshape_image_processor",
|
"hy3dshape_image_processor",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
def validate_disagg_role(self, role: RoleType) -> None:
|
||||||
|
if role == RoleType.MONOLITHIC:
|
||||||
|
return
|
||||||
|
config = self.server_args.pipeline_config
|
||||||
|
if not isinstance(config, Hunyuan3D2PipelineConfig):
|
||||||
|
raise TypeError(
|
||||||
|
"Hunyuan3D2Pipeline requires Hunyuan3D2PipelineConfig, "
|
||||||
|
f"got {type(config)}"
|
||||||
|
)
|
||||||
|
if config.paint_enable:
|
||||||
|
raise ValueError(
|
||||||
|
"Hunyuan3D2Pipeline only supports shape-only disaggregation. "
|
||||||
|
"Disable paint_enable when launching encoder/denoiser/decoder roles."
|
||||||
|
)
|
||||||
|
|
||||||
def _load_config(self) -> dict[str, Any]:
|
def _load_config(self) -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"_class_name": self.pipeline_name,
|
"_class_name": self.pipeline_name,
|
||||||
@@ -357,6 +373,8 @@ class Hunyuan3D2Pipeline(ComposedPipelineBase):
|
|||||||
def create_pipeline_stages(self, server_args: ServerArgs):
|
def create_pipeline_stages(self, server_args: ServerArgs):
|
||||||
config = server_args.pipeline_config
|
config = server_args.pipeline_config
|
||||||
assert isinstance(config, Hunyuan3D2PipelineConfig)
|
assert isinstance(config, Hunyuan3D2PipelineConfig)
|
||||||
|
latent_shape = tuple(config.vae_config.arch_config.latent_shape)
|
||||||
|
guidance_embed = bool(config.dit_config.arch_config.guidance_embed)
|
||||||
|
|
||||||
# Shape: 4 stages
|
# Shape: 4 stages
|
||||||
self.add_stage(
|
self.add_stage(
|
||||||
@@ -364,10 +382,10 @@ class Hunyuan3D2Pipeline(ComposedPipelineBase):
|
|||||||
stage=Hunyuan3DShapeBeforeDenoisingStage(
|
stage=Hunyuan3DShapeBeforeDenoisingStage(
|
||||||
image_processor=self.get_module("hy3dshape_image_processor"),
|
image_processor=self.get_module("hy3dshape_image_processor"),
|
||||||
conditioner=self.get_module("hy3dshape_conditioner"),
|
conditioner=self.get_module("hy3dshape_conditioner"),
|
||||||
vae=self.get_module("hy3dshape_vae"),
|
|
||||||
model=self.get_module("hy3dshape_model"),
|
|
||||||
scheduler=self.get_module("hy3dshape_scheduler"),
|
scheduler=self.get_module("hy3dshape_scheduler"),
|
||||||
config=config,
|
config=config,
|
||||||
|
latent_shape=latent_shape,
|
||||||
|
guidance_embed=guidance_embed,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
self.add_stage(
|
self.add_stage(
|
||||||
|
|||||||
@@ -63,7 +63,12 @@ class MOVAPipeline(ComposedPipelineBase):
|
|||||||
self.add_stage(InputValidationStage())
|
self.add_stage(InputValidationStage())
|
||||||
self.add_standard_text_encoding_stage()
|
self.add_standard_text_encoding_stage()
|
||||||
if getattr(self.get_module("video_dit"), "require_vae_embedding", True):
|
if getattr(self.get_module("video_dit"), "require_vae_embedding", True):
|
||||||
self.add_stage(ImageVAEEncodingStage(vae=self.get_module("video_vae")))
|
self.add_stage(
|
||||||
|
ImageVAEEncodingStage(
|
||||||
|
vae=self.get_module("video_vae"),
|
||||||
|
component_name="video_vae",
|
||||||
|
)
|
||||||
|
)
|
||||||
self.add_stage(
|
self.add_stage(
|
||||||
MOVALatentPreparationStage(
|
MOVALatentPreparationStage(
|
||||||
audio_vae=self.get_module("audio_vae"),
|
audio_vae=self.get_module("audio_vae"),
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
# 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.disaggregation.roles import RoleType
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
|
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
ComposedPipelineBase,
|
ComposedPipelineBase,
|
||||||
@@ -115,9 +116,10 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
|
|||||||
]
|
]
|
||||||
|
|
||||||
def create_pipeline_stages(self, server_args: ServerArgs):
|
def create_pipeline_stages(self, server_args: ServerArgs):
|
||||||
self.add_stage(
|
def create_before_denoising_stage():
|
||||||
QwenImageLayeredBeforeDenoisingStage(
|
return QwenImageLayeredBeforeDenoisingStage(
|
||||||
vae=self.get_module("vae"),
|
vae=self.get_module("vae"),
|
||||||
|
text_encoder=None,
|
||||||
tokenizer=self.get_module("tokenizer"),
|
tokenizer=self.get_module("tokenizer"),
|
||||||
processor=self.get_module("processor"),
|
processor=self.get_module("processor"),
|
||||||
transformer=self.get_module("transformer"),
|
transformer=self.get_module("transformer"),
|
||||||
@@ -128,6 +130,11 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
|
|||||||
server_args.pipeline_config.text_encoder_precisions[0]
|
server_args.pipeline_config.text_encoder_precisions[0]
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.add_stage_factory(
|
||||||
|
RoleType.ENCODER,
|
||||||
|
create_before_denoising_stage,
|
||||||
|
"QwenImageLayeredBeforeDenoisingStage",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.add_standard_timestep_preparation_stage(
|
self.add_standard_timestep_preparation_stage(
|
||||||
|
|||||||
@@ -99,6 +99,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
"""
|
"""
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self._disagg_role = server_args.disagg_role
|
self._disagg_role = server_args.disagg_role
|
||||||
|
self.validate_disagg_role(self._disagg_role)
|
||||||
|
|
||||||
self.model_path: str = model_path
|
self.model_path: str = model_path
|
||||||
self._stages: list[PipelineStage] = []
|
self._stages: list[PipelineStage] = []
|
||||||
@@ -107,17 +108,26 @@ class ComposedPipelineBase(ABC):
|
|||||||
self.executor = executor or self.build_executor(server_args=server_args)
|
self.executor = executor or self.build_executor(server_args=server_args)
|
||||||
self.component_residency_manager: ComponentResidencyManager | None = None
|
self.component_residency_manager: ComponentResidencyManager | None = None
|
||||||
|
|
||||||
if required_config_modules is not None:
|
base_required_config_modules = (
|
||||||
self._required_config_modules = required_config_modules
|
required_config_modules
|
||||||
|
if required_config_modules is not None
|
||||||
if self._required_config_modules is None:
|
else self._required_config_modules
|
||||||
|
)
|
||||||
|
if base_required_config_modules is None:
|
||||||
raise NotImplementedError("Subclass must set _required_config_modules")
|
raise NotImplementedError("Subclass must set _required_config_modules")
|
||||||
|
self._required_config_modules = list(base_required_config_modules)
|
||||||
|
self._extra_config_module_map = dict(self._extra_config_module_map)
|
||||||
|
|
||||||
# Filter modules based on disaggregation role
|
# Filter modules based on disaggregation role
|
||||||
if self._disagg_role != RoleType.MONOLITHIC:
|
if self._disagg_role != RoleType.MONOLITHIC:
|
||||||
original_modules = list(self._required_config_modules)
|
original_modules = list(self._required_config_modules)
|
||||||
|
task_name = self.server_args.pipeline_config.task_type.name.lower()
|
||||||
self._required_config_modules = filter_modules_for_role(
|
self._required_config_modules = filter_modules_for_role(
|
||||||
self._required_config_modules, self._disagg_role
|
self._required_config_modules,
|
||||||
|
self._disagg_role,
|
||||||
|
extra_allowed_modules=self._get_extra_allowed_modules_for_role(
|
||||||
|
self._disagg_role, task_name
|
||||||
|
),
|
||||||
)
|
)
|
||||||
skipped = set(original_modules) - set(self._required_config_modules)
|
skipped = set(original_modules) - set(self._required_config_modules)
|
||||||
if skipped:
|
if skipped:
|
||||||
@@ -202,6 +212,44 @@ class ComposedPipelineBase(ABC):
|
|||||||
"""
|
"""
|
||||||
return
|
return
|
||||||
|
|
||||||
|
def validate_disagg_role(self, role: RoleType) -> None:
|
||||||
|
"""Validate whether the requested disaggregation role is supported."""
|
||||||
|
return
|
||||||
|
|
||||||
|
def _get_extra_allowed_modules_for_role(
|
||||||
|
self, role: RoleType, task_name: str
|
||||||
|
) -> set[str]:
|
||||||
|
role_to_pipeline_modules: dict[RoleType, dict[str, set[str]]] = {
|
||||||
|
RoleType.ENCODER: {
|
||||||
|
"Flux2Pipeline": {"vae"},
|
||||||
|
"Flux2KleinPipeline": {"vae"},
|
||||||
|
"QwenImageEditPipeline": {"vae"},
|
||||||
|
"QwenImageEditPlusPipeline": {"vae"},
|
||||||
|
"QwenImageLayeredPipeline": {"vae", "transformer"},
|
||||||
|
"GlmImagePipeline": {"vae", "transformer"},
|
||||||
|
"WanImageToVideoPipeline": {"vae"},
|
||||||
|
"WanImageToVideoDmdPipeline": {"vae"},
|
||||||
|
"MOVA": {"video_vae", "audio_vae"},
|
||||||
|
"MOVAPipeline": {"video_vae", "audio_vae"},
|
||||||
|
},
|
||||||
|
RoleType.DENOISER: {},
|
||||||
|
RoleType.DECODER: {},
|
||||||
|
}
|
||||||
|
extra_allowed_modules = set(
|
||||||
|
role_to_pipeline_modules.get(role, {}).get(self.pipeline_name, set())
|
||||||
|
)
|
||||||
|
|
||||||
|
if role == RoleType.DENOISER and task_name == "ti2v":
|
||||||
|
if self.pipeline_name in {
|
||||||
|
"WanImageToVideoPipeline",
|
||||||
|
"WanImageToVideoDmdPipeline",
|
||||||
|
}:
|
||||||
|
extra_allowed_modules.add("vae")
|
||||||
|
elif self.pipeline_name == "LTX2Pipeline":
|
||||||
|
extra_allowed_modules.update({"vae", "audio_vae"})
|
||||||
|
|
||||||
|
return extra_allowed_modules
|
||||||
|
|
||||||
# --- Config-name → pipeline_config attribute mapping ---
|
# --- Config-name → pipeline_config attribute mapping ---
|
||||||
_CONFIG_ATTR_MAP: dict[str, str] = {
|
_CONFIG_ATTR_MAP: dict[str, str] = {
|
||||||
"vae": "vae_config",
|
"vae": "vae_config",
|
||||||
@@ -495,6 +543,22 @@ class ComposedPipelineBase(ABC):
|
|||||||
self.memory_usages,
|
self.memory_usages,
|
||||||
round(current_platform.get_available_gpu_memory(), 2),
|
round(current_platform.get_available_gpu_memory(), 2),
|
||||||
)
|
)
|
||||||
|
total_consumed_gb = sum(
|
||||||
|
usage
|
||||||
|
for usage in self.memory_usages.values()
|
||||||
|
if isinstance(usage, (int, float))
|
||||||
|
)
|
||||||
|
available_after_gb = current_platform.get_available_gpu_memory()
|
||||||
|
logger.debug(
|
||||||
|
"Module load summary: required_modules=%s loaded_modules=%s "
|
||||||
|
"memory_usages_gb=%s total_consumed_gb=%.2f "
|
||||||
|
"available_after_gb=%.2f",
|
||||||
|
list(required_modules),
|
||||||
|
list(loaded_components.keys()),
|
||||||
|
self.memory_usages,
|
||||||
|
total_consumed_gb,
|
||||||
|
available_after_gb,
|
||||||
|
)
|
||||||
|
|
||||||
return loaded_components
|
return loaded_components
|
||||||
|
|
||||||
@@ -502,6 +566,24 @@ class ComposedPipelineBase(ABC):
|
|||||||
def _infer_stage_name(stage: PipelineStage) -> str:
|
def _infer_stage_name(stage: PipelineStage) -> str:
|
||||||
return stage.__class__.__name__
|
return stage.__class__.__name__
|
||||||
|
|
||||||
|
def _should_add_stage_for_role(
|
||||||
|
self,
|
||||||
|
role_affinity: RoleType,
|
||||||
|
stage_name: str,
|
||||||
|
) -> bool:
|
||||||
|
if self._disagg_role == RoleType.MONOLITHIC:
|
||||||
|
return True
|
||||||
|
if role_affinity == self._disagg_role:
|
||||||
|
return True
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Disagg role=%s: skipping stage %s (affinity=%s)",
|
||||||
|
self._disagg_role.value,
|
||||||
|
stage_name,
|
||||||
|
role_affinity.value,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
def _profile_stage_name(self, stage: PipelineStage, stage_name: str) -> str:
|
def _profile_stage_name(self, stage: PipelineStage, stage_name: str) -> str:
|
||||||
class_name = stage.__class__.__name__
|
class_name = stage.__class__.__name__
|
||||||
if any(existing.__class__.__name__ == class_name for existing in self._stages):
|
if any(existing.__class__.__name__ == class_name for existing in self._stages):
|
||||||
@@ -513,22 +595,13 @@ class ComposedPipelineBase(ABC):
|
|||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
|
|
||||||
assert self.modules is not None, "No modules are registered"
|
assert self.modules is not None, "No modules are registered"
|
||||||
|
if stage_name is None:
|
||||||
|
stage_name = self._infer_stage_name(stage)
|
||||||
|
|
||||||
# Filter stages based on disaggregation role
|
# Filter stages based on disaggregation role
|
||||||
if self._disagg_role != RoleType.MONOLITHIC:
|
if not self._should_add_stage_for_role(stage.role_affinity, stage_name):
|
||||||
if stage.role_affinity != self._disagg_role:
|
|
||||||
if stage_name is None:
|
|
||||||
stage_name = self._infer_stage_name(stage)
|
|
||||||
logger.info(
|
|
||||||
"Disagg role=%s: skipping stage %s (affinity=%s)",
|
|
||||||
self._disagg_role.value,
|
|
||||||
stage_name,
|
|
||||||
stage.role_affinity.value,
|
|
||||||
)
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
if stage_name is None:
|
|
||||||
stage_name = self._infer_stage_name(stage)
|
|
||||||
if stage_name in self._stage_name_mapping:
|
if stage_name in self._stage_name_mapping:
|
||||||
raise ValueError(f"Duplicate stage name detected: {stage_name}")
|
raise ValueError(f"Duplicate stage name detected: {stage_name}")
|
||||||
|
|
||||||
@@ -538,6 +611,17 @@ class ComposedPipelineBase(ABC):
|
|||||||
self._stage_name_mapping[stage_name] = stage
|
self._stage_name_mapping[stage_name] = stage
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
def add_stage_factory(
|
||||||
|
self,
|
||||||
|
role_affinity: RoleType,
|
||||||
|
stage_factory: Callable[[], PipelineStage],
|
||||||
|
stage_name: str,
|
||||||
|
) -> "ComposedPipelineBase":
|
||||||
|
assert self.modules is not None, "No modules are registered"
|
||||||
|
if not self._should_add_stage_for_role(role_affinity, stage_name):
|
||||||
|
return self
|
||||||
|
return self.add_stage(stage_factory(), stage_name)
|
||||||
|
|
||||||
def add_stages(
|
def add_stages(
|
||||||
self, stages: list[PipelineStage | tuple[PipelineStage, str]]
|
self, stages: list[PipelineStage | tuple[PipelineStage, str]]
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
@@ -579,12 +663,12 @@ class ComposedPipelineBase(ABC):
|
|||||||
def add_standard_timestep_preparation_stage(
|
def add_standard_timestep_preparation_stage(
|
||||||
self,
|
self,
|
||||||
scheduler_key: str = "scheduler",
|
scheduler_key: str = "scheduler",
|
||||||
prepare_extra_kwargs: list[Callable] | None = [],
|
prepare_extra_kwargs: list[Callable] | None = None,
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
return self.add_stage(
|
return self.add_stage(
|
||||||
TimestepPreparationStage(
|
TimestepPreparationStage(
|
||||||
scheduler=self.get_module(scheduler_key),
|
scheduler=self.get_module(scheduler_key),
|
||||||
prepare_extra_set_timesteps_kwargs=prepare_extra_kwargs,
|
prepare_extra_set_timesteps_kwargs=list(prepare_extra_kwargs or []),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -606,8 +690,10 @@ class ComposedPipelineBase(ABC):
|
|||||||
transformer_2_key: str | None = "transformer_2",
|
transformer_2_key: str | None = "transformer_2",
|
||||||
scheduler_key: str = "scheduler",
|
scheduler_key: str = "scheduler",
|
||||||
vae_key: str | None = "vae",
|
vae_key: str | None = "vae",
|
||||||
|
stage_name: str = "denoising_stage",
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
|
|
||||||
|
def create_stage() -> PipelineStage:
|
||||||
kwargs = {
|
kwargs = {
|
||||||
"transformer": self.get_module(transformer_key),
|
"transformer": self.get_module(transformer_key),
|
||||||
"scheduler": self.get_module(scheduler_key),
|
"scheduler": self.get_module(scheduler_key),
|
||||||
@@ -624,25 +710,37 @@ class ComposedPipelineBase(ABC):
|
|||||||
kwargs["vae"] = vae
|
kwargs["vae"] = vae
|
||||||
kwargs["pipeline"] = self
|
kwargs["pipeline"] = self
|
||||||
|
|
||||||
return self.add_stage(DenoisingStage(**kwargs))
|
return DenoisingStage(**kwargs)
|
||||||
|
|
||||||
|
return self.add_stage_factory(
|
||||||
|
RoleType.DENOISER,
|
||||||
|
create_stage,
|
||||||
|
stage_name,
|
||||||
|
)
|
||||||
|
|
||||||
def add_standard_decoding_stage(
|
def add_standard_decoding_stage(
|
||||||
self,
|
self,
|
||||||
vae_key: str = "vae",
|
vae_key: str = "vae",
|
||||||
|
stage_name: str = "decoding_stage",
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
|
|
||||||
return self.add_stage(
|
def create_stage() -> PipelineStage:
|
||||||
DecodingStage(
|
return DecodingStage(
|
||||||
vae=self.get_module(vae_key),
|
vae=self.get_module(vae_key),
|
||||||
pipeline=self,
|
pipeline=self,
|
||||||
component_name=vae_key,
|
component_name=vae_key,
|
||||||
),
|
)
|
||||||
|
|
||||||
|
return self.add_stage_factory(
|
||||||
|
RoleType.DECODER,
|
||||||
|
create_stage,
|
||||||
|
stage_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
def add_standard_t2i_stages(
|
def add_standard_t2i_stages(
|
||||||
self,
|
self,
|
||||||
include_input_validation: bool = True,
|
include_input_validation: bool = True,
|
||||||
prepare_extra_timestep_kwargs: list[Callable] | None = [],
|
prepare_extra_timestep_kwargs: list[Callable] | None = None,
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
|
|
||||||
if include_input_validation:
|
if include_input_validation:
|
||||||
@@ -671,7 +769,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
prompt_text_encoder_key: str = "text_encoder",
|
prompt_text_encoder_key: str = "text_encoder",
|
||||||
image_vae_key: str = "vae",
|
image_vae_key: str = "vae",
|
||||||
image_vae_stage_kwargs: dict[str, Any] | None = None,
|
image_vae_stage_kwargs: dict[str, Any] | None = None,
|
||||||
prepare_extra_timestep_kwargs: list[Callable] | None = [],
|
prepare_extra_timestep_kwargs: list[Callable] | None = None,
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
if include_input_validation:
|
if include_input_validation:
|
||||||
self.add_stage(
|
self.add_stage(
|
||||||
@@ -696,7 +794,10 @@ class ComposedPipelineBase(ABC):
|
|||||||
self.add_stage(
|
self.add_stage(
|
||||||
ImageVAEEncodingStage(
|
ImageVAEEncodingStage(
|
||||||
vae=self.get_module(image_vae_key),
|
vae=self.get_module(image_vae_key),
|
||||||
|
**{
|
||||||
|
"component_name": image_vae_key,
|
||||||
**(image_vae_stage_kwargs or {}),
|
**(image_vae_stage_kwargs or {}),
|
||||||
|
},
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -723,8 +824,9 @@ class ComposedPipelineBase(ABC):
|
|||||||
image_vae_encoding_position: Literal[
|
image_vae_encoding_position: Literal[
|
||||||
"before_timestep", "after_latent"
|
"before_timestep", "after_latent"
|
||||||
] = "before_timestep",
|
] = "before_timestep",
|
||||||
prepare_extra_timestep_kwargs: list[Callable] | None = [],
|
prepare_extra_timestep_kwargs: list[Callable] | None = None,
|
||||||
denoising_stage_factory: Callable[[], PipelineStage] | None = None,
|
denoising_stage_factory: Callable[[], PipelineStage] | None = None,
|
||||||
|
denoising_stage_name: str = "denoising_stage",
|
||||||
) -> "ComposedPipelineBase":
|
) -> "ComposedPipelineBase":
|
||||||
if include_input_validation:
|
if include_input_validation:
|
||||||
self.add_stage(
|
self.add_stage(
|
||||||
@@ -750,7 +852,10 @@ class ComposedPipelineBase(ABC):
|
|||||||
self.add_stage(
|
self.add_stage(
|
||||||
ImageVAEEncodingStage(
|
ImageVAEEncodingStage(
|
||||||
vae=self.get_module(image_vae_key),
|
vae=self.get_module(image_vae_key),
|
||||||
|
**{
|
||||||
|
"component_name": image_vae_key,
|
||||||
**(image_vae_stage_kwargs or {}),
|
**(image_vae_stage_kwargs or {}),
|
||||||
|
},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -762,7 +867,10 @@ class ComposedPipelineBase(ABC):
|
|||||||
self.add_stage(
|
self.add_stage(
|
||||||
ImageVAEEncodingStage(
|
ImageVAEEncodingStage(
|
||||||
vae=self.get_module(image_vae_key),
|
vae=self.get_module(image_vae_key),
|
||||||
|
**{
|
||||||
|
"component_name": image_vae_key,
|
||||||
**(image_vae_stage_kwargs or {}),
|
**(image_vae_stage_kwargs or {}),
|
||||||
|
},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
elif image_vae_encoding_position != "before_timestep":
|
elif image_vae_encoding_position != "before_timestep":
|
||||||
@@ -773,7 +881,11 @@ class ComposedPipelineBase(ABC):
|
|||||||
if denoising_stage_factory is None:
|
if denoising_stage_factory is None:
|
||||||
self.add_standard_denoising_stage()
|
self.add_standard_denoising_stage()
|
||||||
else:
|
else:
|
||||||
self.add_stage(denoising_stage_factory())
|
self.add_stage_factory(
|
||||||
|
RoleType.DENOISER,
|
||||||
|
denoising_stage_factory,
|
||||||
|
denoising_stage_name,
|
||||||
|
)
|
||||||
|
|
||||||
self.add_standard_decoding_stage()
|
self.add_standard_decoding_stage()
|
||||||
return self
|
return self
|
||||||
|
|||||||
@@ -128,18 +128,18 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
|
|||||||
self,
|
self,
|
||||||
image_processor: Any,
|
image_processor: Any,
|
||||||
conditioner: Any,
|
conditioner: Any,
|
||||||
vae: Any,
|
|
||||||
model: Any,
|
|
||||||
scheduler: Any,
|
scheduler: Any,
|
||||||
config: Hunyuan3D2PipelineConfig,
|
config: Hunyuan3D2PipelineConfig,
|
||||||
|
latent_shape: tuple[int, ...],
|
||||||
|
guidance_embed: bool,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.image_processor = image_processor
|
self.image_processor = image_processor
|
||||||
self.conditioner = conditioner
|
self.conditioner = conditioner
|
||||||
self.vae = vae
|
|
||||||
self.model = model
|
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
self.config = config
|
self.config = config
|
||||||
|
self.latent_shape = latent_shape
|
||||||
|
self.guidance_embed = guidance_embed
|
||||||
|
|
||||||
def _validate_input(self, batch: Req, server_args: ServerArgs) -> None:
|
def _validate_input(self, batch: Req, server_args: ServerArgs) -> None:
|
||||||
if batch.image_path is None:
|
if batch.image_path is None:
|
||||||
@@ -160,10 +160,41 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
|
|||||||
def _prepare_latents(self, batch_size, dtype, device, generator, scheduler):
|
def _prepare_latents(self, batch_size, dtype, device, generator, scheduler):
|
||||||
from diffusers.utils.torch_utils import randn_tensor
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
|
|
||||||
shape = (batch_size, *self.vae.latent_shape)
|
shape = (batch_size, *self.latent_shape)
|
||||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||||
return latents * getattr(scheduler, "init_noise_sigma", 1.0)
|
return latents * getattr(scheduler, "init_noise_sigma", 1.0)
|
||||||
|
|
||||||
|
def _find_conditioner_dtype(self, items_fn_name: str) -> torch.dtype | None:
|
||||||
|
items_fn = getattr(self.conditioner, items_fn_name, None)
|
||||||
|
if not callable(items_fn):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
for item in items_fn():
|
||||||
|
if isinstance(item, torch.Tensor) and torch.is_floating_point(item):
|
||||||
|
return item.dtype
|
||||||
|
except TypeError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to inspect Hunyuan3D conditioner %s() for runtime dtype; "
|
||||||
|
"falling back to the sample tensor dtype. error=%s",
|
||||||
|
items_fn_name,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _resolve_runtime_dtype(
|
||||||
|
self, sample_tensor: torch.Tensor | None = None
|
||||||
|
) -> torch.dtype:
|
||||||
|
for items_fn_name in ("parameters", "buffers"):
|
||||||
|
dtype = self._find_conditioner_dtype(items_fn_name)
|
||||||
|
if dtype is not None:
|
||||||
|
return dtype
|
||||||
|
|
||||||
|
if isinstance(sample_tensor, torch.Tensor) and torch.is_floating_point(
|
||||||
|
sample_tensor
|
||||||
|
):
|
||||||
|
return sample_tensor.dtype
|
||||||
|
return torch.float32
|
||||||
|
|
||||||
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
|
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
|
||||||
# 1. Input validation
|
# 1. Input validation
|
||||||
self._validate_input(batch, server_args)
|
self._validate_input(batch, server_args)
|
||||||
@@ -173,14 +204,14 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
|
|||||||
image = cond_inputs.pop("image")
|
image = cond_inputs.pop("image")
|
||||||
|
|
||||||
device = self.device
|
device = self.device
|
||||||
dtype = next(self.model.parameters()).dtype
|
dtype = self._resolve_runtime_dtype(
|
||||||
|
image if isinstance(image, torch.Tensor) else None
|
||||||
|
)
|
||||||
image = _move_to_device(image, device, dtype)
|
image = _move_to_device(image, device, dtype)
|
||||||
cond_inputs = _move_to_device(cond_inputs, device, dtype)
|
cond_inputs = _move_to_device(cond_inputs, device, dtype)
|
||||||
|
|
||||||
# 3. Conditioning with CFG
|
# 3. Conditioning with CFG
|
||||||
do_cfg = batch.guidance_scale >= 0 and not (
|
do_cfg = batch.guidance_scale >= 0 and not self.guidance_embed
|
||||||
hasattr(self.model, "guidance_embed") and self.model.guidance_embed is True
|
|
||||||
)
|
|
||||||
|
|
||||||
cond = self.conditioner(image=image, **cond_inputs)
|
cond = self.conditioner(image=image, **cond_inputs)
|
||||||
if do_cfg:
|
if do_cfg:
|
||||||
@@ -216,7 +247,7 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
|
|||||||
latents = self._prepare_latents(batch_size, dtype, device, generator, scheduler)
|
latents = self._prepare_latents(batch_size, dtype, device, generator, scheduler)
|
||||||
|
|
||||||
guidance = None
|
guidance = None
|
||||||
if hasattr(self.model, "guidance_embed") and self.model.guidance_embed is True:
|
if self.guidance_embed:
|
||||||
guidance = torch.tensor(
|
guidance = torch.tensor(
|
||||||
[batch.guidance_scale] * batch_size, device=device, dtype=dtype
|
[batch.guidance_scale] * batch_size, device=device, dtype=dtype
|
||||||
)
|
)
|
||||||
@@ -416,6 +447,12 @@ class Hunyuan3DShapeExportStage(PipelineStage):
|
|||||||
self.vae = vae
|
self.vae = vae
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role_affinity(self):
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
|
|
||||||
|
return RoleType.DECODER
|
||||||
|
|
||||||
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
|
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
|
||||||
if self.config.shape_mc_algo is not None:
|
if self.config.shape_mc_algo is not None:
|
||||||
try:
|
try:
|
||||||
@@ -473,6 +510,12 @@ class Hunyuan3DShapeSaveStage(PipelineStage):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role_affinity(self):
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
|
|
||||||
|
return RoleType.DECODER
|
||||||
|
|
||||||
def _get_output_paths(self, batch: Req) -> tuple[str, str]:
|
def _get_output_paths(self, batch: Req) -> tuple[str, str]:
|
||||||
output_path = batch.output_file_path() or os.path.join(
|
output_path = batch.output_file_path() or os.path.join(
|
||||||
batch.output_path, "output.obj"
|
batch.output_path, "output.obj"
|
||||||
|
|||||||
@@ -803,9 +803,15 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
"vae_image_sizes",
|
"vae_image_sizes",
|
||||||
)
|
)
|
||||||
|
|
||||||
def __init__(self, vae: ParallelTiledVAE, **kwargs) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
vae: ParallelTiledVAE,
|
||||||
|
component_name: str = "vae",
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.vae: ParallelTiledVAE = vae
|
self.vae: ParallelTiledVAE = vae
|
||||||
|
self.component_name = component_name
|
||||||
|
|
||||||
def component_uses(
|
def component_uses(
|
||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
self, server_args: ServerArgs, stage_name: str | None = None
|
||||||
@@ -815,7 +821,7 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
return [
|
return [
|
||||||
ComponentUse(
|
ComponentUse(
|
||||||
stage_name,
|
stage_name,
|
||||||
"vae",
|
self.component_name,
|
||||||
target_dtype=vae_dtype,
|
target_dtype=vae_dtype,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
@@ -851,7 +857,10 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
vae_dtype != torch.float32
|
vae_dtype != torch.float32
|
||||||
) and not server_args.disable_autocast
|
) and not server_args.disable_autocast
|
||||||
|
|
||||||
with self.use_declared_component(component_name="vae", module=self.vae) as vae:
|
with self.use_declared_component(
|
||||||
|
component_name=self.component_name,
|
||||||
|
module=self.vae,
|
||||||
|
) as vae:
|
||||||
assert vae is not None
|
assert vae is not None
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
|
|
||||||
|
|||||||
+5
@@ -13,6 +13,7 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
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.managers.memory_managers.component_manager import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||||
ComponentUse,
|
ComponentUse,
|
||||||
@@ -105,6 +106,10 @@ class HeliosChunkedDenoisingStage(PipelineStage):
|
|||||||
self.transformer = transformer
|
self.transformer = transformer
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role_affinity(self) -> RoleType:
|
||||||
|
return RoleType.DENOISER
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def parallelism_type(self):
|
def parallelism_type(self):
|
||||||
return StageParallelismType.REPLICATED
|
return StageParallelismType.REPLICATED
|
||||||
|
|||||||
+16
@@ -21,6 +21,7 @@ import torch.nn as nn
|
|||||||
from diffusers.utils.torch_utils import randn_tensor
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
get_local_torch_device,
|
get_local_torch_device,
|
||||||
get_world_group,
|
get_world_group,
|
||||||
@@ -187,6 +188,10 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
return uses
|
return uses
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role_affinity(self) -> RoleType:
|
||||||
|
return RoleType.DENOISER
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def parallelism_type(self) -> StageParallelismType:
|
def parallelism_type(self) -> StageParallelismType:
|
||||||
if get_global_server_args().enable_cfg_parallel:
|
if get_global_server_args().enable_cfg_parallel:
|
||||||
@@ -240,6 +245,13 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
"""
|
"""
|
||||||
if not server_args.enable_torch_compile or not isinstance(module, nn.Module):
|
if not server_args.enable_torch_compile or not isinstance(module, nn.Module):
|
||||||
return
|
return
|
||||||
|
if current_platform.is_hip():
|
||||||
|
logger.warning(
|
||||||
|
"Skipping torch.compile for %s on ROCm because the current "
|
||||||
|
"HIPRTC/Inductor path can emit invalid bf16 kernels.",
|
||||||
|
module.__class__.__name__,
|
||||||
|
)
|
||||||
|
return
|
||||||
compile_kwargs: dict[str, object] = {"fullgraph": False, "dynamic": None}
|
compile_kwargs: dict[str, object] = {"fullgraph": False, "dynamic": None}
|
||||||
|
|
||||||
if current_platform.is_npu():
|
if current_platform.is_npu():
|
||||||
@@ -954,6 +966,10 @@ class MOVADecodingStage(PipelineStage):
|
|||||||
ComponentUse(stage_name, "audio_vae"),
|
ComponentUse(stage_name, "audio_vae"),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role_affinity(self) -> RoleType:
|
||||||
|
return RoleType.DECODER
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def parallelism_type(self) -> StageParallelismType:
|
def parallelism_type(self) -> StageParallelismType:
|
||||||
if get_global_server_args().enable_cfg_parallel:
|
if get_global_server_args().enable_cfg_parallel:
|
||||||
|
|||||||
+40
-8
@@ -21,6 +21,31 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_text_encoder_dtype(
|
||||||
|
text_encoder: object, fallback: torch.dtype = torch.bfloat16
|
||||||
|
) -> torch.dtype:
|
||||||
|
module_dtype = getattr(text_encoder, "dtype", None)
|
||||||
|
if isinstance(module_dtype, torch.dtype):
|
||||||
|
return module_dtype
|
||||||
|
|
||||||
|
for tensor_source in ("parameters", "buffers"):
|
||||||
|
tensors = getattr(text_encoder, tensor_source, None)
|
||||||
|
if not callable(tensors):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
for tensor in tensors():
|
||||||
|
if isinstance(tensor, torch.Tensor) and torch.is_floating_point(tensor):
|
||||||
|
return tensor.dtype
|
||||||
|
except TypeError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to inspect text encoder %s() for dtype: %s",
|
||||||
|
tensor_source,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
return fallback
|
||||||
|
|
||||||
|
|
||||||
def _seq_lens_from_optional_mask(
|
def _seq_lens_from_optional_mask(
|
||||||
prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None
|
prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None
|
||||||
) -> list[int]:
|
) -> list[int]:
|
||||||
@@ -125,6 +150,7 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
vae,
|
vae,
|
||||||
|
text_encoder,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
transformer,
|
transformer,
|
||||||
@@ -137,14 +163,14 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
|
|||||||
self.vae = vae.to(dtype=vae_dtype)
|
self.vae = vae.to(dtype=vae_dtype)
|
||||||
self.vae_dtype = vae_dtype
|
self.vae_dtype = vae_dtype
|
||||||
self.text_encoder_dtype = text_encoder_dtype
|
self.text_encoder_dtype = text_encoder_dtype
|
||||||
|
if text_encoder is None:
|
||||||
from transformers import Qwen2_5_VLForConditionalGeneration
|
from transformers import Qwen2_5_VLForConditionalGeneration
|
||||||
|
|
||||||
self.text_encoder = (
|
text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
||||||
Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
|
||||||
model_path, subfolder="text_encoder"
|
model_path, subfolder="text_encoder"
|
||||||
)
|
)
|
||||||
.to(get_local_torch_device())
|
self.text_encoder = text_encoder.to(
|
||||||
.to(dtype=self.text_encoder_dtype)
|
device=get_local_torch_device(), dtype=self.text_encoder_dtype
|
||||||
)
|
)
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.processor = processor
|
self.processor = processor
|
||||||
@@ -186,9 +212,15 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
stage_name = self._component_stage_name(stage_name)
|
stage_name = self._component_stage_name(stage_name)
|
||||||
return [
|
return [
|
||||||
ComponentUse(
|
ComponentUse(
|
||||||
stage_name, "qwen_layered_text_encoder", target_dtype=torch.bfloat16
|
stage_name,
|
||||||
|
"text_encoder",
|
||||||
|
target_dtype=self.text_encoder_dtype,
|
||||||
|
),
|
||||||
|
ComponentUse(
|
||||||
|
stage_name,
|
||||||
|
"vae",
|
||||||
|
target_dtype=self.vae_dtype,
|
||||||
),
|
),
|
||||||
ComponentUse(stage_name, "vae", target_dtype=torch.bfloat16),
|
|
||||||
]
|
]
|
||||||
|
|
||||||
# Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage.QwenImagePipeline._extract_masked_hidden
|
# Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage.QwenImagePipeline._extract_masked_hidden
|
||||||
@@ -232,7 +264,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
device: Optional[torch.device] = None,
|
device: Optional[torch.device] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
):
|
):
|
||||||
dtype = dtype or self.text_encoder.dtype
|
dtype = dtype or _resolve_text_encoder_dtype(self.text_encoder)
|
||||||
|
|
||||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||||
|
|
||||||
@@ -482,7 +514,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
|
|
||||||
prompt = batch.prompt
|
prompt = batch.prompt
|
||||||
with self.use_declared_component(
|
with self.use_declared_component(
|
||||||
component_name="qwen_layered_text_encoder",
|
component_name="text_encoder",
|
||||||
module=self.text_encoder,
|
module=self.text_encoder,
|
||||||
) as text_encoder:
|
) as text_encoder:
|
||||||
assert text_encoder is not None
|
assert text_encoder is not None
|
||||||
|
|||||||
@@ -60,13 +60,13 @@ class TimestepPreparationStage(PipelineStage):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
scheduler,
|
scheduler,
|
||||||
prepare_extra_set_timesteps_kwargs: list[
|
prepare_extra_set_timesteps_kwargs: (
|
||||||
Callable[[Req, ServerArgs], Tuple[str, Any]]
|
list[Callable[[Req, ServerArgs], Tuple[str, Any]]] | None
|
||||||
] = [],
|
) = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
self.prepare_extra_set_timesteps_kwargs = (
|
self.prepare_extra_set_timesteps_kwargs = list(
|
||||||
prepare_extra_set_timesteps_kwargs or []
|
prepare_extra_set_timesteps_kwargs or []
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,641 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Unit tests for disaggregation role-based module filtering."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
|
||||||
|
Hunyuan3D2PipelineConfig,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime import server_args as server_args_module
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import (
|
||||||
|
RoleType,
|
||||||
|
filter_modules_for_role,
|
||||||
|
get_module_role,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.flux_2 import Flux2Pipeline
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.glm_image import GlmImagePipeline
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.hunyuan3d_pipeline import (
|
||||||
|
Hunyuan3D2Pipeline,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import LTX2Pipeline
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.mova_pipeline import (
|
||||||
|
MOVAPipeline,
|
||||||
|
MOVAPipelineAlias,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.qwen_image import (
|
||||||
|
QwenImageEditPipeline,
|
||||||
|
QwenImageLayeredPipeline,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.wan_i2v_dmd_pipeline import (
|
||||||
|
WanImageToVideoDmdPipeline,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines.wan_i2v_pipeline import (
|
||||||
|
WanImageToVideoPipeline,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
|
ComposedPipelineBase,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising_av import (
|
||||||
|
LTX2RefinementStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.hunyuan3d_shape import (
|
||||||
|
Hunyuan3DShapeBeforeDenoisingStage,
|
||||||
|
Hunyuan3DShapeExportStage,
|
||||||
|
Hunyuan3DShapeSaveStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
|
||||||
|
ImageVAEEncodingStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.helios_denoising import (
|
||||||
|
HeliosChunkedDenoisingStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.mova import (
|
||||||
|
MOVADecodingStage,
|
||||||
|
MOVADenoisingStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered import (
|
||||||
|
QwenImageLayeredBeforeDenoisingStage,
|
||||||
|
_resolve_text_encoder_dtype,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import set_global_server_args
|
||||||
|
|
||||||
|
|
||||||
|
class _GlobalStageArgsMixin:
|
||||||
|
def _install_stage_server_args(self, **kwargs):
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
comfyui_mode=False,
|
||||||
|
enable_torch_compile=False,
|
||||||
|
enable_cfg_parallel=False,
|
||||||
|
attention_backend=None,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
set_global_server_args(server_args)
|
||||||
|
return server_args
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
super().setUp()
|
||||||
|
self._prev_global_server_args = server_args_module._global_server_args
|
||||||
|
self._install_stage_server_args()
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
set_global_server_args(self._prev_global_server_args)
|
||||||
|
super().tearDown()
|
||||||
|
|
||||||
|
|
||||||
|
class TestRoleType(unittest.TestCase):
|
||||||
|
def test_from_string(self):
|
||||||
|
self.assertEqual(RoleType.from_string("monolithic"), RoleType.MONOLITHIC)
|
||||||
|
self.assertEqual(RoleType.from_string("encoder"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(RoleType.from_string("denoiser"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(RoleType.from_string("decoder"), RoleType.DECODER)
|
||||||
|
self.assertEqual(RoleType.from_string("ENCODER"), RoleType.ENCODER)
|
||||||
|
|
||||||
|
def test_from_string_backward_compat(self):
|
||||||
|
self.assertEqual(RoleType.from_string("denoising"), RoleType.DENOISER)
|
||||||
|
|
||||||
|
def test_from_string_invalid(self):
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
RoleType.from_string("invalid")
|
||||||
|
|
||||||
|
def test_choices(self):
|
||||||
|
choices = RoleType.choices()
|
||||||
|
self.assertIn("monolithic", choices)
|
||||||
|
self.assertIn("encoder", choices)
|
||||||
|
self.assertIn("denoiser", choices)
|
||||||
|
self.assertIn("denoising", choices)
|
||||||
|
self.assertIn("decoder", choices)
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetModuleRole(unittest.TestCase):
|
||||||
|
def test_encoder_modules(self):
|
||||||
|
self.assertEqual(get_module_role("text_encoder"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("text_encoder_2"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("tokenizer"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("tokenizer_2"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("image_encoder"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("image_processor"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("connectors"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("vision_language_encoder"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("hy3dshape_conditioner"), RoleType.ENCODER)
|
||||||
|
self.assertEqual(get_module_role("hy3dshape_image_processor"), RoleType.ENCODER)
|
||||||
|
|
||||||
|
def test_denoiser_modules(self):
|
||||||
|
self.assertEqual(get_module_role("transformer"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(get_module_role("transformer_2"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(get_module_role("video_dit"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(get_module_role("video_dit_2"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(get_module_role("audio_dit"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(get_module_role("dual_tower_bridge"), RoleType.DENOISER)
|
||||||
|
self.assertEqual(get_module_role("hy3dshape_model"), RoleType.DENOISER)
|
||||||
|
|
||||||
|
def test_decoder_modules(self):
|
||||||
|
self.assertEqual(get_module_role("vae"), RoleType.DECODER)
|
||||||
|
self.assertEqual(get_module_role("audio_vae"), RoleType.DECODER)
|
||||||
|
self.assertEqual(get_module_role("video_vae"), RoleType.DECODER)
|
||||||
|
self.assertEqual(get_module_role("vocoder"), RoleType.DECODER)
|
||||||
|
self.assertEqual(get_module_role("hy3dshape_vae"), RoleType.DECODER)
|
||||||
|
|
||||||
|
def test_shared_modules(self):
|
||||||
|
self.assertIsNone(get_module_role("scheduler"))
|
||||||
|
self.assertIsNone(get_module_role("hy3dshape_scheduler"))
|
||||||
|
|
||||||
|
|
||||||
|
class TestFilterModulesForRole(unittest.TestCase):
|
||||||
|
WAN_MODULES = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
|
||||||
|
|
||||||
|
def test_monolithic_keeps_all(self):
|
||||||
|
result = filter_modules_for_role(self.WAN_MODULES, RoleType.MONOLITHIC)
|
||||||
|
self.assertEqual(result, self.WAN_MODULES)
|
||||||
|
|
||||||
|
def test_encoder_does_not_keep_decoder_modules_by_default(self):
|
||||||
|
result = filter_modules_for_role(self.WAN_MODULES, RoleType.ENCODER)
|
||||||
|
self.assertEqual(result, ["text_encoder", "tokenizer", "scheduler"])
|
||||||
|
|
||||||
|
def test_encoder_can_keep_explicit_cross_role_modules(self):
|
||||||
|
result = filter_modules_for_role(
|
||||||
|
self.WAN_MODULES,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules={"vae"},
|
||||||
|
)
|
||||||
|
self.assertEqual(result, ["text_encoder", "tokenizer", "vae", "scheduler"])
|
||||||
|
|
||||||
|
def test_denoiser_skips_encoders_and_vae(self):
|
||||||
|
result = filter_modules_for_role(self.WAN_MODULES, RoleType.DENOISER)
|
||||||
|
self.assertEqual(result, ["transformer", "scheduler"])
|
||||||
|
|
||||||
|
def test_decoder_keeps_vae_and_scheduler(self):
|
||||||
|
result = filter_modules_for_role(self.WAN_MODULES, RoleType.DECODER)
|
||||||
|
self.assertEqual(result, ["vae", "scheduler"])
|
||||||
|
|
||||||
|
|
||||||
|
class TestFilterModulesLTX2(unittest.TestCase):
|
||||||
|
LTX2_MODULES = [
|
||||||
|
"transformer",
|
||||||
|
"text_encoder",
|
||||||
|
"tokenizer",
|
||||||
|
"scheduler",
|
||||||
|
"vae",
|
||||||
|
"audio_vae",
|
||||||
|
"vocoder",
|
||||||
|
"connectors",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_decoder_includes_audio(self):
|
||||||
|
result = filter_modules_for_role(self.LTX2_MODULES, RoleType.DECODER)
|
||||||
|
self.assertEqual(result, ["scheduler", "vae", "audio_vae", "vocoder"])
|
||||||
|
|
||||||
|
def test_encoder_does_not_keep_decoder_modules_by_default(self):
|
||||||
|
result = filter_modules_for_role(self.LTX2_MODULES, RoleType.ENCODER)
|
||||||
|
self.assertEqual(
|
||||||
|
result, ["text_encoder", "tokenizer", "scheduler", "connectors"]
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_denoiser_can_keep_ti2v_decoder_components(self):
|
||||||
|
result = filter_modules_for_role(
|
||||||
|
self.LTX2_MODULES,
|
||||||
|
RoleType.DENOISER,
|
||||||
|
extra_allowed_modules={"vae", "audio_vae"},
|
||||||
|
)
|
||||||
|
self.assertEqual(result, ["transformer", "scheduler", "vae", "audio_vae"])
|
||||||
|
|
||||||
|
|
||||||
|
# Consolidated from test_pipeline_stage_role_filter.py.
|
||||||
|
class _FakePipeline(ComposedPipelineBase):
|
||||||
|
pipeline_name = "FakePipeline"
|
||||||
|
_required_config_modules = []
|
||||||
|
|
||||||
|
def initialize_pipeline(self, server_args):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def create_pipeline_stages(self, server_args) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _make_pipeline(role: RoleType) -> _FakePipeline:
|
||||||
|
pipeline = object.__new__(_FakePipeline)
|
||||||
|
pipeline.modules = {}
|
||||||
|
pipeline._stages = []
|
||||||
|
pipeline._stage_name_mapping = {}
|
||||||
|
pipeline._disagg_role = role
|
||||||
|
return pipeline
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeStage:
|
||||||
|
def __init__(self, role_affinity: RoleType):
|
||||||
|
self.role_affinity = role_affinity
|
||||||
|
self.registered_stage_name = None
|
||||||
|
self.profile_stage_name = None
|
||||||
|
|
||||||
|
def set_registered_stage_name(self, stage_name: str) -> None:
|
||||||
|
self.registered_stage_name = stage_name
|
||||||
|
|
||||||
|
def set_profile_stage_name(self, stage_name: str) -> None:
|
||||||
|
self.profile_stage_name = stage_name
|
||||||
|
|
||||||
|
|
||||||
|
class TestPipelineStageRoleFilter(unittest.TestCase):
|
||||||
|
def test_stage_factory_skips_without_constructing_for_other_role(self):
|
||||||
|
pipeline = _make_pipeline(RoleType.ENCODER)
|
||||||
|
|
||||||
|
def should_not_construct():
|
||||||
|
raise AssertionError("stage factory should have been skipped")
|
||||||
|
|
||||||
|
pipeline.add_stage_factory(
|
||||||
|
RoleType.DENOISER,
|
||||||
|
should_not_construct,
|
||||||
|
"denoising_stage",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(pipeline.stages, [])
|
||||||
|
|
||||||
|
def test_stage_factory_constructs_for_matching_role(self):
|
||||||
|
pipeline = _make_pipeline(RoleType.DENOISER)
|
||||||
|
stage = _FakeStage(RoleType.DENOISER)
|
||||||
|
events = []
|
||||||
|
|
||||||
|
def create_stage():
|
||||||
|
events.append("called")
|
||||||
|
return stage
|
||||||
|
|
||||||
|
pipeline.add_stage_factory(
|
||||||
|
RoleType.DENOISER,
|
||||||
|
create_stage,
|
||||||
|
"denoising_stage",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(events, ["called"])
|
||||||
|
self.assertIs(pipeline.get_stage("denoising_stage"), stage)
|
||||||
|
self.assertEqual(stage.registered_stage_name, "denoising_stage")
|
||||||
|
|
||||||
|
def test_encoder_role_does_not_construct_standard_denoising_stage(self):
|
||||||
|
pipeline = _make_pipeline(RoleType.ENCODER)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base.DenoisingStage",
|
||||||
|
side_effect=AssertionError("DenoisingStage should not be constructed"),
|
||||||
|
):
|
||||||
|
pipeline.add_standard_denoising_stage()
|
||||||
|
|
||||||
|
self.assertEqual(pipeline.stages, [])
|
||||||
|
|
||||||
|
def test_encoder_role_does_not_construct_standard_decoding_stage(self):
|
||||||
|
pipeline = _make_pipeline(RoleType.ENCODER)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base.DecodingStage",
|
||||||
|
side_effect=AssertionError("DecodingStage should not be constructed"),
|
||||||
|
):
|
||||||
|
pipeline.add_standard_decoding_stage()
|
||||||
|
|
||||||
|
self.assertEqual(pipeline.stages, [])
|
||||||
|
|
||||||
|
|
||||||
|
# Consolidated from test_disagg_pipeline_alignment.py.
|
||||||
|
class TestPipelineSpecificExtraModules(unittest.TestCase):
|
||||||
|
def _get_extra_modules(
|
||||||
|
self, pipeline_cls, role: RoleType, task_name: str
|
||||||
|
) -> set[str]:
|
||||||
|
pipeline = object.__new__(pipeline_cls)
|
||||||
|
return pipeline._get_extra_allowed_modules_for_role(role, task_name)
|
||||||
|
|
||||||
|
def test_flux_encoder_keeps_vae(self):
|
||||||
|
extras = self._get_extra_modules(Flux2Pipeline, RoleType.ENCODER, "ti2i")
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
Flux2Pipeline._required_config_modules,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"vae"})
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered), {"text_encoder", "tokenizer", "vae", "scheduler"}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_qwen_image_edit_encoder_keeps_vae(self):
|
||||||
|
extras = self._get_extra_modules(
|
||||||
|
QwenImageEditPipeline, RoleType.ENCODER, "ti2i"
|
||||||
|
)
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
QwenImageEditPipeline._required_config_modules,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"vae"})
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered),
|
||||||
|
{"processor", "scheduler", "text_encoder", "tokenizer", "vae"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_qwen_image_layered_encoder_keeps_required_cross_role_modules(self):
|
||||||
|
extras = self._get_extra_modules(
|
||||||
|
QwenImageLayeredPipeline, RoleType.ENCODER, "ti2i"
|
||||||
|
)
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
QwenImageLayeredPipeline._required_config_modules,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"vae", "transformer"})
|
||||||
|
self.assertNotIn(
|
||||||
|
"text_encoder", QwenImageLayeredPipeline._required_config_modules
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered),
|
||||||
|
{
|
||||||
|
"vae",
|
||||||
|
"tokenizer",
|
||||||
|
"processor",
|
||||||
|
"transformer",
|
||||||
|
"scheduler",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_glm_image_encoder_keeps_vae_and_transformer(self):
|
||||||
|
extras = self._get_extra_modules(GlmImagePipeline, RoleType.ENCODER, "ti2i")
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
GlmImagePipeline._required_config_modules,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"vae", "transformer"})
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered),
|
||||||
|
{
|
||||||
|
"text_encoder",
|
||||||
|
"tokenizer",
|
||||||
|
"vae",
|
||||||
|
"vision_language_encoder",
|
||||||
|
"processor",
|
||||||
|
"transformer",
|
||||||
|
"scheduler",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_wan_ti2v_denoiser_keeps_vae(self):
|
||||||
|
for pipeline_cls in (WanImageToVideoPipeline, WanImageToVideoDmdPipeline):
|
||||||
|
extras = self._get_extra_modules(pipeline_cls, RoleType.DENOISER, "ti2v")
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
pipeline_cls._required_config_modules,
|
||||||
|
RoleType.DENOISER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"vae"})
|
||||||
|
self.assertEqual(set(filtered), {"vae", "transformer", "scheduler"})
|
||||||
|
|
||||||
|
def test_ltx2_encoder_does_not_keep_decoder_modules(self):
|
||||||
|
extras = self._get_extra_modules(LTX2Pipeline, RoleType.ENCODER, "ti2v")
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
LTX2Pipeline._required_config_modules,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, set())
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered),
|
||||||
|
{"text_encoder", "tokenizer", "scheduler", "connectors"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_ltx2_ti2v_denoiser_keeps_vae_and_audio_vae(self):
|
||||||
|
extras = self._get_extra_modules(LTX2Pipeline, RoleType.DENOISER, "ti2v")
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
LTX2Pipeline._required_config_modules,
|
||||||
|
RoleType.DENOISER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"vae", "audio_vae"})
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered), {"transformer", "scheduler", "vae", "audio_vae"}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_mova_encoder_keeps_video_and_audio_vaes(self):
|
||||||
|
extras = self._get_extra_modules(MOVAPipeline, RoleType.ENCODER, "i2v")
|
||||||
|
filtered = filter_modules_for_role(
|
||||||
|
MOVAPipeline._required_config_modules,
|
||||||
|
RoleType.ENCODER,
|
||||||
|
extra_allowed_modules=extras,
|
||||||
|
)
|
||||||
|
self.assertEqual(extras, {"video_vae", "audio_vae"})
|
||||||
|
self.assertEqual(
|
||||||
|
set(filtered),
|
||||||
|
{"video_vae", "audio_vae", "text_encoder", "tokenizer", "scheduler"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_mova_alias_uses_same_encoder_extras(self):
|
||||||
|
extras = self._get_extra_modules(MOVAPipelineAlias, RoleType.ENCODER, "i2v")
|
||||||
|
self.assertEqual(extras, {"video_vae", "audio_vae"})
|
||||||
|
|
||||||
|
|
||||||
|
class TestQwenImageLayeredDtype(_GlobalStageArgsMixin, unittest.TestCase):
|
||||||
|
def test_text_encoder_dtype_uses_parameter_dtype_without_dtype_attr(self):
|
||||||
|
text_encoder = torch.nn.Linear(1, 1, bias=False).to(dtype=torch.bfloat16)
|
||||||
|
self.assertEqual(
|
||||||
|
_resolve_text_encoder_dtype(text_encoder),
|
||||||
|
torch.bfloat16,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_component_uses_keep_standard_text_encoder_and_configured_dtypes(self):
|
||||||
|
class _DummyVAE:
|
||||||
|
temperal_downsample = []
|
||||||
|
z_dim = 16
|
||||||
|
|
||||||
|
def to(self, *args, **kwargs):
|
||||||
|
return self
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered.get_local_torch_device",
|
||||||
|
return_value=torch.device("cpu"),
|
||||||
|
):
|
||||||
|
stage = QwenImageLayeredBeforeDenoisingStage(
|
||||||
|
vae=_DummyVAE(),
|
||||||
|
text_encoder=torch.nn.Linear(1, 1),
|
||||||
|
tokenizer=object(),
|
||||||
|
processor=object(),
|
||||||
|
transformer=object(),
|
||||||
|
scheduler=object(),
|
||||||
|
model_path="/unused",
|
||||||
|
vae_dtype=torch.float32,
|
||||||
|
text_encoder_dtype=torch.float16,
|
||||||
|
)
|
||||||
|
|
||||||
|
uses = stage.component_uses(SimpleNamespace(), "qwen_layered")
|
||||||
|
self.assertEqual(
|
||||||
|
[(use.component_name, use.target_dtype) for use in uses],
|
||||||
|
[("text_encoder", torch.float16), ("vae", torch.float32)],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestImageVAEEncodingStageComponentName(_GlobalStageArgsMixin, unittest.TestCase):
|
||||||
|
def test_component_name_can_follow_non_default_vae_key(self):
|
||||||
|
stage = ImageVAEEncodingStage(vae=object(), component_name="video_vae")
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
pipeline_config=SimpleNamespace(vae_precision="bf16")
|
||||||
|
)
|
||||||
|
|
||||||
|
uses = stage.component_uses(server_args, "image_vae_encoding")
|
||||||
|
self.assertEqual(len(uses), 1)
|
||||||
|
self.assertEqual(uses[0].component_name, "video_vae")
|
||||||
|
self.assertEqual(uses[0].target_dtype, torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStageAffinityAndValidation(_GlobalStageArgsMixin, unittest.TestCase):
|
||||||
|
def _make_hunyuan_pipeline(
|
||||||
|
self, role: RoleType, *, paint_enable: bool
|
||||||
|
) -> Hunyuan3D2Pipeline:
|
||||||
|
pipeline = object.__new__(Hunyuan3D2Pipeline)
|
||||||
|
pipeline.server_args = self._install_stage_server_args(
|
||||||
|
pipeline_config=Hunyuan3D2PipelineConfig(paint_enable=paint_enable)
|
||||||
|
)
|
||||||
|
pipeline._disagg_role = role
|
||||||
|
pipeline.modules = {
|
||||||
|
"hy3dshape_image_processor": object(),
|
||||||
|
"hy3dshape_conditioner": object(),
|
||||||
|
"hy3dshape_scheduler": object(),
|
||||||
|
"hy3dshape_model": torch.nn.Linear(1, 1),
|
||||||
|
"hy3dshape_vae": object(),
|
||||||
|
}
|
||||||
|
pipeline._stages = []
|
||||||
|
pipeline._stage_name_mapping = {}
|
||||||
|
return pipeline
|
||||||
|
|
||||||
|
def test_helios_denoising_stage_is_denoiser_affine(self):
|
||||||
|
stage = object.__new__(HeliosChunkedDenoisingStage)
|
||||||
|
self.assertEqual(stage.role_affinity, RoleType.DENOISER)
|
||||||
|
|
||||||
|
def test_mova_denoising_stage_is_denoiser_affine(self):
|
||||||
|
stage = object.__new__(MOVADenoisingStage)
|
||||||
|
self.assertEqual(stage.role_affinity, RoleType.DENOISER)
|
||||||
|
|
||||||
|
def test_mova_decoding_stage_is_decoder_affine(self):
|
||||||
|
stage = object.__new__(MOVADecodingStage)
|
||||||
|
self.assertEqual(stage.role_affinity, RoleType.DECODER)
|
||||||
|
|
||||||
|
def test_mova_skips_torch_compile_on_rocm(self):
|
||||||
|
class _CompileTrackingModule(torch.nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.compile_called = False
|
||||||
|
|
||||||
|
def compile(self, *args, **kwargs):
|
||||||
|
self.compile_called = True
|
||||||
|
|
||||||
|
stage = object.__new__(MOVADenoisingStage)
|
||||||
|
module = _CompileTrackingModule()
|
||||||
|
server_args = SimpleNamespace(enable_torch_compile=True)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.mova.current_platform.is_hip",
|
||||||
|
return_value=True,
|
||||||
|
):
|
||||||
|
stage._maybe_enable_torch_compile(module, server_args)
|
||||||
|
|
||||||
|
self.assertFalse(module.compile_called)
|
||||||
|
|
||||||
|
def test_hunyuan3d_shape_only_disagg_accepts_non_monolithic_roles(self):
|
||||||
|
pipeline = self._make_hunyuan_pipeline(RoleType.ENCODER, paint_enable=False)
|
||||||
|
pipeline.validate_disagg_role(RoleType.ENCODER)
|
||||||
|
pipeline.validate_disagg_role(RoleType.MONOLITHIC)
|
||||||
|
|
||||||
|
def test_hunyuan3d_disagg_rejects_paint_pipeline(self):
|
||||||
|
pipeline = self._make_hunyuan_pipeline(RoleType.ENCODER, paint_enable=True)
|
||||||
|
with self.assertRaisesRegex(ValueError, "shape-only disaggregation"):
|
||||||
|
pipeline.validate_disagg_role(RoleType.ENCODER)
|
||||||
|
|
||||||
|
def test_hunyuan3d_shape_export_and_save_are_decoder_affine(self):
|
||||||
|
export_stage = Hunyuan3DShapeExportStage(
|
||||||
|
vae=object(),
|
||||||
|
config=Hunyuan3D2PipelineConfig(paint_enable=False),
|
||||||
|
)
|
||||||
|
save_stage = Hunyuan3DShapeSaveStage(
|
||||||
|
config=Hunyuan3D2PipelineConfig(paint_enable=False),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(export_stage.role_affinity, RoleType.DECODER)
|
||||||
|
self.assertEqual(save_stage.role_affinity, RoleType.DECODER)
|
||||||
|
|
||||||
|
def test_hunyuan3d_stage_filtering_matches_shape_only_roles(self):
|
||||||
|
expected = {
|
||||||
|
RoleType.ENCODER: ["shape_before_denoising"],
|
||||||
|
RoleType.DENOISER: ["shape_denoising"],
|
||||||
|
RoleType.DECODER: ["shape_export", "shape_save"],
|
||||||
|
}
|
||||||
|
|
||||||
|
for role, stage_names in expected.items():
|
||||||
|
pipeline = self._make_hunyuan_pipeline(role, paint_enable=False)
|
||||||
|
pipeline.create_pipeline_stages(pipeline.server_args)
|
||||||
|
self.assertEqual(list(pipeline._stage_name_mapping.keys()), stage_names)
|
||||||
|
|
||||||
|
def test_hunyuan3d_shape_stage_no_longer_stores_model_dtype(self):
|
||||||
|
pipeline = self._make_hunyuan_pipeline(RoleType.ENCODER, paint_enable=False)
|
||||||
|
pipeline.create_pipeline_stages(pipeline.server_args)
|
||||||
|
stage = pipeline._stage_name_mapping["shape_before_denoising"]
|
||||||
|
self.assertIsInstance(stage, Hunyuan3DShapeBeforeDenoisingStage)
|
||||||
|
self.assertFalse(hasattr(stage, "model_dtype"))
|
||||||
|
|
||||||
|
def test_ltx2_refinement_stage_keeps_class_name_stage_key(self):
|
||||||
|
stage = object.__new__(LTX2RefinementStage)
|
||||||
|
self.assertEqual(
|
||||||
|
ComposedPipelineBase._infer_stage_name(stage), "LTX2RefinementStage"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestHunyuan3DShapeStageRuntimeDtype(_GlobalStageArgsMixin, unittest.TestCase):
|
||||||
|
def test_conditioner_parameter_dtype_wins_over_sample_dtype(self):
|
||||||
|
conditioner = torch.nn.Linear(4, 4, bias=False).to(dtype=torch.float32)
|
||||||
|
stage = Hunyuan3DShapeBeforeDenoisingStage(
|
||||||
|
image_processor=object(),
|
||||||
|
conditioner=conditioner,
|
||||||
|
scheduler=SimpleNamespace(init_noise_sigma=1.0),
|
||||||
|
config=Hunyuan3D2PipelineConfig(),
|
||||||
|
latent_shape=(1, 2, 2),
|
||||||
|
guidance_embed=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
stage._resolve_runtime_dtype(torch.zeros(1, dtype=torch.float16)),
|
||||||
|
torch.float32,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_runtime_dtype_falls_back_to_sample_tensor_without_module_dtype(self):
|
||||||
|
stage = Hunyuan3DShapeBeforeDenoisingStage(
|
||||||
|
image_processor=object(),
|
||||||
|
conditioner=object(),
|
||||||
|
scheduler=SimpleNamespace(init_noise_sigma=1.0),
|
||||||
|
config=Hunyuan3D2PipelineConfig(),
|
||||||
|
latent_shape=(1, 2, 2),
|
||||||
|
guidance_embed=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
stage._resolve_runtime_dtype(torch.zeros(1, dtype=torch.bfloat16)),
|
||||||
|
torch.bfloat16,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_runtime_dtype_warns_and_falls_back_after_non_iterable_parameters(self):
|
||||||
|
conditioner = SimpleNamespace(
|
||||||
|
parameters=lambda: (_ for _ in ()).throw(TypeError("not iterable")),
|
||||||
|
buffers=lambda: iter(()),
|
||||||
|
)
|
||||||
|
stage = Hunyuan3DShapeBeforeDenoisingStage(
|
||||||
|
image_processor=object(),
|
||||||
|
conditioner=conditioner,
|
||||||
|
scheduler=SimpleNamespace(init_noise_sigma=1.0),
|
||||||
|
config=Hunyuan3D2PipelineConfig(),
|
||||||
|
latent_shape=(1, 2, 2),
|
||||||
|
guidance_embed=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.pipelines_core.stages.hunyuan3d_shape.logger.warning"
|
||||||
|
) as mock_warning:
|
||||||
|
dtype = stage._resolve_runtime_dtype(torch.zeros(1, dtype=torch.float16))
|
||||||
|
|
||||||
|
self.assertEqual(dtype, torch.float16)
|
||||||
|
mock_warning.assert_called_once()
|
||||||
|
self.assertEqual(mock_warning.call_args.args[1], "parameters")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user