[diffusion] chore: make malformed component execution options fail-fast (#37049)
This commit is contained in:
@@ -309,7 +309,7 @@ Selectors match exact loaded component keys from `model_index.json`, including n
|
|||||||
|
|
||||||
The existing `--dit-cpu-offload`, `--text-encoder-cpu-offload`, `--image-encoder-cpu-offload`, `--vae-cpu-offload`, and `--cpu-offload-components` options remain supported. New and legacy options may be mixed: `--component-residency` wins only for components it matches, while unmatched legacy settings remain effective. Legacy layerwise selectors take precedence over legacy component-offload selectors for the same component. Explicit `--dit-layerwise-offload false` makes the DiT resident unless another explicit DiT selector, such as `--dit-cpu-offload true` or `--component-residency dit=component-offload`, selects a different mode.
|
The existing `--dit-cpu-offload`, `--text-encoder-cpu-offload`, `--image-encoder-cpu-offload`, `--vae-cpu-offload`, and `--cpu-offload-components` options remain supported. New and legacy options may be mixed: `--component-residency` wins only for components it matches, while unmatched legacy settings remain effective. Legacy layerwise selectors take precedence over legacy component-offload selectors for the same component. Explicit `--dit-layerwise-offload false` makes the DiT resident unless another explicit DiT selector, such as `--dit-cpu-offload true` or `--component-residency dit=component-offload`, selects a different mode.
|
||||||
|
|
||||||
Layerwise selection is strict. A native weighted component selected for `layerwise-offload` must declare its layer structure; otherwise startup fails with the unsupported component name instead of silently changing modes. FSDP applies only to resident components. The Diffusers backend supports only pipeline-wide `all=resident` and `all=component-offload`.
|
Layerwise selection is strict. A native weighted component selected for `layerwise-offload` must declare its layer structure; otherwise startup fails with the unsupported component name instead of silently changing modes. Explicit non-resident placement also requires a request-time component-use declaration, so it cannot silently select a module that the pipeline does not manage. FSDP applies only to resident components. The Diffusers backend supports only pipeline-wide `all=resident` and `all=component-offload`.
|
||||||
|
|
||||||
### Layerwise Offload Tuning
|
### Layerwise Offload Tuning
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ This document describes the attention backends available in sglang diffusion (`s
|
|||||||
|
|
||||||
Attention backends are defined by `AttentionBackendEnum` (`sglang.multimodal_gen.runtime.platforms.interface.AttentionBackendEnum`) and selected via the CLI flag `--attention-backend`.
|
Attention backends are defined by `AttentionBackendEnum` (`sglang.multimodal_gen.runtime.platforms.interface.AttentionBackendEnum`) and selected via the CLI flag `--attention-backend`.
|
||||||
|
|
||||||
Backend selection is performed by the shared attention layers (e.g. `LocalAttention` / `USPAttention` / `UlyssesAttention` in `sglang.multimodal_gen.runtime.layers.attention.layer`). `--attention-backend` is strict for the diffusion transformer / DiT. Auxiliary components such as encoders and VAEs use it when compatible, then fall back to a component default or a platform-compatible backend. Use `--component-attention-backends` when an auxiliary component must use a specific backend; incompatible component overrides fail unless a sparse backend is being replaced for cross-attention.
|
Backend selection is performed by the shared attention layers (e.g. `LocalAttention` / `USPAttention` / `UlyssesAttention` in `sglang.multimodal_gen.runtime.layers.attention.layer`). `--attention-backend` is strict for the diffusion transformer / DiT. Auxiliary components such as encoders and VAEs use it when compatible, then fall back to a component default or a platform-compatible backend. Use `--component-attention-backends` when an auxiliary component must use a specific backend; incompatible overrides, including an override for a component that constructs no SGLang attention layer, fail unless a sparse backend is being replaced for cross-attention.
|
||||||
|
|
||||||
When using the diffusers backend, `--attention-backend` is passed through to diffusers'
|
When using the diffusers backend, `--attention-backend` is passed through to diffusers'
|
||||||
`set_attention_backend` (e.g., `flash`, `_flash_3_hub`, `sage`, `xformers`, `native`).
|
`set_attention_backend` (e.g., `flash`, `_flash_3_hub`, `sage`, `xformers`, `native`).
|
||||||
|
|||||||
@@ -200,6 +200,9 @@ class PipelineConfig:
|
|||||||
|
|
||||||
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
|
||||||
|
# implementation because their pipeline requires SGLang-specific behavior.
|
||||||
|
native_only_components: tuple[str, ...] = ()
|
||||||
|
|
||||||
model_path: str = ""
|
model_path: str = ""
|
||||||
pipeline_config_path: str | None = None
|
pipeline_config_path: str | None = None
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ class Hunyuan3D2PipelineConfig(PipelineConfig):
|
|||||||
default_factory=lambda: (CLIPTextConfig(),)
|
default_factory=lambda: (CLIPTextConfig(),)
|
||||||
)
|
)
|
||||||
text_encoder_precisions: tuple[str, ...] = ("fp16",)
|
text_encoder_precisions: tuple[str, ...] = ("fp16",)
|
||||||
native_only_components = ("delight_text_encoder",)
|
native_only_components: tuple[str, ...] = ("delight_text_encoder",)
|
||||||
|
|
||||||
# Shape model configuration
|
# Shape model configuration
|
||||||
shape_model_path: Optional[str] = None
|
shape_model_path: Optional[str] = None
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ class LTX25PipelineConfig(LTX2PipelineConfig):
|
|||||||
# One checkpoint drives both T2V and image-conditioned generation, so this
|
# One checkpoint drives both T2V and image-conditioned generation, so this
|
||||||
# must stay TI2V -- T2V rejects `--image-path` outright.
|
# must stay TI2V -- T2V rejects `--image-path` outright.
|
||||||
task_type: ModelTaskType = ModelTaskType.TI2V
|
task_type: ModelTaskType = ModelTaskType.TI2V
|
||||||
native_only_components = ("diffusion_decoder",)
|
native_only_components: tuple[str, ...] = ("diffusion_decoder",)
|
||||||
|
|
||||||
dit_config: LTX25Config = field(default_factory=LTX25Config)
|
dit_config: LTX25Config = field(default_factory=LTX25Config)
|
||||||
vae_config: LTX25VideoVAEConfig = field(default_factory=LTX25VideoVAEConfig)
|
vae_config: LTX25VideoVAEConfig = field(default_factory=LTX25VideoVAEConfig)
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ class MiniMaxH3PipelineConfig(PipelineConfig):
|
|||||||
# generic TI2V image resize would both duplicate that work and overwrite
|
# generic TI2V image resize would both duplicate that work and overwrite
|
||||||
# the already-resolved target canvas.
|
# the already-resolved target canvas.
|
||||||
skip_input_image_preprocess: bool = True
|
skip_input_image_preprocess: bool = True
|
||||||
native_only_components = (
|
native_only_components: tuple[str, ...] = (
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
"transformer",
|
"transformer",
|
||||||
"video_vae",
|
"video_vae",
|
||||||
|
|||||||
@@ -130,6 +130,13 @@ def _record_component_attn_backend(backend_name: str, reason: str | None) -> boo
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def record_component_attn_backend(
|
||||||
|
backend: AttentionBackendEnum, reason: str | None = None
|
||||||
|
) -> bool:
|
||||||
|
"""Record a component backend selected outside layer construction."""
|
||||||
|
return _record_component_attn_backend(backend.name.lower(), reason)
|
||||||
|
|
||||||
|
|
||||||
def _log_component_attn_backend_summary(
|
def _log_component_attn_backend_summary(
|
||||||
context: ComponentAttnBackendContext | None,
|
context: ComponentAttnBackendContext | None,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -365,6 +372,7 @@ def component_attn_backend_context_manager(
|
|||||||
attn_backend: AttentionBackendEnum | None,
|
attn_backend: AttentionBackendEnum | None,
|
||||||
component_name: str | None = None,
|
component_name: str | None = None,
|
||||||
allow_global_backend_fallback: bool = False,
|
allow_global_backend_fallback: bool = False,
|
||||||
|
require_component_backend_selection: bool = True,
|
||||||
) -> Generator[None, None, None]:
|
) -> Generator[None, None, None]:
|
||||||
if attn_backend is None and component_name is None:
|
if attn_backend is None and component_name is None:
|
||||||
yield
|
yield
|
||||||
@@ -378,12 +386,35 @@ def component_attn_backend_context_manager(
|
|||||||
allow_global_backend_fallback,
|
allow_global_backend_fallback,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
unused_component_name: str | None = None
|
||||||
|
unused_backend_name: str | None = None
|
||||||
|
completed = False
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
|
completed = True
|
||||||
finally:
|
finally:
|
||||||
context = component_attn_backend_context.get()
|
context = component_attn_backend_context.get()
|
||||||
|
unused_component_override = (
|
||||||
|
completed
|
||||||
|
and require_component_backend_selection
|
||||||
|
and (
|
||||||
|
context is not None
|
||||||
|
and context.backend is not None
|
||||||
|
and context.component_name is not None
|
||||||
|
and not context.selected_backends
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if unused_component_override:
|
||||||
|
unused_component_name = context.component_name
|
||||||
|
unused_backend_name = context.backend.name.lower()
|
||||||
_log_component_attn_backend_summary(context)
|
_log_component_attn_backend_summary(context)
|
||||||
component_attn_backend_context.reset(token)
|
component_attn_backend_context.reset(token)
|
||||||
|
if unused_component_name is not None and unused_backend_name is not None:
|
||||||
|
raise ValueError(
|
||||||
|
f"Attention backend {unused_backend_name!r} was requested for component "
|
||||||
|
f"{unused_component_name!r}, but that component "
|
||||||
|
"did not construct an SGLang attention layer."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
|
|||||||
@@ -180,10 +180,7 @@ class ComponentLoader(ABC):
|
|||||||
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:
|
||||||
native_only_components = getattr(
|
return component_name in server_args.pipeline_config.native_only_components
|
||||||
server_args.pipeline_config, "native_only_components", ()
|
|
||||||
)
|
|
||||||
return component_name in native_only_components
|
|
||||||
|
|
||||||
def validate_native_fallback(
|
def validate_native_fallback(
|
||||||
self, _server_args: ServerArgs, _component_name: str
|
self, _server_args: ServerArgs, _component_name: str
|
||||||
@@ -213,6 +210,12 @@ class ComponentLoader(ABC):
|
|||||||
attn_backend,
|
attn_backend,
|
||||||
component_name=component_attn_name,
|
component_name=component_attn_name,
|
||||||
allow_global_backend_fallback=allow_global_backend_fallback,
|
allow_global_backend_fallback=allow_global_backend_fallback,
|
||||||
|
require_component_backend_selection=(
|
||||||
|
attn_backend is None
|
||||||
|
or not server_args.is_component_attention_backend_automatic(
|
||||||
|
component_attn_name
|
||||||
|
)
|
||||||
|
),
|
||||||
):
|
):
|
||||||
load_kwargs = self.customized_load_kwargs_for_component(
|
load_kwargs = self.customized_load_kwargs_for_component(
|
||||||
server_args, component_name
|
server_args, component_name
|
||||||
@@ -235,6 +238,12 @@ class ComponentLoader(ABC):
|
|||||||
attn_backend,
|
attn_backend,
|
||||||
component_name=component_attn_name,
|
component_name=component_attn_name,
|
||||||
allow_global_backend_fallback=allow_global_backend_fallback,
|
allow_global_backend_fallback=allow_global_backend_fallback,
|
||||||
|
require_component_backend_selection=(
|
||||||
|
attn_backend is None
|
||||||
|
or not server_args.is_component_attention_backend_automatic(
|
||||||
|
component_attn_name
|
||||||
|
)
|
||||||
|
),
|
||||||
):
|
):
|
||||||
component = self.load_native(
|
component = self.load_native(
|
||||||
component_model_path,
|
component_model_path,
|
||||||
@@ -735,6 +744,12 @@ class PipelineComponentLoader:
|
|||||||
allow_global_backend_fallback=(
|
allow_global_backend_fallback=(
|
||||||
loader.allow_global_attention_backend_fallback
|
loader.allow_global_attention_backend_fallback
|
||||||
),
|
),
|
||||||
|
require_component_backend_selection=(
|
||||||
|
component_attn_backend is None
|
||||||
|
or not server_args.is_component_attention_backend_automatic(
|
||||||
|
component_attn_name
|
||||||
|
)
|
||||||
|
),
|
||||||
):
|
):
|
||||||
return loader.load(
|
return loader.load(
|
||||||
component_model_path,
|
component_model_path,
|
||||||
|
|||||||
+10
-1
@@ -196,6 +196,7 @@ class TransformerLoader(ComponentLoader):
|
|||||||
def validate_native_fallback(
|
def validate_native_fallback(
|
||||||
self, server_args: ServerArgs, component_name: str
|
self, server_args: ServerArgs, component_name: str
|
||||||
) -> None:
|
) -> None:
|
||||||
|
super().validate_native_fallback(server_args, component_name)
|
||||||
requested_distributed_execution = []
|
requested_distributed_execution = []
|
||||||
if server_args.tp_size is not None and server_args.tp_size > 1:
|
if server_args.tp_size is not None and server_args.tp_size > 1:
|
||||||
requested_distributed_execution.append(f"tp_size={server_args.tp_size}")
|
requested_distributed_execution.append(f"tp_size={server_args.tp_size}")
|
||||||
@@ -209,6 +210,13 @@ class TransformerLoader(ComponentLoader):
|
|||||||
requested_distributed_execution.append(
|
requested_distributed_execution.append(
|
||||||
f"ring_degree={server_args.ring_degree}"
|
f"ring_degree={server_args.ring_degree}"
|
||||||
)
|
)
|
||||||
|
if (
|
||||||
|
server_args.kv_gather_degree is not None
|
||||||
|
and server_args.kv_gather_degree > 1
|
||||||
|
):
|
||||||
|
requested_distributed_execution.append(
|
||||||
|
f"kv_gather_degree={server_args.kv_gather_degree}"
|
||||||
|
)
|
||||||
if server_args.should_use_fsdp_for_component(component_name):
|
if server_args.should_use_fsdp_for_component(component_name):
|
||||||
requested_distributed_execution.append("FSDP")
|
requested_distributed_execution.append("FSDP")
|
||||||
if requested_distributed_execution:
|
if requested_distributed_execution:
|
||||||
@@ -217,7 +225,8 @@ class TransformerLoader(ComponentLoader):
|
|||||||
f"{component_name!r} cannot honor requested distributed execution: "
|
f"{component_name!r} cannot honor requested distributed execution: "
|
||||||
f"{', '.join(requested_distributed_execution)}. Use an SGLang-native "
|
f"{', '.join(requested_distributed_execution)}. Use an SGLang-native "
|
||||||
"transformer implementation or set tp_size, sp_degree, "
|
"transformer implementation or set tp_size, sp_degree, "
|
||||||
"ulysses_degree, and ring_degree to 1 without FSDP."
|
"ulysses_degree, ring_degree, and kv_gather_degree to 1 without "
|
||||||
|
"FSDP."
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_customized(
|
def load_customized(
|
||||||
|
|||||||
@@ -178,6 +178,35 @@ class ComponentResidencyManager:
|
|||||||
self._ordered_uses = tuple(
|
self._ordered_uses = tuple(
|
||||||
use for uses in self._stage_uses_by_index for use in uses
|
use for uses in self._stage_uses_by_index for use in uses
|
||||||
)
|
)
|
||||||
|
self._validate_explicit_nonresident_components()
|
||||||
|
|
||||||
|
def _validate_explicit_nonresident_components(self) -> None:
|
||||||
|
"""Reject explicit offload selectors with no request-time use site.
|
||||||
|
|
||||||
|
Component placement is enacted by the request timeline, not merely by
|
||||||
|
choosing an initial load device. An explicit non-resident module with
|
||||||
|
no declared ``ComponentUse`` would otherwise be accepted but never
|
||||||
|
moved to the device before a forward pass.
|
||||||
|
"""
|
||||||
|
if not isinstance(self.server_args, ServerArgs):
|
||||||
|
return
|
||||||
|
|
||||||
|
declared_components = {use.component_name for use in self._ordered_uses}
|
||||||
|
unmanaged_components = sorted(
|
||||||
|
component_name
|
||||||
|
for component_name, module in self.pipeline.modules.items()
|
||||||
|
if isinstance(module, nn.Module)
|
||||||
|
and self.server_args.explicit_residency_mode(component_name)
|
||||||
|
in (COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD)
|
||||||
|
and component_name not in declared_components
|
||||||
|
)
|
||||||
|
if unmanaged_components:
|
||||||
|
names = ", ".join(repr(name) for name in unmanaged_components)
|
||||||
|
raise ComponentResidencyError(
|
||||||
|
"Explicit component residency requires "
|
||||||
|
f"{names} to have a request-time ComponentUse declaration; "
|
||||||
|
"none appears in this pipeline"
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _is_warmup_batch(batch: ResidencyBatch | list[ResidencyBatch]) -> bool:
|
def _is_warmup_batch(batch: ResidencyBatch | list[ResidencyBatch]) -> bool:
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
|||||||
get_attn_backend,
|
get_attn_backend,
|
||||||
get_component_forced_attn_backend,
|
get_component_forced_attn_backend,
|
||||||
get_global_forced_attn_backend,
|
get_global_forced_attn_backend,
|
||||||
|
record_component_attn_backend,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
ColumnParallelLinear,
|
ColumnParallelLinear,
|
||||||
@@ -2011,6 +2012,11 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
# Component overrides disappear when the loader context exits. Preserve
|
# Component overrides disappear when the loader context exits. Preserve
|
||||||
# only that selection; process-wide overrides are resolved at first use.
|
# only that selection; process-wide overrides are resolved at first use.
|
||||||
self._component_attention_backend_override = get_component_forced_attn_backend()
|
self._component_attention_backend_override = get_component_forced_attn_backend()
|
||||||
|
if self._component_attention_backend_override is not None:
|
||||||
|
record_component_attn_backend(
|
||||||
|
self._component_attention_backend_override,
|
||||||
|
"deferred model-specific resolution",
|
||||||
|
)
|
||||||
self._resolved_attention_backend: AttentionBackendEnum | None = None
|
self._resolved_attention_backend: AttentionBackendEnum | None = None
|
||||||
self._mark_missing_params_required()
|
self._mark_missing_params_required()
|
||||||
|
|
||||||
|
|||||||
@@ -371,6 +371,9 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
pin_cpu_memory: bool = True
|
pin_cpu_memory: bool = True
|
||||||
ltx2_two_stage_device_mode: str | None = None
|
ltx2_two_stage_device_mode: str | None = None
|
||||||
_explicit_arg_names: set[str] = field(default_factory=set, repr=False)
|
_explicit_arg_names: set[str] = field(default_factory=set, repr=False)
|
||||||
|
_automatic_component_attention_backend_keys: set[str] = field(
|
||||||
|
default_factory=set, init=False, repr=False
|
||||||
|
)
|
||||||
_required_resident_components: set[str] = field(
|
_required_resident_components: set[str] = field(
|
||||||
default_factory=set, init=False, repr=False
|
default_factory=set, init=False, repr=False
|
||||||
)
|
)
|
||||||
@@ -952,6 +955,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
logger.info(
|
logger.info(
|
||||||
"Automatically set torch_sdpa backend for component text_encoder to preserve LTX2 official attention semantics"
|
"Automatically set torch_sdpa backend for component text_encoder to preserve LTX2 official attention semantics"
|
||||||
)
|
)
|
||||||
|
self._automatic_component_attention_backend_keys.add("text_encoder")
|
||||||
else:
|
else:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Overriding %s backend with torch_sdpa for component text_encoder to preserve LTX2 official attention semantics",
|
"Overriding %s backend with torch_sdpa for component text_encoder to preserve LTX2 official attention semantics",
|
||||||
@@ -975,6 +979,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
"encoder; laser_attn applies to the transformer"
|
"encoder; laser_attn applies to the transformer"
|
||||||
)
|
)
|
||||||
self.component_attention_backends["text_encoder"] = "torch_sdpa"
|
self.component_attention_backends["text_encoder"] = "torch_sdpa"
|
||||||
|
self._automatic_component_attention_backend_keys.add("text_encoder")
|
||||||
|
|
||||||
if self.ring_degree > 1:
|
if self.ring_degree > 1:
|
||||||
if (
|
if (
|
||||||
@@ -1161,6 +1166,14 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
return AttentionBackendEnum[backend.upper()], backend_key
|
return AttentionBackendEnum[backend.upper()], backend_key
|
||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
|
def is_component_attention_backend_automatic(
|
||||||
|
self, component_name: str | None
|
||||||
|
) -> bool:
|
||||||
|
return (
|
||||||
|
component_name is not None
|
||||||
|
and component_name in self._automatic_component_attention_backend_keys
|
||||||
|
)
|
||||||
|
|
||||||
def _adjust_warmup(self):
|
def _adjust_warmup(self):
|
||||||
if self.warmup_mode is not None and self.warmup_mode not in WARMUP_MODES:
|
if self.warmup_mode is not None and self.warmup_mode not in WARMUP_MODES:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -164,6 +164,36 @@ class TestAttentionBackendFallback(unittest.TestCase):
|
|||||||
self.assertIs(backend, _FakeFABackend)
|
self.assertIs(backend, _FakeFABackend)
|
||||||
self.assertIsNone(_FakePlatform.selected_backend)
|
self.assertIsNone(_FakePlatform.selected_backend)
|
||||||
|
|
||||||
|
def test_component_override_requires_an_sglang_attention_layer(self):
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ValueError, "did not construct an SGLang attention layer"
|
||||||
|
):
|
||||||
|
with component_attn_backend_context_manager(
|
||||||
|
AttentionBackendEnum.FA, component_name="vae"
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
self.assertIsNone(get_component_attn_backend_context())
|
||||||
|
|
||||||
|
def test_component_override_preserves_load_error_and_resets_context(self):
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "component failed"):
|
||||||
|
with component_attn_backend_context_manager(
|
||||||
|
AttentionBackendEnum.FA, component_name="vae"
|
||||||
|
):
|
||||||
|
raise RuntimeError("component failed")
|
||||||
|
|
||||||
|
self.assertIsNone(get_component_attn_backend_context())
|
||||||
|
|
||||||
|
def test_automatic_component_backend_may_skip_sglang_attention_layer(self):
|
||||||
|
with component_attn_backend_context_manager(
|
||||||
|
AttentionBackendEnum.TORCH_SDPA,
|
||||||
|
component_name="text_encoder",
|
||||||
|
require_component_backend_selection=False,
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
self.assertIsNone(get_component_attn_backend_context())
|
||||||
|
|
||||||
def test_implicit_preference_falls_back_for_missing_capability(self):
|
def test_implicit_preference_falls_back_for_missing_capability(self):
|
||||||
backend = self._resolve(
|
backend = self._resolve(
|
||||||
AttentionBackendEnum.AITER,
|
AttentionBackendEnum.AITER,
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import Mock
|
from unittest.mock import Mock
|
||||||
|
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||||
@@ -8,6 +9,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im
|
|||||||
ComponentUse,
|
ComponentUse,
|
||||||
ResidencyState,
|
ResidencyState,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
||||||
|
ComponentResidencyError,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import (
|
||||||
ComponentOffloadStrategy,
|
ComponentOffloadStrategy,
|
||||||
ResidentStrategy,
|
ResidentStrategy,
|
||||||
@@ -15,6 +19,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_
|
|||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime.text_encoding import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime.text_encoding import (
|
||||||
RealtimeTextEncodingStage,
|
RealtimeTextEncodingStage,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
|
||||||
|
|
||||||
def test_component_offload_releases_preferred_component_after_request():
|
def test_component_offload_releases_preferred_component_after_request():
|
||||||
@@ -171,6 +176,44 @@ def _manager_for_stage(stage, modules):
|
|||||||
return manager, server_args
|
return manager, server_args
|
||||||
|
|
||||||
|
|
||||||
|
def _server_args_with_component_offload(component_name):
|
||||||
|
server_args = ServerArgs.__new__(ServerArgs)
|
||||||
|
server_args.component_residency = {component_name: "component-offload"}
|
||||||
|
return server_args
|
||||||
|
|
||||||
|
|
||||||
|
def test_explicit_component_offload_requires_a_declared_request_use():
|
||||||
|
stage = _Stage()
|
||||||
|
pipeline = SimpleNamespace(
|
||||||
|
modules={"auxiliary": torch.nn.Linear(2, 2)},
|
||||||
|
_stage_name_mapping={"stage": stage},
|
||||||
|
component_residency_strategies={},
|
||||||
|
)
|
||||||
|
server_args = _server_args_with_component_offload("auxiliary")
|
||||||
|
manager = ComponentResidencyManager(pipeline, server_args)
|
||||||
|
manager.refresh_pipeline(pipeline)
|
||||||
|
|
||||||
|
with pytest.raises(
|
||||||
|
ComponentResidencyError,
|
||||||
|
match="'auxiliary'.*ComponentUse declaration",
|
||||||
|
):
|
||||||
|
manager.begin_request([stage], SimpleNamespace(is_warmup=False), server_args)
|
||||||
|
|
||||||
|
|
||||||
|
def test_declared_component_use_admits_explicit_component_offload():
|
||||||
|
stage = _Stage(ComponentUse("stage", "auxiliary"))
|
||||||
|
pipeline = SimpleNamespace(
|
||||||
|
modules={"auxiliary": torch.nn.Linear(2, 2)},
|
||||||
|
_stage_name_mapping={"stage": stage},
|
||||||
|
component_residency_strategies={},
|
||||||
|
)
|
||||||
|
server_args = _server_args_with_component_offload("auxiliary")
|
||||||
|
manager = ComponentResidencyManager(pipeline, server_args)
|
||||||
|
manager.refresh_pipeline(pipeline)
|
||||||
|
|
||||||
|
manager.begin_request([stage], SimpleNamespace(is_warmup=False), server_args)
|
||||||
|
|
||||||
|
|
||||||
def test_single_component_stage_is_prepared_at_stage_entry():
|
def test_single_component_stage_is_prepared_at_stage_entry():
|
||||||
module = torch.nn.Linear(2, 2)
|
module = torch.nn.Linear(2, 2)
|
||||||
use = ComponentUse("stage", "text_encoder")
|
use = ComponentUse("stage", "text_encoder")
|
||||||
|
|||||||
@@ -232,6 +232,28 @@ class TestServerArgsPathExpansion(unittest.TestCase):
|
|||||||
self.assertEqual(backend.name, "TORCH_SDPA")
|
self.assertEqual(backend.name, "TORCH_SDPA")
|
||||||
self.assertEqual(matched_key, "text_encoder")
|
self.assertEqual(matched_key, "text_encoder")
|
||||||
|
|
||||||
|
def test_ltx_automatic_text_encoder_backend_is_not_explicit(self):
|
||||||
|
args = _from_dict_without_model_resolution(
|
||||||
|
{"model_path": "Lightricks/LTX-2.3"},
|
||||||
|
pipeline_config=LTX2PipelineConfig(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
args.component_attention_backends, {"text_encoder": "torch_sdpa"}
|
||||||
|
)
|
||||||
|
self.assertTrue(args.is_component_attention_backend_automatic("text_encoder"))
|
||||||
|
|
||||||
|
def test_ltx_explicit_text_encoder_backend_remains_explicit(self):
|
||||||
|
args = _from_dict_without_model_resolution(
|
||||||
|
{
|
||||||
|
"model_path": "Lightricks/LTX-2.3",
|
||||||
|
"component_attention_backends": {"text_encoder": "torch_sdpa"},
|
||||||
|
},
|
||||||
|
pipeline_config=LTX2PipelineConfig(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertFalse(args.is_component_attention_backend_automatic("text_encoder"))
|
||||||
|
|
||||||
def test_invalid_component_attention_backend_raises(self):
|
def test_invalid_component_attention_backend_raises(self):
|
||||||
with self.assertRaises(ValueError):
|
with self.assertRaises(ValueError):
|
||||||
self._from_dict_without_model_resolution(
|
self._from_dict_without_model_resolution(
|
||||||
|
|||||||
@@ -1,11 +1,14 @@
|
|||||||
|
import re
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.bridge_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.bridge_loader import (
|
||||||
BridgeLoader,
|
BridgeLoader,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentLoader,
|
ComponentLoader,
|
||||||
|
NativeComponentLoaderRequired,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
|
||||||
TransformerLoader,
|
TransformerLoader,
|
||||||
@@ -20,21 +23,110 @@ class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _server_args(*, fsdp_requested=False, **overrides):
|
def _server_args(*, fsdp_requested=False, **overrides):
|
||||||
values = {
|
values = {
|
||||||
|
"component_precisions": {},
|
||||||
|
"component_quantizations": {},
|
||||||
|
"component_weights_paths": {},
|
||||||
|
"component_quantization_ignored_layers": {},
|
||||||
|
"transformer_weights_path": None,
|
||||||
|
"nunchaku_config": None,
|
||||||
|
"quantization": None,
|
||||||
|
"pipeline_config": SimpleNamespace(native_only_components=()),
|
||||||
"tp_size": 1,
|
"tp_size": 1,
|
||||||
"sp_degree": 1,
|
"sp_degree": 1,
|
||||||
"ulysses_degree": 1,
|
"ulysses_degree": 1,
|
||||||
"ring_degree": 1,
|
"ring_degree": 1,
|
||||||
"should_use_fsdp_for_component": lambda _component: fsdp_requested,
|
"kv_gather_degree": 1,
|
||||||
|
"enable_cfg_parallel": False,
|
||||||
|
"dp_size": 1,
|
||||||
|
"use_fsdp_inference": False,
|
||||||
|
"resolve_component_attention_backend": mock.Mock(return_value=(None, None)),
|
||||||
|
"should_use_fsdp_for_component": mock.Mock(return_value=fsdp_requested),
|
||||||
}
|
}
|
||||||
values.update(overrides)
|
values.update(overrides)
|
||||||
return SimpleNamespace(**values)
|
return SimpleNamespace(**values)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _mocked_load(loader):
|
||||||
|
customized_load = mock.patch.object(
|
||||||
|
loader,
|
||||||
|
"_load_customized_with_context",
|
||||||
|
side_effect=NativeComponentLoaderRequired("native loader required"),
|
||||||
|
)
|
||||||
|
native_load = mock.patch.object(
|
||||||
|
loader, "_load_native_with_context", return_value=object()
|
||||||
|
)
|
||||||
|
available_memory = mock.patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"component_loader.current_platform.get_available_gpu_memory",
|
||||||
|
return_value=0.0,
|
||||||
|
)
|
||||||
|
return customized_load, native_load, available_memory
|
||||||
|
|
||||||
|
def test_distributed_execution_rejects_before_native_load(self):
|
||||||
|
cases = (
|
||||||
|
("tp_size", 2, "tp_size=2"),
|
||||||
|
("sp_degree", 2, "sp_degree=2"),
|
||||||
|
("ulysses_degree", 2, "ulysses_degree=2"),
|
||||||
|
("ring_degree", 2, "ring_degree=2"),
|
||||||
|
("kv_gather_degree", 2, "kv_gather_degree=2"),
|
||||||
|
("fsdp_requested", True, "FSDP"),
|
||||||
|
)
|
||||||
|
|
||||||
|
for field, value, expected_error in cases:
|
||||||
|
with self.subTest(field=field):
|
||||||
|
loader = TransformerLoader()
|
||||||
|
server_args = self._server_args(**{field: value})
|
||||||
|
customized_load, native_load, available_memory = self._mocked_load(
|
||||||
|
loader
|
||||||
|
)
|
||||||
|
|
||||||
|
with customized_load, native_load as native, available_memory:
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
RuntimeError, re.escape(expected_error)
|
||||||
|
):
|
||||||
|
loader.load(
|
||||||
|
"/model/transformer_2",
|
||||||
|
server_args,
|
||||||
|
"transformer_2",
|
||||||
|
"diffusers",
|
||||||
|
)
|
||||||
|
|
||||||
|
native.assert_not_called()
|
||||||
|
server_args.should_use_fsdp_for_component.assert_called_once_with(
|
||||||
|
"transformer_2"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_replicated_cfg_and_dp_keep_native_fallback(self):
|
||||||
|
loader = TransformerLoader()
|
||||||
|
server_args = self._server_args(
|
||||||
|
enable_cfg_parallel=True,
|
||||||
|
dp_size=2,
|
||||||
|
use_fsdp_inference=True,
|
||||||
|
)
|
||||||
|
customized_load, native_load, available_memory = self._mocked_load(loader)
|
||||||
|
|
||||||
|
with customized_load, native_load as native, available_memory:
|
||||||
|
component, consumed = loader.load(
|
||||||
|
"/model/transformer_2",
|
||||||
|
server_args,
|
||||||
|
"transformer_2",
|
||||||
|
"diffusers",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNotNone(component)
|
||||||
|
self.assertEqual(consumed, 0.0)
|
||||||
|
native.assert_called_once()
|
||||||
|
server_args.should_use_fsdp_for_component.assert_called_once_with(
|
||||||
|
"transformer_2"
|
||||||
|
)
|
||||||
|
|
||||||
def test_parallel_execution_rejects_native_fallback(self):
|
def test_parallel_execution_rejects_native_fallback(self):
|
||||||
cases = (
|
cases = (
|
||||||
({"tp_size": 2}, "tp_size=2"),
|
({"tp_size": 2}, "tp_size=2"),
|
||||||
({"sp_degree": 2}, "sp_degree=2"),
|
({"sp_degree": 2}, "sp_degree=2"),
|
||||||
({"ulysses_degree": 2}, "ulysses_degree=2"),
|
({"ulysses_degree": 2}, "ulysses_degree=2"),
|
||||||
({"ring_degree": 2}, "ring_degree=2"),
|
({"ring_degree": 2}, "ring_degree=2"),
|
||||||
|
({"kv_gather_degree": 2}, "kv_gather_degree=2"),
|
||||||
({"fsdp_requested": True}, "FSDP"),
|
({"fsdp_requested": True}, "FSDP"),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -55,6 +55,9 @@ class _FakeServerArgs:
|
|||||||
def should_start_component_on_cpu(self, _component_name):
|
def should_start_component_on_cpu(self, _component_name):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def should_use_fsdp_for_component(self, _component_name):
|
||||||
|
return False
|
||||||
|
|
||||||
def should_configure_layerwise_offload_for_lazy_component(self, component_name):
|
def should_configure_layerwise_offload_for_lazy_component(self, component_name):
|
||||||
return component_name in self.layerwise_components
|
return component_name in self.layerwise_components
|
||||||
|
|
||||||
@@ -412,6 +415,12 @@ class TestVAELoader(unittest.TestCase):
|
|||||||
|
|
||||||
native_load.assert_not_called()
|
native_load.assert_not_called()
|
||||||
|
|
||||||
|
def test_pipeline_config_declares_an_empty_native_only_default(self):
|
||||||
|
loader = vae_loader.VAELoader()
|
||||||
|
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
||||||
|
|
||||||
|
self.assertFalse(loader.should_raise_customized_load_error(server_args, "vae"))
|
||||||
|
|
||||||
def test_backfill_ltx2_audio_vae_latent_stats_maps_official_keys(self):
|
def test_backfill_ltx2_audio_vae_latent_stats_maps_official_keys(self):
|
||||||
loaded = {
|
loaded = {
|
||||||
"per_channel_statistics.mean-of-means": torch.tensor([1.0, 2.0]),
|
"per_channel_statistics.mean-of-means": torch.tensor([1.0, 2.0]),
|
||||||
|
|||||||
Reference in New Issue
Block a user