[diffusion] chore: preserve exact component identity during loading (#36875)
This commit is contained in:
@@ -7,7 +7,7 @@ import os
|
|||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from dataclasses import asdict, dataclass, field, fields
|
from dataclasses import asdict, dataclass, field, fields
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from typing import Any
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import PIL
|
import PIL
|
||||||
@@ -198,6 +198,7 @@ def maybe_unpad_latents(latents, batch):
|
|||||||
class PipelineConfig:
|
class PipelineConfig:
|
||||||
"""The base configuration class for a generation pipeline."""
|
"""The base configuration class for a generation pipeline."""
|
||||||
|
|
||||||
|
native_only_components: ClassVar[tuple[str, ...]] = ()
|
||||||
task_type: ModelTaskType = ModelTaskType.I2I
|
task_type: ModelTaskType = ModelTaskType.I2I
|
||||||
skip_input_image_preprocess: bool = False
|
skip_input_image_preprocess: bool = False
|
||||||
# Components that cannot fall back to a native Transformers/Diffusers
|
# Components that cannot fall back to a native Transformers/Diffusers
|
||||||
|
|||||||
@@ -1,11 +1,19 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
"""Role definitions for diffusion pipeline disaggregation."""
|
"""Role definitions for diffusion pipeline disaggregation."""
|
||||||
|
|
||||||
|
from collections.abc import Mapping
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|
||||||
_ROLE_ALIASES = {"denoising": "denoiser"}
|
_ROLE_ALIASES = {"denoising": "denoiser"}
|
||||||
|
|
||||||
|
|
||||||
|
def _matches_component_type(component_name: str, component_type: str) -> bool:
|
||||||
|
prefix = f"{component_type}_"
|
||||||
|
return component_name == component_type or (
|
||||||
|
component_name.startswith(prefix) and component_name[len(prefix) :].isdigit()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class RoleType(str, Enum):
|
class RoleType(str, Enum):
|
||||||
MONOLITHIC = "monolithic"
|
MONOLITHIC = "monolithic"
|
||||||
ENCODER = "encoder"
|
ENCODER = "encoder"
|
||||||
@@ -86,21 +94,31 @@ def filter_modules_for_role(
|
|||||||
role: "RoleType",
|
role: "RoleType",
|
||||||
*,
|
*,
|
||||||
extra_allowed_modules: set[str] | None = None,
|
extra_allowed_modules: set[str] | None = None,
|
||||||
|
structural_component_names: Mapping[str, str] | None = None,
|
||||||
) -> list[str]:
|
) -> 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()
|
extra_allowed_modules = extra_allowed_modules or set()
|
||||||
|
structural_component_names = structural_component_names or {}
|
||||||
filtered = []
|
filtered = []
|
||||||
for name in module_names:
|
for name in module_names:
|
||||||
module_role = get_module_role(name)
|
structural_name = structural_component_names.get(name, name)
|
||||||
|
module_role = get_module_role(structural_name)
|
||||||
|
|
||||||
if module_role is None:
|
if module_role is None:
|
||||||
filtered.append(name)
|
filtered.append(name)
|
||||||
elif module_role == role:
|
elif module_role == role:
|
||||||
filtered.append(name)
|
filtered.append(name)
|
||||||
elif name in extra_allowed_modules:
|
elif (
|
||||||
|
name in extra_allowed_modules
|
||||||
|
or structural_name in extra_allowed_modules
|
||||||
|
or any(
|
||||||
|
_matches_component_type(structural_name, component_type)
|
||||||
|
for component_type in extra_allowed_modules
|
||||||
|
)
|
||||||
|
):
|
||||||
filtered.append(name)
|
filtered.append(name)
|
||||||
|
|
||||||
return filtered
|
return filtered
|
||||||
|
|||||||
@@ -71,7 +71,8 @@ class AdapterLoader(PlainStateDictComponentLoader):
|
|||||||
server_args, component_name, precision_attr="dit_precision"
|
server_args, component_name, precision_attr="dit_precision"
|
||||||
)
|
)
|
||||||
|
|
||||||
config_cls = self._CONFIG_CLASSES[component_name]
|
component_type = self.structural_component_type(component_name)
|
||||||
|
config_cls = self._CONFIG_CLASSES[component_type]
|
||||||
with set_default_torch_dtype(default_dtype), skip_init_modules():
|
with set_default_torch_dtype(default_dtype), skip_init_modules():
|
||||||
adapter_cfg = config_cls()
|
adapter_cfg = config_cls()
|
||||||
adapter_cfg.update_model_arch(config)
|
adapter_cfg.update_model_arch(config)
|
||||||
|
|||||||
@@ -156,8 +156,17 @@ class ComponentLoader(ABC):
|
|||||||
def __init__(self, device=None) -> None:
|
def __init__(self, device=None) -> None:
|
||||||
self.device = device
|
self.device = device
|
||||||
self.component_architecture: str | None = None
|
self.component_architecture: str | None = None
|
||||||
|
self.component_type: str | None = None
|
||||||
self._native_load_manages_placement = False
|
self._native_load_manages_placement = False
|
||||||
|
|
||||||
|
def structural_component_name(self, component_name: str) -> str:
|
||||||
|
"""Return the config slot without changing the exact policy key."""
|
||||||
|
return self.component_type or component_name
|
||||||
|
|
||||||
|
def structural_component_type(self, component_name: str) -> str:
|
||||||
|
"""Return the normalized loader role for an exact component key."""
|
||||||
|
return _normalize_component_type(self.structural_component_name(component_name))
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def target_device(component_starts_on_cpu: bool) -> torch.device:
|
def target_device(component_starts_on_cpu: bool) -> torch.device:
|
||||||
if component_starts_on_cpu:
|
if component_starts_on_cpu:
|
||||||
@@ -190,10 +199,23 @@ class ComponentLoader(ABC):
|
|||||||
) -> bool:
|
) -> bool:
|
||||||
return self.supports_direct_gpu_weight_loading
|
return self.supports_direct_gpu_weight_loading
|
||||||
|
|
||||||
|
def is_native_only_component(
|
||||||
|
self, server_args: ServerArgs, component_name: str
|
||||||
|
) -> bool:
|
||||||
|
native_only_components = server_args.pipeline_config.native_only_components
|
||||||
|
return any(
|
||||||
|
name in native_only_components
|
||||||
|
for name in (
|
||||||
|
component_name,
|
||||||
|
self.structural_component_name(component_name),
|
||||||
|
self.structural_component_type(component_name),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def should_raise_customized_load_error(
|
def should_raise_customized_load_error(
|
||||||
self, server_args: ServerArgs, component_name: str
|
self, server_args: ServerArgs, component_name: str
|
||||||
) -> bool:
|
) -> bool:
|
||||||
return component_name in server_args.pipeline_config.native_only_components
|
return self.is_native_only_component(server_args, component_name)
|
||||||
|
|
||||||
def validate_native_fallback(
|
def validate_native_fallback(
|
||||||
self, _server_args: ServerArgs, _component_name: str
|
self, _server_args: ServerArgs, _component_name: str
|
||||||
@@ -429,11 +451,19 @@ class ComponentLoader(ABC):
|
|||||||
"""
|
"""
|
||||||
Load the component using the native library (transformers/diffusers).
|
Load the component using the native library (transformers/diffusers).
|
||||||
"""
|
"""
|
||||||
precision = (
|
precision = None
|
||||||
resolve_component_precision(server_args, component_name)
|
if component_name is not None:
|
||||||
if component_name is not None
|
precision_names = dict.fromkeys(
|
||||||
else None
|
(
|
||||||
|
component_name,
|
||||||
|
self.structural_component_name(component_name),
|
||||||
|
self.structural_component_type(component_name),
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
for precision_name in precision_names:
|
||||||
|
precision = resolve_component_precision(server_args, precision_name)
|
||||||
|
if precision is not None:
|
||||||
|
break
|
||||||
load_kwargs = {}
|
load_kwargs = {}
|
||||||
if precision is not None:
|
if precision is not None:
|
||||||
load_kwargs["torch_dtype"] = precision
|
load_kwargs["torch_dtype"] = precision
|
||||||
@@ -571,7 +601,7 @@ class ComponentLoader(ABC):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def for_component_type(
|
def for_component_type(
|
||||||
cls,
|
cls,
|
||||||
component_name: str,
|
component_type: str,
|
||||||
transformers_or_diffusers: str,
|
transformers_or_diffusers: str,
|
||||||
component_architecture: str | None = None,
|
component_architecture: str | None = None,
|
||||||
) -> "ComponentLoader":
|
) -> "ComponentLoader":
|
||||||
@@ -579,37 +609,43 @@ class ComponentLoader(ABC):
|
|||||||
Factory method to create a component loader for a specific component type.
|
Factory method to create a component loader for a specific component type.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
component_name: Type of component (e.g., "vae", "text_encoder", "transformer", "scheduler")
|
component_type: Structural role (e.g. "vae" or "text_encoder")
|
||||||
transformers_or_diffusers: Whether the component is from transformers or diffusers
|
transformers_or_diffusers: Whether the component is from transformers or diffusers
|
||||||
"""
|
"""
|
||||||
cls._ensure_loaders_registered()
|
cls._ensure_loaders_registered()
|
||||||
|
|
||||||
# Map of component types to their loader classes and expected library
|
# Map of component types to their loader classes and expected library
|
||||||
component_name = _normalize_component_type(component_name)
|
structural_component_name = component_type
|
||||||
|
loader_type = _normalize_component_type(component_type)
|
||||||
|
|
||||||
transformers_or_diffusers = cls.resolve_transformers_or_diffusers(
|
transformers_or_diffusers = cls.resolve_transformers_or_diffusers(
|
||||||
transformers_or_diffusers, component_name
|
transformers_or_diffusers, loader_type
|
||||||
)
|
)
|
||||||
|
|
||||||
if component_name in component_name_to_loader_cls:
|
if loader_type in component_name_to_loader_cls:
|
||||||
loader_cls: Type[ComponentLoader] = component_name_to_loader_cls[
|
loader_cls: Type[ComponentLoader] = component_name_to_loader_cls[
|
||||||
component_name
|
loader_type
|
||||||
]
|
]
|
||||||
expected_library = loader_cls.expected_library
|
expected_library = loader_cls.expected_library
|
||||||
# Assert that the library matches what's expected for this component type
|
# Assert that the library matches what's expected for this component type
|
||||||
assert (
|
assert (
|
||||||
transformers_or_diffusers == expected_library
|
transformers_or_diffusers == expected_library
|
||||||
), f"{component_name} must be loaded from {expected_library}, got {transformers_or_diffusers}"
|
), f"{loader_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
|
||||||
loader = loader_cls()
|
loader = loader_cls()
|
||||||
|
loader.component_type = structural_component_name
|
||||||
loader.component_architecture = component_architecture
|
loader.component_architecture = component_architecture
|
||||||
return loader
|
return loader
|
||||||
|
|
||||||
# For unknown component types, use a generic loader
|
# For unknown component types, use a generic loader
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"No specific loader found for component type: %s. Using generic loader.",
|
"No specific loader found for component type: %s. Using generic loader.",
|
||||||
component_name,
|
loader_type,
|
||||||
)
|
)
|
||||||
return GenericComponentLoader(transformers_or_diffusers, component_architecture)
|
loader = GenericComponentLoader(
|
||||||
|
transformers_or_diffusers, component_architecture
|
||||||
|
)
|
||||||
|
loader.component_type = structural_component_name
|
||||||
|
return loader
|
||||||
|
|
||||||
|
|
||||||
class PlainStateDictComponentLoader(ComponentLoader):
|
class PlainStateDictComponentLoader(ComponentLoader):
|
||||||
@@ -759,6 +795,7 @@ class PipelineComponentLoader:
|
|||||||
component_architecture: str | None = None,
|
component_architecture: str | None = None,
|
||||||
component_attn_backend: Any = None,
|
component_attn_backend: Any = None,
|
||||||
component_attn_name: str | None = None,
|
component_attn_name: str | None = None,
|
||||||
|
component_type: str | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Load a pipeline component.
|
Load a pipeline component.
|
||||||
@@ -768,11 +805,14 @@ class PipelineComponentLoader:
|
|||||||
component_model_path: Path to the component model
|
component_model_path: Path to the component model
|
||||||
transformers_or_diffusers: Whether the component is from transformers or diffusers
|
transformers_or_diffusers: Whether the component is from transformers or diffusers
|
||||||
component_architecture: the class name of the module
|
component_architecture: the class name of the module
|
||||||
|
component_type: structural config slot when it differs from the exact key
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Get the appropriate loader for this component type
|
# Get the appropriate loader for this component type
|
||||||
loader = ComponentLoader.for_component_type(
|
loader = ComponentLoader.for_component_type(
|
||||||
component_name, transformers_or_diffusers, component_architecture
|
component_type or component_name,
|
||||||
|
transformers_or_diffusers,
|
||||||
|
component_architecture,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
+3
-1
@@ -744,7 +744,9 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# TODO(mick): had to throw an exception for different text-encoder arch
|
# TODO(mick): had to throw an exception for different text-encoder arch
|
||||||
encoder_index = self._extract_encoder_index(component_name)
|
encoder_index = self._extract_encoder_index(
|
||||||
|
self.structural_component_name(component_name)
|
||||||
|
)
|
||||||
assert encoder_index < len(
|
assert encoder_index < len(
|
||||||
server_args.pipeline_config.text_encoder_configs
|
server_args.pipeline_config.text_encoder_configs
|
||||||
) and encoder_index < len(server_args.pipeline_config.text_encoder_precisions)
|
) and encoder_index < len(server_args.pipeline_config.text_encoder_precisions)
|
||||||
|
|||||||
+33
-26
@@ -29,7 +29,6 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
|||||||
resolve_transformer_gguf_to_load,
|
resolve_transformer_gguf_to_load,
|
||||||
resolve_transformer_quant_load_spec,
|
resolve_transformer_quant_load_spec,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import _normalize_component_type
|
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
|
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
|
||||||
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
|
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
|
||||||
from sglang.multimodal_gen.runtime.platforms import (
|
from sglang.multimodal_gen.runtime.platforms import (
|
||||||
@@ -96,20 +95,43 @@ def _warn_if_expected_param_dtype_missing(
|
|||||||
|
|
||||||
|
|
||||||
def _server_args_for_transformer_component(
|
def _server_args_for_transformer_component(
|
||||||
server_args: ServerArgs, component_name: str
|
server_args: ServerArgs,
|
||||||
|
component_name: str,
|
||||||
|
structural_component_name: str | None = None,
|
||||||
) -> ServerArgs:
|
) -> ServerArgs:
|
||||||
"""Mask global quantized override flags for secondary transformer components."""
|
"""Mask global quantized override flags for secondary transformer components."""
|
||||||
|
structural_component_name = structural_component_name or component_name
|
||||||
|
_, separator, suffix = structural_component_name.rpartition("_")
|
||||||
|
is_secondary = structural_component_name == "unconditional_transformer" or (
|
||||||
|
bool(separator) and suffix.isdigit() and int(suffix) >= 2
|
||||||
|
)
|
||||||
component_weights_path = server_args.component_weights_paths.get(component_name)
|
component_weights_path = server_args.component_weights_paths.get(component_name)
|
||||||
component_quantization = server_args.component_quantizations.get(component_name)
|
component_quantization = server_args.component_quantizations.get(component_name)
|
||||||
component_ignored_layers = server_args.component_quantization_ignored_layers.get(
|
component_ignored_layers = server_args.component_quantization_ignored_layers.get(
|
||||||
component_name
|
component_name
|
||||||
)
|
)
|
||||||
if (
|
has_exact_override = (
|
||||||
component_weights_path is not None
|
component_weights_path is not None
|
||||||
or component_quantization is not None
|
or component_quantization is not None
|
||||||
or component_ignored_layers is not None
|
or component_ignored_layers is not None
|
||||||
):
|
)
|
||||||
|
has_global_weights = (
|
||||||
|
server_args.transformer_weights_path is not None
|
||||||
|
or server_args.nunchaku_config is not None
|
||||||
|
)
|
||||||
|
if not has_exact_override and not (is_secondary and has_global_weights):
|
||||||
|
return server_args
|
||||||
|
|
||||||
component_server_args = copy.copy(server_args)
|
component_server_args = copy.copy(server_args)
|
||||||
|
if is_secondary:
|
||||||
|
component_server_args.transformer_weights_path = None
|
||||||
|
component_server_args.nunchaku_config = None
|
||||||
|
if has_global_weights:
|
||||||
|
logger.info(
|
||||||
|
"Ignoring global transformer weight overrides for %s; keep them "
|
||||||
|
"on the primary component unless an exact override is provided.",
|
||||||
|
component_name,
|
||||||
|
)
|
||||||
if component_weights_path is not None:
|
if component_weights_path is not None:
|
||||||
component_server_args.transformer_weights_path = component_weights_path
|
component_server_args.transformer_weights_path = component_weights_path
|
||||||
component_server_args.nunchaku_config = None
|
component_server_args.nunchaku_config = None
|
||||||
@@ -129,25 +151,6 @@ def _server_args_for_transformer_component(
|
|||||||
component_server_args.quantization_ignored_layers = component_ignored_layers
|
component_server_args.quantization_ignored_layers = component_ignored_layers
|
||||||
return component_server_args
|
return component_server_args
|
||||||
|
|
||||||
if component_name not in ("transformer_2", "unconditional_transformer"):
|
|
||||||
return server_args
|
|
||||||
|
|
||||||
if (
|
|
||||||
server_args.transformer_weights_path is None
|
|
||||||
and server_args.nunchaku_config is None
|
|
||||||
):
|
|
||||||
return server_args
|
|
||||||
|
|
||||||
component_server_args = copy.copy(server_args)
|
|
||||||
component_server_args.transformer_weights_path = None
|
|
||||||
component_server_args.nunchaku_config = None
|
|
||||||
logger.info(
|
|
||||||
"Ignoring global transformer_weights_path for %s; keep it on the base "
|
|
||||||
"checkpoint unless a per-component override path is provided.",
|
|
||||||
component_name,
|
|
||||||
)
|
|
||||||
return component_server_args
|
|
||||||
|
|
||||||
|
|
||||||
class TransformerLoader(ComponentLoader):
|
class TransformerLoader(ComponentLoader):
|
||||||
"""Shared loader for (video/audio) DiT transformers."""
|
"""Shared loader for (video/audio) DiT transformers."""
|
||||||
@@ -183,7 +186,9 @@ class TransformerLoader(ComponentLoader):
|
|||||||
self, server_args: ServerArgs, component_name: str
|
self, server_args: ServerArgs, component_name: str
|
||||||
) -> bool:
|
) -> bool:
|
||||||
component_server_args = _server_args_for_transformer_component(
|
component_server_args = _server_args_for_transformer_component(
|
||||||
server_args, component_name
|
server_args,
|
||||||
|
component_name,
|
||||||
|
self.structural_component_name(component_name),
|
||||||
)
|
)
|
||||||
# Don't let a quantized load quietly fall back to the unquantized native
|
# Don't let a quantized load quietly fall back to the unquantized native
|
||||||
# model. That would drop the requested precision and bury the real error.
|
# model. That would drop the requested precision and bury the real error.
|
||||||
@@ -238,7 +243,9 @@ class TransformerLoader(ComponentLoader):
|
|||||||
):
|
):
|
||||||
"""Load the transformer based on the model path, and inference args."""
|
"""Load the transformer based on the model path, and inference args."""
|
||||||
component_server_args = _server_args_for_transformer_component(
|
component_server_args = _server_args_for_transformer_component(
|
||||||
server_args, component_name
|
server_args,
|
||||||
|
component_name,
|
||||||
|
self.structural_component_name(component_name),
|
||||||
)
|
)
|
||||||
|
|
||||||
# 1. hf config
|
# 1. hf config
|
||||||
@@ -261,7 +268,7 @@ class TransformerLoader(ComponentLoader):
|
|||||||
|
|
||||||
# 2. dit config
|
# 2. dit config
|
||||||
# Config from Diffusers supersedes sgl_diffusion's model config
|
# Config from Diffusers supersedes sgl_diffusion's model config
|
||||||
component_type = _normalize_component_type(component_name)
|
component_type = self.structural_component_type(component_name)
|
||||||
server_args.model_paths[component_name] = component_model_path
|
server_args.model_paths[component_name] = component_model_path
|
||||||
if component_type in (
|
if component_type in (
|
||||||
"transformer",
|
"transformer",
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
_list_safetensors_files,
|
_list_safetensors_files,
|
||||||
|
_normalize_component_type,
|
||||||
checkpoint_bytes,
|
checkpoint_bytes,
|
||||||
keep_checkpoint_mapped,
|
keep_checkpoint_mapped,
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
@@ -96,9 +97,9 @@ def _require_native_loader_for_quantized_vae(
|
|||||||
|
|
||||||
|
|
||||||
def _backfill_ltx2_audio_vae_latent_stats(
|
def _backfill_ltx2_audio_vae_latent_stats(
|
||||||
loaded: dict[str, torch.Tensor], component_name: str
|
loaded: dict[str, torch.Tensor], component_type: str
|
||||||
) -> None:
|
) -> None:
|
||||||
if component_name != "audio_vae":
|
if component_type != "audio_vae":
|
||||||
return
|
return
|
||||||
mean_key = "per_channel_statistics.mean-of-means"
|
mean_key = "per_channel_statistics.mean-of-means"
|
||||||
std_key = "per_channel_statistics.std-of-means"
|
std_key = "per_channel_statistics.std-of-means"
|
||||||
@@ -128,9 +129,9 @@ def _convert_conv3d_weights_to_channels_last_3d(module: nn.Module) -> int:
|
|||||||
|
|
||||||
|
|
||||||
def _should_use_channels_last_3d(
|
def _should_use_channels_last_3d(
|
||||||
server_args: ServerArgs | None, component_name: str
|
server_args: ServerArgs | None, component_type: str
|
||||||
) -> bool:
|
) -> bool:
|
||||||
if component_name not in (
|
if component_type not in (
|
||||||
"vae",
|
"vae",
|
||||||
"video_vae",
|
"video_vae",
|
||||||
) or not (current_platform.is_cuda() or current_platform.is_rocm()):
|
) or not (current_platform.is_cuda() or current_platform.is_rocm()):
|
||||||
@@ -232,7 +233,12 @@ def _rehome_cast_weights_to_file(
|
|||||||
|
|
||||||
|
|
||||||
def _hold_decoder_weights_in_decode_dtype(
|
def _hold_decoder_weights_in_decode_dtype(
|
||||||
vae, server_args: ServerArgs, component_name: str, component_model_path: str = ""
|
vae,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
component_name: str,
|
||||||
|
component_model_path: str = "",
|
||||||
|
*,
|
||||||
|
component_type: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Round decoder weights to their decode compute dtype at load.
|
"""Round decoder weights to their decode compute dtype at load.
|
||||||
|
|
||||||
@@ -244,7 +250,8 @@ def _hold_decoder_weights_in_decode_dtype(
|
|||||||
restreaming a third of it per tile and holding all 36 blocks on a 12 GiB
|
restreaming a third of it per tile and holding all 36 blocks on a 12 GiB
|
||||||
card for the decode.
|
card for the decode.
|
||||||
"""
|
"""
|
||||||
if component_name not in ("vae", "video_vae"):
|
component_type = component_type or _normalize_component_type(component_name)
|
||||||
|
if component_type not in ("vae", "video_vae"):
|
||||||
return
|
return
|
||||||
if envs.SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST:
|
if envs.SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST:
|
||||||
return
|
return
|
||||||
@@ -398,6 +405,21 @@ class VAELoader(ComponentLoader):
|
|||||||
) -> bool:
|
) -> bool:
|
||||||
return component_name in ("vae", "video_vae")
|
return component_name in ("vae", "video_vae")
|
||||||
|
|
||||||
|
def select_weight_files(
|
||||||
|
self,
|
||||||
|
safetensors_list: list[str],
|
||||||
|
component_model_path: str,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
component_name: str,
|
||||||
|
vae_precision: str,
|
||||||
|
) -> list[str]:
|
||||||
|
return server_args.pipeline_config.select_vae_weight_files(
|
||||||
|
safetensors_list=safetensors_list,
|
||||||
|
component_model_path=component_model_path,
|
||||||
|
component_name=self.structural_component_type(component_name),
|
||||||
|
vae_precision=vae_precision,
|
||||||
|
)
|
||||||
|
|
||||||
def component_load_precision(
|
def component_load_precision(
|
||||||
self, server_args: ServerArgs, component_name: str
|
self, server_args: ServerArgs, component_name: str
|
||||||
) -> str | None:
|
) -> str | None:
|
||||||
@@ -459,9 +481,7 @@ class VAELoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
config = get_diffusers_component_config(component_path=component_model_path)
|
config = get_diffusers_component_config(component_path=component_model_path)
|
||||||
server_args.model_paths[component_name] = component_model_path
|
server_args.model_paths[component_name] = component_model_path
|
||||||
native_only = component_name in getattr(
|
native_only = self.is_native_only_component(server_args, component_name)
|
||||||
server_args.pipeline_config, "native_only_components", ()
|
|
||||||
)
|
|
||||||
_require_native_loader_for_quantized_vae(
|
_require_native_loader_for_quantized_vae(
|
||||||
config,
|
config,
|
||||||
component_name,
|
component_name,
|
||||||
@@ -474,10 +494,11 @@ class VAELoader(ComponentLoader):
|
|||||||
class_name is not None
|
class_name is not None
|
||||||
), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||||
|
|
||||||
if component_name in ("vae", "video_vae"):
|
component_type = self.structural_component_type(component_name)
|
||||||
|
if component_type in ("vae", "video_vae"):
|
||||||
pipeline_vae_config_attr = "vae_config"
|
pipeline_vae_config_attr = "vae_config"
|
||||||
pipeline_vae_precision = "vae_precision"
|
pipeline_vae_precision = "vae_precision"
|
||||||
elif component_name in ("audio_vae",):
|
elif component_type == "audio_vae":
|
||||||
pipeline_vae_config_attr = "audio_vae_config"
|
pipeline_vae_config_attr = "audio_vae_config"
|
||||||
pipeline_vae_precision = "audio_vae_precision"
|
pipeline_vae_precision = "audio_vae_precision"
|
||||||
else:
|
else:
|
||||||
@@ -529,14 +550,18 @@ class VAELoader(ComponentLoader):
|
|||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
)
|
)
|
||||||
vae = vae.to(device=target_device, dtype=vae_dtype)
|
vae = vae.to(device=target_device, dtype=vae_dtype)
|
||||||
if _should_use_channels_last_3d(server_args, component_name):
|
if _should_use_channels_last_3d(server_args, component_type):
|
||||||
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
||||||
if n > 0:
|
if n > 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"VAE: converted %d Conv3d weights to channels_last_3d", n
|
"VAE: converted %d Conv3d weights to channels_last_3d", n
|
||||||
)
|
)
|
||||||
_hold_decoder_weights_in_decode_dtype(
|
_hold_decoder_weights_in_decode_dtype(
|
||||||
vae, server_args, component_name, component_model_path
|
vae,
|
||||||
|
server_args,
|
||||||
|
component_name,
|
||||||
|
component_model_path,
|
||||||
|
component_type=component_type,
|
||||||
)
|
)
|
||||||
vae = current_platform.optimize_vae(vae)
|
vae = current_platform.optimize_vae(vae)
|
||||||
return vae
|
return vae
|
||||||
@@ -567,11 +592,12 @@ class VAELoader(ComponentLoader):
|
|||||||
safetensors_list = [component_weights_path]
|
safetensors_list = [component_weights_path]
|
||||||
else:
|
else:
|
||||||
safetensors_list = _list_safetensors_files(component_weights_path)
|
safetensors_list = _list_safetensors_files(component_weights_path)
|
||||||
safetensors_list = server_args.pipeline_config.select_vae_weight_files(
|
safetensors_list = self.select_weight_files(
|
||||||
safetensors_list=safetensors_list,
|
safetensors_list,
|
||||||
component_model_path=component_weights_path,
|
component_weights_path,
|
||||||
component_name=component_name,
|
server_args,
|
||||||
vae_precision=vae_precision,
|
component_name,
|
||||||
|
vae_precision,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
@@ -596,7 +622,7 @@ class VAELoader(ComponentLoader):
|
|||||||
loaded = {}
|
loaded = {}
|
||||||
for sf_path in safetensors_list:
|
for sf_path in safetensors_list:
|
||||||
loaded.update(safetensors_load_file(sf_path))
|
loaded.update(safetensors_load_file(sf_path))
|
||||||
_backfill_ltx2_audio_vae_latent_stats(loaded, component_name)
|
_backfill_ltx2_audio_vae_latent_stats(loaded, component_type)
|
||||||
strict_load = native_only
|
strict_load = native_only
|
||||||
# `loaded` holds views into the safetensors mapping. When the component
|
# `loaded` holds views into the safetensors mapping. When the component
|
||||||
# starts on the CPU and the host cannot afford copies of the whole
|
# starts on the CPU and the host cannot afford copies of the whole
|
||||||
@@ -639,13 +665,17 @@ class VAELoader(ComponentLoader):
|
|||||||
if unexpected_keys:
|
if unexpected_keys:
|
||||||
logger.warning("VAE unexpected keys: %s", unexpected_keys)
|
logger.warning("VAE unexpected keys: %s", unexpected_keys)
|
||||||
|
|
||||||
if _should_use_channels_last_3d(server_args, component_name):
|
if _should_use_channels_last_3d(server_args, component_type):
|
||||||
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
||||||
if n > 0:
|
if n > 0:
|
||||||
logger.info("VAE: converted %d Conv3d weights to channels_last_3d", n)
|
logger.info("VAE: converted %d Conv3d weights to channels_last_3d", n)
|
||||||
|
|
||||||
_hold_decoder_weights_in_decode_dtype(
|
_hold_decoder_weights_in_decode_dtype(
|
||||||
vae, server_args, component_name, component_weights_path
|
vae,
|
||||||
|
server_args,
|
||||||
|
component_name,
|
||||||
|
component_weights_path,
|
||||||
|
component_type=component_type,
|
||||||
)
|
)
|
||||||
vae = current_platform.optimize_vae(vae)
|
vae = current_platform.optimize_vae(vae)
|
||||||
return vae
|
return vae
|
||||||
|
|||||||
@@ -22,9 +22,9 @@ class VisionLanguageEncoderLoader(ComponentLoader):
|
|||||||
self,
|
self,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
transformers_or_diffusers: str = "vision_language_encoder",
|
component_name: str = "vision_language_encoder",
|
||||||
) -> Any:
|
) -> Any:
|
||||||
if transformers_or_diffusers == "vision_language_encoder":
|
if self.structural_component_type(component_name) == "vision_language_encoder":
|
||||||
|
|
||||||
if server_args.srt_encoder_url is not None:
|
if server_args.srt_encoder_url is not None:
|
||||||
health_url = server_args.srt_encoder_url.rstrip("/") + "/health"
|
health_url = server_args.srt_encoder_url.rstrip("/") + "/health"
|
||||||
@@ -59,7 +59,7 @@ class VisionLanguageEncoderLoader(ComponentLoader):
|
|||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
)
|
)
|
||||||
target_device = self.target_device(
|
target_device = self.target_device(
|
||||||
server_args.should_start_component_on_cpu("vision_language_encoder")
|
server_args.should_start_component_on_cpu(component_name)
|
||||||
)
|
)
|
||||||
model = GlmImageForConditionalGeneration.from_pretrained(
|
model = GlmImageForConditionalGeneration.from_pretrained(
|
||||||
component_model_path,
|
component_model_path,
|
||||||
@@ -70,5 +70,6 @@ class VisionLanguageEncoderLoader(ComponentLoader):
|
|||||||
return model
|
return model
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported library for VisionLanguageEncoder: {transformers_or_diffusers}"
|
f"Unsupported component type for VisionLanguageEncoder: "
|
||||||
|
f"{self.structural_component_type(component_name)}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import (
|
|||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
PipelineComponentLoader,
|
PipelineComponentLoader,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.utils import _normalize_component_type
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_loading_order import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_loading_order import (
|
||||||
ComponentLoadSpec,
|
ComponentLoadSpec,
|
||||||
order_component_load_specs,
|
order_component_load_specs,
|
||||||
@@ -73,6 +74,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
is_video_pipeline: bool = False # To be overridden by video pipelines
|
is_video_pipeline: bool = False # To be overridden by video pipelines
|
||||||
# should contains only the modules to be loaded
|
# should contains only the modules to be loaded
|
||||||
_required_config_modules: list[str] = []
|
_required_config_modules: list[str] = []
|
||||||
|
_unfiltered_required_config_modules: tuple[str, ...] = ()
|
||||||
_extra_config_module_map: dict[str, str] = {}
|
_extra_config_module_map: dict[str, str] = {}
|
||||||
server_args: ServerArgs | None = None
|
server_args: ServerArgs | None = None
|
||||||
modules: dict[str, Any] = {}
|
modules: dict[str, Any] = {}
|
||||||
@@ -118,7 +120,8 @@ class ComposedPipelineBase(ABC):
|
|||||||
)
|
)
|
||||||
if base_required_config_modules is None:
|
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._unfiltered_required_config_modules = tuple(base_required_config_modules)
|
||||||
|
self._required_config_modules = list(self._unfiltered_required_config_modules)
|
||||||
self._extra_config_module_map = dict(self._extra_config_module_map)
|
self._extra_config_module_map = dict(self._extra_config_module_map)
|
||||||
|
|
||||||
# Filter modules based on disaggregation role
|
# Filter modules based on disaggregation role
|
||||||
@@ -131,6 +134,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
extra_allowed_modules=self._get_extra_allowed_modules_for_role(
|
extra_allowed_modules=self._get_extra_allowed_modules_for_role(
|
||||||
self._disagg_role, task_name
|
self._disagg_role, task_name
|
||||||
),
|
),
|
||||||
|
structural_component_names=self._extra_config_module_map,
|
||||||
)
|
)
|
||||||
skipped = set(original_modules) - set(self._required_config_modules)
|
skipped = set(original_modules) - set(self._required_config_modules)
|
||||||
if skipped:
|
if skipped:
|
||||||
@@ -308,11 +312,36 @@ class ComposedPipelineBase(ABC):
|
|||||||
get_diffusers_component_config,
|
get_diffusers_component_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
required = set(self.required_config_modules)
|
loaded_components = set(self.required_config_modules)
|
||||||
for module_name in full_model_index:
|
loaded_structural_components = {
|
||||||
if module_name in required:
|
self._extra_config_module_map.get(name, name)
|
||||||
continue # will be loaded normally
|
for name in self.required_config_modules
|
||||||
cfg_attr = self._CONFIG_ATTR_MAP.get(module_name)
|
}
|
||||||
|
skipped_components: list[tuple[str, str]] = []
|
||||||
|
seen_component_keys = loaded_components | loaded_structural_components
|
||||||
|
for component_name in self._unfiltered_required_config_modules:
|
||||||
|
structural_name = self._extra_config_module_map.get(
|
||||||
|
component_name, component_name
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
component_name in seen_component_keys
|
||||||
|
or structural_name in seen_component_keys
|
||||||
|
or (
|
||||||
|
component_name not in full_model_index
|
||||||
|
and structural_name not in full_model_index
|
||||||
|
)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
skipped_components.append((component_name, structural_name))
|
||||||
|
seen_component_keys.update((component_name, structural_name))
|
||||||
|
for structural_name in full_model_index:
|
||||||
|
if structural_name not in seen_component_keys:
|
||||||
|
skipped_components.append((structural_name, structural_name))
|
||||||
|
|
||||||
|
for component_name, structural_name in skipped_components:
|
||||||
|
cfg_attr = self._CONFIG_ATTR_MAP.get(
|
||||||
|
_normalize_component_type(structural_name)
|
||||||
|
)
|
||||||
if cfg_attr is None:
|
if cfg_attr is None:
|
||||||
continue # not a config we need to patch
|
continue # not a config we need to patch
|
||||||
|
|
||||||
@@ -322,7 +351,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
component_path = self._resolve_component_path(
|
component_path = self._resolve_component_path(
|
||||||
server_args, module_name, module_name
|
server_args, component_name, structural_name
|
||||||
)
|
)
|
||||||
hf_config = get_diffusers_component_config(
|
hf_config = get_diffusers_component_config(
|
||||||
component_path=component_path
|
component_path=component_path
|
||||||
@@ -336,7 +365,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
"Disagg role=%s: initialized %s config from HF JSON "
|
"Disagg role=%s: initialized %s config from HF JSON "
|
||||||
"(spatial_compression_ratio=%s)",
|
"(spatial_compression_ratio=%s)",
|
||||||
self._disagg_role.value,
|
self._disagg_role.value,
|
||||||
module_name,
|
component_name,
|
||||||
getattr(
|
getattr(
|
||||||
getattr(pipeline_cfg, "arch_config", None),
|
getattr(pipeline_cfg, "arch_config", None),
|
||||||
"spatial_compression_ratio",
|
"spatial_compression_ratio",
|
||||||
@@ -348,7 +377,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
"Disagg role=%s: failed to read HF config for skipped "
|
"Disagg role=%s: failed to read HF config for skipped "
|
||||||
"component %s: %s",
|
"component %s: %s",
|
||||||
self._disagg_role.value,
|
self._disagg_role.value,
|
||||||
module_name,
|
component_name,
|
||||||
e,
|
e,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -460,9 +489,11 @@ class ComposedPipelineBase(ABC):
|
|||||||
if self._disagg_role != RoleType.MONOLITHIC:
|
if self._disagg_role != RoleType.MONOLITHIC:
|
||||||
self._init_skipped_component_configs(model_index, server_args)
|
self._init_skipped_component_configs(model_index, server_args)
|
||||||
|
|
||||||
|
declared_modules = model_index
|
||||||
model_index = {
|
model_index = {
|
||||||
required_module: model_index[required_module]
|
required_module: model_index[required_module]
|
||||||
for required_module in self.required_config_modules
|
for required_module in self.required_config_modules
|
||||||
|
if required_module in model_index
|
||||||
}
|
}
|
||||||
|
|
||||||
for module_name in self.required_config_modules:
|
for module_name in self.required_config_modules:
|
||||||
@@ -477,15 +508,17 @@ class ComposedPipelineBase(ABC):
|
|||||||
module_name,
|
module_name,
|
||||||
extra_module_value,
|
extra_module_value,
|
||||||
)
|
)
|
||||||
if extra_module_value in model_index:
|
if extra_module_value in declared_modules:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Using module %s for %s", extra_module_value, module_name
|
"Using module %s for %s", extra_module_value, module_name
|
||||||
)
|
)
|
||||||
model_index[module_name] = model_index[extra_module_value]
|
model_index[module_name] = declared_modules[extra_module_value]
|
||||||
continue
|
continue
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Required module key: {module_name} value: {model_index.get(module_name)} was not found in loaded modules {model_index.keys()}"
|
f"Required module key: {module_name} value: "
|
||||||
|
f"{declared_modules.get(module_name)} was not found in "
|
||||||
|
f"declared modules {declared_modules.keys()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# all the component models used by the pipeline
|
# all the component models used by the pipeline
|
||||||
@@ -581,7 +614,8 @@ class ComposedPipelineBase(ABC):
|
|||||||
matched_backend_key,
|
matched_backend_key,
|
||||||
)
|
)
|
||||||
module, memory_usage = PipelineComponentLoader.load_component(
|
module, memory_usage = PipelineComponentLoader.load_component(
|
||||||
component_name=load_module_name,
|
component_name=module_name,
|
||||||
|
component_type=load_module_name,
|
||||||
component_model_path=component_model_path,
|
component_model_path=component_model_path,
|
||||||
transformers_or_diffusers=transformers_or_diffusers,
|
transformers_or_diffusers=transformers_or_diffusers,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -590,7 +624,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
component_attn_name=matched_backend_key or module_name,
|
component_attn_name=matched_backend_key or module_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.memory_usages[load_module_name] = memory_usage
|
self.memory_usages[module_name] = memory_usage
|
||||||
|
|
||||||
if module_name in loaded_components:
|
if module_name in loaded_components:
|
||||||
logger.warning("Overwriting module %s", module_name)
|
logger.warning("Overwriting module %s", module_name)
|
||||||
|
|||||||
@@ -0,0 +1,263 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.disaggregation.roles import (
|
||||||
|
RoleType,
|
||||||
|
filter_modules_for_role,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.adapter_loader import (
|
||||||
|
AdapterLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
|
ComponentLoader,
|
||||||
|
PipelineComponentLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
|
||||||
|
_server_args_for_transformer_component,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import VAELoader
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.vl_encoder_loader import (
|
||||||
|
VisionLanguageEncoderLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
|
ComposedPipelineBase,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _AliasPipeline(ComposedPipelineBase):
|
||||||
|
def initialize_pipeline(self, _server_args):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def create_pipeline_stages(self, _server_args):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class TestComponentLoaderIdentity(unittest.TestCase):
|
||||||
|
def test_structural_identity_selects_roles_and_preserves_exact_policy_keys(self):
|
||||||
|
aliases = {
|
||||||
|
"conditioning": "text_encoder_2",
|
||||||
|
"refiner": "transformer_3",
|
||||||
|
"decoder": "video_vae",
|
||||||
|
}
|
||||||
|
self.assertEqual(
|
||||||
|
filter_modules_for_role(
|
||||||
|
aliases, RoleType.ENCODER, structural_component_names=aliases
|
||||||
|
),
|
||||||
|
["conditioning"],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
filter_modules_for_role(
|
||||||
|
aliases, RoleType.DENOISER, structural_component_names=aliases
|
||||||
|
),
|
||||||
|
["refiner"],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
filter_modules_for_role(
|
||||||
|
aliases, RoleType.DECODER, structural_component_names=aliases
|
||||||
|
),
|
||||||
|
["decoder"],
|
||||||
|
)
|
||||||
|
|
||||||
|
loader = ComponentLoader.for_component_type(
|
||||||
|
"duration_head_2", "ltx2", "LTX2DurationHeadModel"
|
||||||
|
)
|
||||||
|
self.assertIsInstance(loader, AdapterLoader)
|
||||||
|
self.assertEqual(
|
||||||
|
loader.structural_component_name("auxiliary_head"), "duration_head_2"
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
loader.structural_component_type("auxiliary_head"), "duration_head"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_declared_alias_loads_by_exact_key_and_structural_source(self):
|
||||||
|
pipeline = object.__new__(_AliasPipeline)
|
||||||
|
pipeline.model_path = "/model"
|
||||||
|
pipeline.memory_usages = {}
|
||||||
|
pipeline._disagg_role = RoleType.MONOLITHIC
|
||||||
|
pipeline._required_config_modules = ["auxiliary_head"]
|
||||||
|
pipeline._extra_config_module_map = {"auxiliary_head": "duration_head_2"}
|
||||||
|
pipeline._load_config = lambda: {
|
||||||
|
"_class_name": "TestPipeline",
|
||||||
|
"_diffusers_version": "0",
|
||||||
|
"duration_head_2": ["ltx2", "LTX2DurationHeadModel"],
|
||||||
|
"scheduler": ["diffusers", "Scheduler"],
|
||||||
|
}
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_paths={},
|
||||||
|
component_direct_gpu_weight_loading=set(),
|
||||||
|
resolve_component_attention_backend=lambda *_names: (None, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
PipelineComponentLoader,
|
||||||
|
"load_component",
|
||||||
|
return_value=(nn.Linear(1, 1), 0.25),
|
||||||
|
) as load_component:
|
||||||
|
loaded = pipeline.load_modules(server_args)
|
||||||
|
|
||||||
|
self.assertIn("auxiliary_head", loaded)
|
||||||
|
self.assertEqual(pipeline.memory_usages, {"auxiliary_head": 0.25})
|
||||||
|
load_component.assert_called_once_with(
|
||||||
|
component_name="auxiliary_head",
|
||||||
|
component_type="duration_head_2",
|
||||||
|
component_model_path="/model/duration_head_2",
|
||||||
|
transformers_or_diffusers="ltx2",
|
||||||
|
server_args=server_args,
|
||||||
|
component_architecture="LTX2DurationHeadModel",
|
||||||
|
component_attn_backend=None,
|
||||||
|
component_attn_name="auxiliary_head",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_skipped_alias_keeps_exact_override_and_structural_config(self):
|
||||||
|
pipeline = object.__new__(_AliasPipeline)
|
||||||
|
pipeline.model_path = "/model"
|
||||||
|
pipeline._disagg_role = RoleType.ENCODER
|
||||||
|
pipeline._required_config_modules = []
|
||||||
|
pipeline._unfiltered_required_config_modules = ("vae_2",)
|
||||||
|
pipeline._extra_config_module_map = {"vae_2": "video_vae_2"}
|
||||||
|
vae_config = Mock()
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_paths={"vae_2": "/exact/decoder"},
|
||||||
|
pipeline_config=SimpleNamespace(vae_config=vae_config),
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.pipelines_core."
|
||||||
|
"composed_pipeline_base.prepare_diffusers_component_path_for_loading",
|
||||||
|
return_value="/resolved/decoder",
|
||||||
|
) as prepare_path,
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.utils.hf_diffusers_utils."
|
||||||
|
"get_diffusers_component_config",
|
||||||
|
return_value={"sample_size": 32},
|
||||||
|
) as get_config,
|
||||||
|
):
|
||||||
|
pipeline._init_skipped_component_configs(
|
||||||
|
{
|
||||||
|
"vae_2": ["diffusers", "VideoVAE"],
|
||||||
|
"video_vae_2": ["diffusers", "VideoVAE"],
|
||||||
|
},
|
||||||
|
server_args,
|
||||||
|
)
|
||||||
|
|
||||||
|
prepare_path.assert_called_once_with("/exact/decoder")
|
||||||
|
get_config.assert_called_once_with(component_path="/resolved/decoder")
|
||||||
|
vae_config.update_model_arch.assert_called_once_with({"sample_size": 32})
|
||||||
|
|
||||||
|
def test_structural_aliases_select_loader_config_and_weight_behavior(self):
|
||||||
|
selected = ["selected.safetensors"]
|
||||||
|
select_weight_files = Mock(return_value=selected)
|
||||||
|
vae_loader = VAELoader()
|
||||||
|
vae_loader.component_type = "video_vae"
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
pipeline_config=SimpleNamespace(select_vae_weight_files=select_weight_files)
|
||||||
|
)
|
||||||
|
self.assertIs(
|
||||||
|
vae_loader.select_weight_files(
|
||||||
|
["candidate.safetensors"],
|
||||||
|
"/decoder",
|
||||||
|
server_args,
|
||||||
|
"decoder",
|
||||||
|
"fp32",
|
||||||
|
),
|
||||||
|
selected,
|
||||||
|
)
|
||||||
|
select_weight_files.assert_called_once_with(
|
||||||
|
safetensors_list=["candidate.safetensors"],
|
||||||
|
component_model_path="/decoder",
|
||||||
|
component_name="video_vae",
|
||||||
|
vae_precision="fp32",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_transformer_exact_override_wins_without_leaking_global_flags(self):
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_weights_paths={},
|
||||||
|
component_quantizations={},
|
||||||
|
component_quantization_ignored_layers={},
|
||||||
|
transformer_weights_path="global.safetensors",
|
||||||
|
nunchaku_config="global-nunchaku",
|
||||||
|
)
|
||||||
|
secondary = _server_args_for_transformer_component(
|
||||||
|
server_args, "refiner", "transformer_3"
|
||||||
|
)
|
||||||
|
self.assertIsNone(secondary.transformer_weights_path)
|
||||||
|
self.assertIsNone(secondary.nunchaku_config)
|
||||||
|
|
||||||
|
server_args.component_weights_paths["refiner"] = "refiner.safetensors"
|
||||||
|
exact = _server_args_for_transformer_component(
|
||||||
|
server_args, "refiner", "transformer_3"
|
||||||
|
)
|
||||||
|
self.assertEqual(exact.transformer_weights_path, "refiner.safetensors")
|
||||||
|
self.assertIs(
|
||||||
|
_server_args_for_transformer_component(
|
||||||
|
server_args, "denoiser", "transformer"
|
||||||
|
),
|
||||||
|
server_args,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_native_fallback_prioritizes_exact_component_precision(self):
|
||||||
|
loader = ComponentLoader()
|
||||||
|
loader.component_type = "video_vae_2"
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
pipeline_config=SimpleNamespace(),
|
||||||
|
revision="test-revision",
|
||||||
|
trust_remote_code=False,
|
||||||
|
)
|
||||||
|
native_model = nn.Linear(1, 1)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"component_loader.resolve_component_precision",
|
||||||
|
side_effect=[torch.float32],
|
||||||
|
) as resolve_precision,
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"component_loader.prepare_diffusers_component_path_for_loading",
|
||||||
|
return_value="/resolved/decoder",
|
||||||
|
),
|
||||||
|
patch("diffusers.AutoModel.from_pretrained", return_value=native_model),
|
||||||
|
):
|
||||||
|
loaded = loader.load_native("/decoder", server_args, "diffusers", "decoder")
|
||||||
|
|
||||||
|
self.assertIs(loaded, native_model)
|
||||||
|
self.assertEqual(
|
||||||
|
[call.args[1] for call in resolve_precision.call_args_list], ["decoder"]
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_vision_language_loader_uses_exact_residency_key(self):
|
||||||
|
requested = []
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
srt_encoder_url=None,
|
||||||
|
trust_remote_code=False,
|
||||||
|
revision=None,
|
||||||
|
should_start_component_on_cpu=lambda name: requested.append(name) or True,
|
||||||
|
)
|
||||||
|
loader = VisionLanguageEncoderLoader()
|
||||||
|
loader.component_type = "vision_language_encoder"
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"vl_encoder_loader.get_hf_config",
|
||||||
|
return_value=object(),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"transformers.GlmImageForConditionalGeneration.from_pretrained",
|
||||||
|
return_value=nn.Linear(1, 1),
|
||||||
|
),
|
||||||
|
):
|
||||||
|
loader.load_customized("unused", server_args, "prompt_conditioner")
|
||||||
|
|
||||||
|
self.assertEqual(requested, ["prompt_conditioner"])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user