[diffusion] feat: support role-based component loading and stage affinity (#25168)

This commit is contained in:
Heyang Huang
2026-05-22 18:50:23 +08:00
committed by GitHub
parent e1dcbca220
commit 8c916a715c
12 changed files with 989 additions and 81 deletions
@@ -25,7 +25,7 @@ class RoleType(str, Enum):
@classmethod
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":
@@ -37,32 +37,53 @@ def get_module_role(module_name: str) -> "RoleType | None":
"image_processor",
"processor",
"connectors",
"vision_language_encoder",
)
if any(
module_name == p or module_name.startswith(p + "_") for p in encoder_prefixes
):
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(
module_name == p or module_name.startswith(p + "_") for p in denoising_prefixes
):
return RoleType.DENOISER
if module_name == "hy3dshape_model":
return RoleType.DENOISER
decoder_prefixes = ("vae", "audio_vae", "video_vae", "vocoder")
if any(
module_name == p or module_name.startswith(p + "_") for p in decoder_prefixes
):
return RoleType.DECODER
if module_name == "hy3dshape_vae":
return RoleType.DECODER
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."""
if role in (RoleType.MONOLITHIC, RoleType.SERVER):
return module_names
extra_allowed_modules = extra_allowed_modules or set()
filtered = []
for name in module_names:
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)
elif module_role == role:
filtered.append(name)
elif role == RoleType.ENCODER and module_role == RoleType.DECODER:
# Encoder also needs VAE for ImageVAEEncoding stages
elif name in extra_allowed_modules:
filtered.append(name)
return filtered
@@ -19,6 +19,7 @@ import torch.nn as nn
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
Hunyuan3D2PipelineConfig,
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
load_model_from_full_model_state_dict,
set_default_torch_dtype,
@@ -58,6 +59,21 @@ class Hunyuan3D2Pipeline(ComposedPipelineBase):
"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]:
return {
"_class_name": self.pipeline_name,
@@ -357,6 +373,8 @@ class Hunyuan3D2Pipeline(ComposedPipelineBase):
def create_pipeline_stages(self, server_args: ServerArgs):
config = server_args.pipeline_config
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
self.add_stage(
@@ -364,10 +382,10 @@ class Hunyuan3D2Pipeline(ComposedPipelineBase):
stage=Hunyuan3DShapeBeforeDenoisingStage(
image_processor=self.get_module("hy3dshape_image_processor"),
conditioner=self.get_module("hy3dshape_conditioner"),
vae=self.get_module("hy3dshape_vae"),
model=self.get_module("hy3dshape_model"),
scheduler=self.get_module("hy3dshape_scheduler"),
config=config,
latent_shape=latent_shape,
guidance_embed=guidance_embed,
),
)
self.add_stage(
@@ -63,7 +63,12 @@ class MOVAPipeline(ComposedPipelineBase):
self.add_stage(InputValidationStage())
self.add_standard_text_encoding_stage()
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(
MOVALatentPreparationStage(
audio_vae=self.get_module("audio_vae"),
@@ -3,6 +3,7 @@
# SPDX-License-Identifier: Apache-2.0
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.composed_pipeline_base import (
ComposedPipelineBase,
@@ -115,9 +116,10 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(
QwenImageLayeredBeforeDenoisingStage(
def create_before_denoising_stage():
return QwenImageLayeredBeforeDenoisingStage(
vae=self.get_module("vae"),
text_encoder=None,
tokenizer=self.get_module("tokenizer"),
processor=self.get_module("processor"),
transformer=self.get_module("transformer"),
@@ -128,6 +130,11 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
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(
@@ -99,6 +99,7 @@ class ComposedPipelineBase(ABC):
"""
self.server_args = server_args
self._disagg_role = server_args.disagg_role
self.validate_disagg_role(self._disagg_role)
self.model_path: str = model_path
self._stages: list[PipelineStage] = []
@@ -107,17 +108,26 @@ class ComposedPipelineBase(ABC):
self.executor = executor or self.build_executor(server_args=server_args)
self.component_residency_manager: ComponentResidencyManager | None = None
if required_config_modules is not None:
self._required_config_modules = required_config_modules
if self._required_config_modules is None:
base_required_config_modules = (
required_config_modules
if required_config_modules is not None
else self._required_config_modules
)
if base_required_config_modules is None:
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
if self._disagg_role != RoleType.MONOLITHIC:
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, 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)
if skipped:
@@ -202,6 +212,44 @@ class ComposedPipelineBase(ABC):
"""
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_ATTR_MAP: dict[str, str] = {
"vae": "vae_config",
@@ -495,6 +543,22 @@ class ComposedPipelineBase(ABC):
self.memory_usages,
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
@@ -502,6 +566,24 @@ class ComposedPipelineBase(ABC):
def _infer_stage_name(stage: PipelineStage) -> str:
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:
class_name = stage.__class__.__name__
if any(existing.__class__.__name__ == class_name for existing in self._stages):
@@ -513,22 +595,13 @@ class ComposedPipelineBase(ABC):
) -> "ComposedPipelineBase":
assert self.modules is not None, "No modules are registered"
# Filter stages based on disaggregation role
if self._disagg_role != RoleType.MONOLITHIC:
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
if stage_name is None:
stage_name = self._infer_stage_name(stage)
# Filter stages based on disaggregation role
if not self._should_add_stage_for_role(stage.role_affinity, stage_name):
return self
if stage_name in self._stage_name_mapping:
raise ValueError(f"Duplicate stage name detected: {stage_name}")
@@ -538,6 +611,17 @@ class ComposedPipelineBase(ABC):
self._stage_name_mapping[stage_name] = stage
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(
self, stages: list[PipelineStage | tuple[PipelineStage, str]]
) -> "ComposedPipelineBase":
@@ -579,12 +663,12 @@ class ComposedPipelineBase(ABC):
def add_standard_timestep_preparation_stage(
self,
scheduler_key: str = "scheduler",
prepare_extra_kwargs: list[Callable] | None = [],
prepare_extra_kwargs: list[Callable] | None = None,
) -> "ComposedPipelineBase":
return self.add_stage(
TimestepPreparationStage(
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,43 +690,57 @@ class ComposedPipelineBase(ABC):
transformer_2_key: str | None = "transformer_2",
scheduler_key: str = "scheduler",
vae_key: str | None = "vae",
stage_name: str = "denoising_stage",
) -> "ComposedPipelineBase":
kwargs = {
"transformer": self.get_module(transformer_key),
"scheduler": self.get_module(scheduler_key),
}
def create_stage() -> PipelineStage:
kwargs = {
"transformer": self.get_module(transformer_key),
"scheduler": self.get_module(scheduler_key),
}
if transformer_2_key:
transformer_2 = self.get_module(transformer_2_key, None)
if transformer_2 is not None:
kwargs["transformer_2"] = transformer_2
if transformer_2_key:
transformer_2 = self.get_module(transformer_2_key, None)
if transformer_2 is not None:
kwargs["transformer_2"] = transformer_2
if vae_key:
vae = self.get_module(vae_key, None)
if vae is not None:
kwargs["vae"] = vae
kwargs["pipeline"] = self
if vae_key:
vae = self.get_module(vae_key, None)
if vae is not None:
kwargs["vae"] = vae
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(
self,
vae_key: str = "vae",
stage_name: str = "decoding_stage",
) -> "ComposedPipelineBase":
return self.add_stage(
DecodingStage(
def create_stage() -> PipelineStage:
return DecodingStage(
vae=self.get_module(vae_key),
pipeline=self,
component_name=vae_key,
),
)
return self.add_stage_factory(
RoleType.DECODER,
create_stage,
stage_name,
)
def add_standard_t2i_stages(
self,
include_input_validation: bool = True,
prepare_extra_timestep_kwargs: list[Callable] | None = [],
prepare_extra_timestep_kwargs: list[Callable] | None = None,
) -> "ComposedPipelineBase":
if include_input_validation:
@@ -671,7 +769,7 @@ class ComposedPipelineBase(ABC):
prompt_text_encoder_key: str = "text_encoder",
image_vae_key: str = "vae",
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":
if include_input_validation:
self.add_stage(
@@ -696,7 +794,10 @@ class ComposedPipelineBase(ABC):
self.add_stage(
ImageVAEEncodingStage(
vae=self.get_module(image_vae_key),
**(image_vae_stage_kwargs or {}),
**{
"component_name": image_vae_key,
**(image_vae_stage_kwargs or {}),
},
),
)
@@ -723,8 +824,9 @@ class ComposedPipelineBase(ABC):
image_vae_encoding_position: Literal[
"before_timestep", "after_latent"
] = "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_name: str = "denoising_stage",
) -> "ComposedPipelineBase":
if include_input_validation:
self.add_stage(
@@ -750,7 +852,10 @@ class ComposedPipelineBase(ABC):
self.add_stage(
ImageVAEEncodingStage(
vae=self.get_module(image_vae_key),
**(image_vae_stage_kwargs or {}),
**{
"component_name": image_vae_key,
**(image_vae_stage_kwargs or {}),
},
)
)
@@ -762,7 +867,10 @@ class ComposedPipelineBase(ABC):
self.add_stage(
ImageVAEEncodingStage(
vae=self.get_module(image_vae_key),
**(image_vae_stage_kwargs or {}),
**{
"component_name": image_vae_key,
**(image_vae_stage_kwargs or {}),
},
)
)
elif image_vae_encoding_position != "before_timestep":
@@ -773,7 +881,11 @@ class ComposedPipelineBase(ABC):
if denoising_stage_factory is None:
self.add_standard_denoising_stage()
else:
self.add_stage(denoising_stage_factory())
self.add_stage_factory(
RoleType.DENOISER,
denoising_stage_factory,
denoising_stage_name,
)
self.add_standard_decoding_stage()
return self
@@ -128,18 +128,18 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
self,
image_processor: Any,
conditioner: Any,
vae: Any,
model: Any,
scheduler: Any,
config: Hunyuan3D2PipelineConfig,
latent_shape: tuple[int, ...],
guidance_embed: bool,
) -> None:
super().__init__()
self.image_processor = image_processor
self.conditioner = conditioner
self.vae = vae
self.model = model
self.scheduler = scheduler
self.config = config
self.latent_shape = latent_shape
self.guidance_embed = guidance_embed
def _validate_input(self, batch: Req, server_args: ServerArgs) -> None:
if batch.image_path is None:
@@ -160,10 +160,41 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
def _prepare_latents(self, batch_size, dtype, device, generator, scheduler):
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)
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:
# 1. Input validation
self._validate_input(batch, server_args)
@@ -173,14 +204,14 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
image = cond_inputs.pop("image")
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)
cond_inputs = _move_to_device(cond_inputs, device, dtype)
# 3. Conditioning with CFG
do_cfg = batch.guidance_scale >= 0 and not (
hasattr(self.model, "guidance_embed") and self.model.guidance_embed is True
)
do_cfg = batch.guidance_scale >= 0 and not self.guidance_embed
cond = self.conditioner(image=image, **cond_inputs)
if do_cfg:
@@ -216,7 +247,7 @@ class Hunyuan3DShapeBeforeDenoisingStage(PipelineStage):
latents = self._prepare_latents(batch_size, dtype, device, generator, scheduler)
guidance = None
if hasattr(self.model, "guidance_embed") and self.model.guidance_embed is True:
if self.guidance_embed:
guidance = torch.tensor(
[batch.guidance_scale] * batch_size, device=device, dtype=dtype
)
@@ -416,6 +447,12 @@ class Hunyuan3DShapeExportStage(PipelineStage):
self.vae = vae
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:
if self.config.shape_mc_algo is not None:
try:
@@ -473,6 +510,12 @@ class Hunyuan3DShapeSaveStage(PipelineStage):
super().__init__()
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]:
output_path = batch.output_file_path() or os.path.join(
batch.output_path, "output.obj"
@@ -803,9 +803,15 @@ class ImageVAEEncodingStage(PipelineStage):
"vae_image_sizes",
)
def __init__(self, vae: ParallelTiledVAE, **kwargs) -> None:
def __init__(
self,
vae: ParallelTiledVAE,
component_name: str = "vae",
**kwargs,
) -> None:
super().__init__()
self.vae: ParallelTiledVAE = vae
self.component_name = component_name
def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None
@@ -815,7 +821,7 @@ class ImageVAEEncodingStage(PipelineStage):
return [
ComponentUse(
stage_name,
"vae",
self.component_name,
target_dtype=vae_dtype,
)
]
@@ -851,7 +857,10 @@ class ImageVAEEncodingStage(PipelineStage):
vae_dtype != torch.float32
) 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
self.vae = vae
@@ -13,6 +13,7 @@ import numpy as np
import torch
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.memory_managers.component_manager import (
ComponentUse,
@@ -105,6 +106,10 @@ class HeliosChunkedDenoisingStage(PipelineStage):
self.transformer = transformer
self.scheduler = scheduler
@property
def role_affinity(self) -> RoleType:
return RoleType.DENOISER
@property
def parallelism_type(self):
return StageParallelismType.REPLICATED
@@ -21,6 +21,7 @@ import torch.nn as nn
from diffusers.utils.torch_utils import randn_tensor
from tqdm.auto import tqdm
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device,
get_world_group,
@@ -187,6 +188,10 @@ class MOVADenoisingStage(PipelineStage):
)
return uses
@property
def role_affinity(self) -> RoleType:
return RoleType.DENOISER
@property
def parallelism_type(self) -> StageParallelismType:
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):
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}
if current_platform.is_npu():
@@ -954,6 +966,10 @@ class MOVADecodingStage(PipelineStage):
ComponentUse(stage_name, "audio_vae"),
]
@property
def role_affinity(self) -> RoleType:
return RoleType.DECODER
@property
def parallelism_type(self) -> StageParallelismType:
if get_global_server_args().enable_cfg_parallel:
@@ -21,6 +21,31 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
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(
prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None
) -> list[int]:
@@ -125,6 +150,7 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
def __init__(
self,
vae,
text_encoder,
tokenizer,
processor,
transformer,
@@ -137,14 +163,14 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
self.vae = vae.to(dtype=vae_dtype)
self.vae_dtype = vae_dtype
self.text_encoder_dtype = text_encoder_dtype
from transformers import Qwen2_5_VLForConditionalGeneration
if text_encoder is None:
from transformers import Qwen2_5_VLForConditionalGeneration
self.text_encoder = (
Qwen2_5_VLForConditionalGeneration.from_pretrained(
text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained(
model_path, subfolder="text_encoder"
)
.to(get_local_torch_device())
.to(dtype=self.text_encoder_dtype)
self.text_encoder = text_encoder.to(
device=get_local_torch_device(), dtype=self.text_encoder_dtype
)
self.tokenizer = tokenizer
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)
return [
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
@@ -232,7 +264,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
device: Optional[torch.device] = 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
@@ -482,7 +514,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
prompt = batch.prompt
with self.use_declared_component(
component_name="qwen_layered_text_encoder",
component_name="text_encoder",
module=self.text_encoder,
) as text_encoder:
assert text_encoder is not None
@@ -60,13 +60,13 @@ class TimestepPreparationStage(PipelineStage):
def __init__(
self,
scheduler,
prepare_extra_set_timesteps_kwargs: list[
Callable[[Req, ServerArgs], Tuple[str, Any]]
] = [],
prepare_extra_set_timesteps_kwargs: (
list[Callable[[Req, ServerArgs], Tuple[str, Any]]] | None
) = None,
) -> None:
super().__init__()
self.scheduler = scheduler
self.prepare_extra_set_timesteps_kwargs = (
self.prepare_extra_set_timesteps_kwargs = list(
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()