[diffusion] chore: enforce component attention backend application (#36907)
This commit is contained in:
@@ -429,7 +429,7 @@ sglang generate \
|
||||
--component-attention-backends text_encoder=torch_sdpa
|
||||
```
|
||||
|
||||
The component key must match a pipeline module key such as `text_encoder`, `text_encoder_2`, `transformer`, `transformer_2`, or `connectors`. Component overrides take precedence over the global `--attention-backend` only while that component is being constructed and otherwise fail if the component cannot satisfy them. Sparse self-attention backends use a compatible dense backend for cross-attention layers. The global backend remains strict for DiT components, while auxiliary components may fall back to a compatible backend.
|
||||
The component key must match a pipeline module key such as `text_encoder`, `text_encoder_2`, `transformer`, `transformer_2`, or `connectors`. Component overrides take precedence over the global `--attention-backend` while that component is being constructed and fail if the component cannot satisfy them. A native component may explicitly defer backend selection until first use; components with fixed attention reject the override. Sparse self-attention backends use a compatible dense backend for cross-attention layers. The global backend remains strict for DiT components, while auxiliary components may fall back to a compatible backend. The Diffusers backend supports only the global backend passthrough.
|
||||
|
||||
You can also pass dotted CLI entries:
|
||||
|
||||
|
||||
@@ -642,6 +642,10 @@ Use this override when the fallback must be pinned: unlike the global backend,
|
||||
an incompatible component override raises an error instead of selecting another
|
||||
backend. The one role-based exception is a sparse self-attention backend, which
|
||||
uses a compatible dense backend for cross-attention layers in the same component.
|
||||
The component must construct SGLang-selectable attention or explicitly defer
|
||||
selection until first use; components with fixed attention reject the override.
|
||||
Per-component overrides apply only to native pipelines. The Diffusers backend
|
||||
accepts the global `--attention-backend` passthrough instead.
|
||||
|
||||
### Per-request override (denoise loop)
|
||||
|
||||
|
||||
@@ -70,6 +70,11 @@ class ComponentAttnBackendContext(NamedTuple):
|
||||
component_name: str | None
|
||||
selected_backends: dict[str, str | None]
|
||||
allow_global_backend_fallback: bool = False
|
||||
require_backend_selection: bool = False
|
||||
|
||||
|
||||
class ComponentAttentionBackendNotAppliedError(ValueError):
|
||||
"""An explicit component backend did not control its attention layers."""
|
||||
|
||||
|
||||
component_attn_backend_context: ContextVar[ComponentAttnBackendContext | None] = (
|
||||
@@ -109,6 +114,17 @@ def get_component_forced_attn_backend() -> AttentionBackendEnum | None:
|
||||
return context.backend if context is not None else None
|
||||
|
||||
|
||||
def claim_deferred_component_attn_backend() -> AttentionBackendEnum | None:
|
||||
"""Capture an override whose compatible backend is resolved on first use."""
|
||||
context = get_component_attn_backend_context()
|
||||
if context is None or context.backend is None:
|
||||
return None
|
||||
_record_component_attn_backend(
|
||||
context.backend.name.lower(), "deferred first-use selection"
|
||||
)
|
||||
return context.backend
|
||||
|
||||
|
||||
def get_component_attn_backend_name() -> str | None:
|
||||
context = get_component_attn_backend_context()
|
||||
return context.component_name if context is not None else None
|
||||
@@ -124,9 +140,11 @@ def _record_component_attn_backend(backend_name: str, reason: str | None) -> boo
|
||||
if context is None or context.component_name is None:
|
||||
return False
|
||||
|
||||
existing_reason = context.selected_backends.get(backend_name)
|
||||
if backend_name not in context.selected_backends or existing_reason is None:
|
||||
if backend_name not in context.selected_backends:
|
||||
context.selected_backends[backend_name] = reason
|
||||
elif reason is None:
|
||||
# unrestricted selection must not be hidden by a later valid fallback
|
||||
context.selected_backends[backend_name] = None
|
||||
return True
|
||||
|
||||
|
||||
@@ -160,6 +178,40 @@ def _log_component_attn_backend_summary(
|
||||
)
|
||||
|
||||
|
||||
def _validate_component_attn_backend_selection(
|
||||
context: ComponentAttnBackendContext,
|
||||
) -> None:
|
||||
if not context.require_backend_selection:
|
||||
return
|
||||
|
||||
requested_backend = context.backend
|
||||
assert requested_backend is not None
|
||||
requested_name = requested_backend.name.lower()
|
||||
component_name = context.component_name or "component"
|
||||
if requested_name not in context.selected_backends:
|
||||
detail = (
|
||||
"did not construct any SGLang-selectable attention layers"
|
||||
if not context.selected_backends
|
||||
else f"selected {', '.join(sorted(context.selected_backends))} instead"
|
||||
)
|
||||
raise ComponentAttentionBackendNotAppliedError(
|
||||
f"Attention backend '{requested_name}' was requested for component "
|
||||
f"'{component_name}', but it {detail}"
|
||||
)
|
||||
|
||||
unexplained = sorted(
|
||||
backend_name
|
||||
for backend_name, reason in context.selected_backends.items()
|
||||
if backend_name != requested_name and reason is None
|
||||
)
|
||||
if unexplained:
|
||||
raise ComponentAttentionBackendNotAppliedError(
|
||||
f"Attention backend '{requested_name}' was requested for component "
|
||||
f"'{component_name}', but it also selected "
|
||||
f"{', '.join(unexplained)} without an allowed fallback"
|
||||
)
|
||||
|
||||
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
@@ -372,49 +424,38 @@ def component_attn_backend_context_manager(
|
||||
attn_backend: AttentionBackendEnum | None,
|
||||
component_name: str | None = None,
|
||||
allow_global_backend_fallback: bool = False,
|
||||
require_component_backend_selection: bool = True,
|
||||
require_backend_selection: bool | None = None,
|
||||
require_component_backend_selection: bool | None = None,
|
||||
) -> Generator[None, None, None]:
|
||||
if attn_backend is None and component_name is None:
|
||||
yield
|
||||
return
|
||||
|
||||
if require_backend_selection is None:
|
||||
require_backend_selection = (
|
||||
require_component_backend_selection
|
||||
if require_component_backend_selection is not None
|
||||
else attn_backend is not None
|
||||
)
|
||||
elif require_component_backend_selection is not None:
|
||||
raise ValueError("Specify only one component backend selection requirement")
|
||||
|
||||
token = component_attn_backend_context.set(
|
||||
ComponentAttnBackendContext(
|
||||
attn_backend,
|
||||
component_name,
|
||||
{},
|
||||
allow_global_backend_fallback,
|
||||
require_backend_selection,
|
||||
)
|
||||
)
|
||||
unused_component_name: str | None = None
|
||||
unused_backend_name: str | None = None
|
||||
completed = False
|
||||
try:
|
||||
yield
|
||||
completed = True
|
||||
finally:
|
||||
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()
|
||||
_validate_component_attn_backend_selection(context)
|
||||
_log_component_attn_backend_summary(context)
|
||||
finally:
|
||||
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
|
||||
|
||||
@@ -23,6 +23,7 @@ from transformers.quantizers import AutoHfQuantizer
|
||||
|
||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
ComponentAttentionBackendNotAppliedError,
|
||||
component_attn_backend_context_manager,
|
||||
get_component_attn_backend_context,
|
||||
)
|
||||
@@ -217,17 +218,13 @@ class ComponentLoader(ABC):
|
||||
attn_backend: Any,
|
||||
component_attn_name: str | None,
|
||||
allow_global_backend_fallback: bool,
|
||||
require_backend_selection: bool,
|
||||
) -> AutoModel:
|
||||
with component_attn_backend_context_manager(
|
||||
attn_backend,
|
||||
component_name=component_attn_name,
|
||||
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
|
||||
)
|
||||
),
|
||||
require_backend_selection=require_backend_selection,
|
||||
):
|
||||
load_kwargs = self.customized_load_kwargs_for_component(
|
||||
server_args, component_name
|
||||
@@ -245,17 +242,13 @@ class ComponentLoader(ABC):
|
||||
attn_backend: Any,
|
||||
component_attn_name: str | None,
|
||||
allow_global_backend_fallback: bool,
|
||||
require_backend_selection: bool,
|
||||
) -> AutoModel:
|
||||
with component_attn_backend_context_manager(
|
||||
attn_backend,
|
||||
component_name=component_attn_name,
|
||||
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
|
||||
)
|
||||
),
|
||||
require_backend_selection=require_backend_selection,
|
||||
):
|
||||
component = self.load_native(
|
||||
component_model_path,
|
||||
@@ -271,6 +264,9 @@ class ComponentLoader(ABC):
|
||||
server_args: ServerArgs,
|
||||
component_name: str,
|
||||
transformers_or_diffusers: str,
|
||||
*,
|
||||
component_attn_backend: Any = None,
|
||||
component_attn_name: str | None = None,
|
||||
) -> tuple[AutoModel, float]:
|
||||
"""
|
||||
Template method that standardizes logging around the core load implementation.
|
||||
@@ -307,32 +303,55 @@ class ComponentLoader(ABC):
|
||||
component_model_path,
|
||||
gpu_mem_before_loading,
|
||||
)
|
||||
attn_backend = None
|
||||
component_attn_name = None
|
||||
if get_component_attn_backend_context() is None:
|
||||
attn_backend, matched_backend_key = (
|
||||
if (
|
||||
component_attn_backend is None
|
||||
and component_attn_name is None
|
||||
and get_component_attn_backend_context() is None
|
||||
):
|
||||
component_attn_backend, matched_backend_key = (
|
||||
server_args.resolve_component_attention_backend(component_name)
|
||||
)
|
||||
component_attn_name = matched_backend_key or component_name
|
||||
if attn_backend is not None:
|
||||
if component_attn_backend is not None:
|
||||
logger.info(
|
||||
"Using %s backend for component: %s",
|
||||
attn_backend.name.lower(),
|
||||
component_attn_backend.name.lower(),
|
||||
matched_backend_key,
|
||||
)
|
||||
requested_backend = (
|
||||
server_args.requested_component_attention_backend(component_attn_name)
|
||||
if component_attn_name is not None
|
||||
else None
|
||||
)
|
||||
require_backend_selection = requested_backend is not None
|
||||
if require_backend_selection and (
|
||||
component_attn_backend is None
|
||||
or component_attn_backend.name.lower() != requested_backend
|
||||
):
|
||||
raise ValueError(
|
||||
f"Component attention backend for {component_attn_name!r} no longer "
|
||||
f"matches the explicit request {requested_backend!r}"
|
||||
)
|
||||
try:
|
||||
component = self._load_customized_with_context(
|
||||
component_model_path,
|
||||
server_args,
|
||||
component_name,
|
||||
attn_backend,
|
||||
component_attn_backend,
|
||||
component_attn_name,
|
||||
self.allow_global_attention_backend_fallback,
|
||||
require_backend_selection,
|
||||
)
|
||||
source = "sgl-diffusion"
|
||||
except (ComponentCheckpointUnsupportedError, ComponentResidencyError):
|
||||
except (
|
||||
ComponentAttentionBackendNotAppliedError,
|
||||
ComponentCheckpointUnsupportedError,
|
||||
ComponentResidencyError,
|
||||
):
|
||||
raise
|
||||
except Exception as e:
|
||||
if require_backend_selection:
|
||||
raise
|
||||
native_loader_required = isinstance(e, NativeComponentLoaderRequired)
|
||||
if self.should_raise_customized_load_error(server_args, component_name):
|
||||
if native_loader_required:
|
||||
@@ -360,9 +379,10 @@ class ComponentLoader(ABC):
|
||||
server_args,
|
||||
component_name,
|
||||
transformers_or_diffusers,
|
||||
attn_backend,
|
||||
component_attn_backend,
|
||||
component_attn_name,
|
||||
self.allow_global_attention_backend_fallback,
|
||||
require_backend_selection,
|
||||
)
|
||||
source = "native"
|
||||
logger.warning(
|
||||
@@ -756,24 +776,13 @@ class PipelineComponentLoader:
|
||||
)
|
||||
|
||||
try:
|
||||
with component_attn_backend_context_manager(
|
||||
component_attn_backend,
|
||||
component_name=component_attn_name,
|
||||
allow_global_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(
|
||||
component_model_path,
|
||||
server_args,
|
||||
component_name,
|
||||
transformers_or_diffusers,
|
||||
component_attn_backend=component_attn_backend,
|
||||
component_attn_name=component_attn_name,
|
||||
)
|
||||
except Exception:
|
||||
logger.error(
|
||||
|
||||
@@ -52,10 +52,9 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
|
||||
AttentionRequirements,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
claim_deferred_component_attn_backend,
|
||||
get_attn_backend,
|
||||
get_component_forced_attn_backend,
|
||||
get_global_forced_attn_backend,
|
||||
record_component_attn_backend,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
@@ -2011,11 +2010,8 @@ class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
||||
)
|
||||
# Component overrides disappear when the loader context exits. Preserve
|
||||
# only that selection; process-wide overrides are resolved at first use.
|
||||
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._component_attention_backend_override = (
|
||||
claim_deferred_component_attn_backend()
|
||||
)
|
||||
self._resolved_attention_backend: AttentionBackendEnum | None = None
|
||||
self._mark_missing_params_required()
|
||||
|
||||
@@ -371,6 +371,12 @@ class DiffusersPipeline(ComposedPipelineBase):
|
||||
loaded_modules: dict[str, torch.nn.Module] | None = None,
|
||||
executor: PipelineExecutor | None = None,
|
||||
):
|
||||
if server_args.has_requested_component_attention_backends():
|
||||
raise ValueError(
|
||||
"--component-attention-backends is supported only by native "
|
||||
"SGLang diffusion pipelines; use --attention-backend with the "
|
||||
"Diffusers backend"
|
||||
)
|
||||
self.server_args = server_args
|
||||
self.model_path = model_path
|
||||
self._stages: list[PipelineStage] = []
|
||||
|
||||
@@ -251,6 +251,9 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
component_attention_backends: dict[str, str] | str | None = field(
|
||||
default_factory=dict
|
||||
)
|
||||
_requested_component_attention_backends: dict[str, str] | None = field(
|
||||
default=None, repr=False, compare=False
|
||||
)
|
||||
cache_dit_config: str | dict[str, Any] | None = (
|
||||
None # cache-dit config for diffusers
|
||||
)
|
||||
@@ -958,6 +961,16 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
self.component_attention_backends
|
||||
)
|
||||
)
|
||||
if self._requested_component_attention_backends is None:
|
||||
self._requested_component_attention_backends = dict(
|
||||
self.component_attention_backends
|
||||
)
|
||||
else:
|
||||
self._requested_component_attention_backends = (
|
||||
self._normalize_component_attention_backends(
|
||||
self._requested_component_attention_backends
|
||||
)
|
||||
)
|
||||
|
||||
# attention_backend_config
|
||||
if self.attention_backend_config is None:
|
||||
@@ -1187,6 +1200,13 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
return AttentionBackendEnum[backend.upper()], backend_key
|
||||
return None, None
|
||||
|
||||
def requested_component_attention_backend(self, component_name: str) -> str | None:
|
||||
assert self._requested_component_attention_backends is not None
|
||||
return self._requested_component_attention_backends.get(component_name)
|
||||
|
||||
def has_requested_component_attention_backends(self) -> bool:
|
||||
return bool(self._requested_component_attention_backends)
|
||||
|
||||
def is_component_attention_backend_automatic(
|
||||
self, component_name: str | None
|
||||
) -> bool:
|
||||
|
||||
@@ -8,7 +8,10 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
|
||||
AttentionRequirements,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
ComponentAttentionBackendNotAppliedError,
|
||||
_cached_get_attn_backend,
|
||||
_record_component_attn_backend,
|
||||
claim_deferred_component_attn_backend,
|
||||
component_attn_backend_context_manager,
|
||||
get_attn_backend,
|
||||
get_component_attn_backend_context,
|
||||
@@ -16,6 +19,7 @@ from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||
ComponentLoader,
|
||||
GenericComponentLoader,
|
||||
NativeComponentLoaderRequired,
|
||||
PipelineComponentLoader,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
|
||||
@@ -25,6 +29,9 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader i
|
||||
TransformerLoader,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import VAELoader
|
||||
from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import (
|
||||
DiffusersPipeline,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
@@ -142,6 +149,7 @@ class TestAttentionBackendFallback(unittest.TestCase):
|
||||
component_backend,
|
||||
component_name="text_encoder",
|
||||
allow_global_backend_fallback=allow_global_backend_fallback,
|
||||
require_backend_selection=component_backend is not None,
|
||||
),
|
||||
):
|
||||
return get_attn_backend(
|
||||
@@ -166,7 +174,7 @@ class TestAttentionBackendFallback(unittest.TestCase):
|
||||
|
||||
def test_component_override_requires_an_sglang_attention_layer(self):
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "did not construct an SGLang attention layer"
|
||||
ValueError, "did not construct any SGLang-selectable attention layers"
|
||||
):
|
||||
with component_attn_backend_context_manager(
|
||||
AttentionBackendEnum.FA, component_name="vae"
|
||||
@@ -269,6 +277,17 @@ class TestAttentionBackendFallback(unittest.TestCase):
|
||||
allow_global_backend_fallback=True,
|
||||
)
|
||||
|
||||
def test_explicit_component_backend_is_consumed(self):
|
||||
backend = self._resolve(
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
explicit=True,
|
||||
is_cross_attention=False,
|
||||
supported={AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA},
|
||||
component_backend=AttentionBackendEnum.FA,
|
||||
)
|
||||
|
||||
self.assertIs(backend, _FakeFABackend)
|
||||
|
||||
def test_sparse_backend_falls_back_for_cross_attention(self):
|
||||
backend = self._resolve(
|
||||
AttentionBackendEnum.LASER_ATTN,
|
||||
@@ -307,21 +326,40 @@ class TestComponentAttentionBackendScope(unittest.TestCase):
|
||||
def _load_with_policy(self, allow_global_backend_fallback: bool):
|
||||
captured_context = None
|
||||
|
||||
class _Loader:
|
||||
def load(self, *_args):
|
||||
class _Loader(ComponentLoader):
|
||||
def load_customized(self, *_args):
|
||||
nonlocal captured_context
|
||||
captured_context = get_component_attn_backend_context()
|
||||
return object(), 0.0
|
||||
return object()
|
||||
|
||||
class _Args:
|
||||
component_quantizations = {}
|
||||
|
||||
@staticmethod
|
||||
def requested_component_attention_backend(_component_name):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def should_direct_gpu_weight_load_component(_component_name):
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def should_use_fsdp_for_component(_component_name):
|
||||
return False
|
||||
|
||||
_Loader.allow_global_attention_backend_fallback = allow_global_backend_fallback
|
||||
with patch.object(
|
||||
ComponentLoader, "for_component_type", return_value=_Loader()
|
||||
with (
|
||||
patch.object(ComponentLoader, "for_component_type", return_value=_Loader()),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.loader.component_loaders.component_loader.current_platform.get_available_gpu_memory",
|
||||
return_value=1.0,
|
||||
),
|
||||
):
|
||||
PipelineComponentLoader.load_component(
|
||||
component_name="text_encoder",
|
||||
component_model_path="unused",
|
||||
transformers_or_diffusers="transformers",
|
||||
server_args=object(),
|
||||
server_args=_Args(),
|
||||
component_attn_name="text_encoder",
|
||||
)
|
||||
return captured_context
|
||||
@@ -344,6 +382,196 @@ class TestComponentAttentionBackendScope(unittest.TestCase):
|
||||
self.assertTrue(TextEncoderLoader.allow_global_attention_backend_fallback)
|
||||
self.assertTrue(VAELoader.allow_global_attention_backend_fallback)
|
||||
|
||||
def test_explicit_backend_must_be_consumed(self):
|
||||
with self.assertRaisesRegex(
|
||||
ComponentAttentionBackendNotAppliedError,
|
||||
"did not construct any SGLang-selectable attention layers",
|
||||
):
|
||||
with component_attn_backend_context_manager(
|
||||
AttentionBackendEnum.FA,
|
||||
component_name="image_encoder",
|
||||
require_backend_selection=True,
|
||||
):
|
||||
pass
|
||||
|
||||
def test_deferred_selection_satisfies_construction_contract(self):
|
||||
with component_attn_backend_context_manager(
|
||||
AttentionBackendEnum.FA,
|
||||
component_name="transformer",
|
||||
require_backend_selection=True,
|
||||
):
|
||||
self.assertIs(
|
||||
claim_deferred_component_attn_backend(),
|
||||
AttentionBackendEnum.FA,
|
||||
)
|
||||
|
||||
def test_fixed_component_load_rejects_explicit_backend(self):
|
||||
class _Loader(ComponentLoader):
|
||||
def load_customized(self, *_args):
|
||||
return object()
|
||||
|
||||
class _Args:
|
||||
component_quantizations = {}
|
||||
|
||||
@staticmethod
|
||||
def requested_component_attention_backend(_component_name):
|
||||
return "fa"
|
||||
|
||||
@staticmethod
|
||||
def should_direct_gpu_weight_load_component(_component_name):
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def should_use_fsdp_for_component(_component_name):
|
||||
return False
|
||||
|
||||
with (
|
||||
patch.object(ComponentLoader, "for_component_type", return_value=_Loader()),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.loader.component_loaders.component_loader.current_platform.get_available_gpu_memory",
|
||||
return_value=1.0,
|
||||
),
|
||||
self.assertRaisesRegex(
|
||||
ComponentAttentionBackendNotAppliedError,
|
||||
"did not construct any SGLang-selectable attention layers",
|
||||
),
|
||||
):
|
||||
PipelineComponentLoader.load_component(
|
||||
component_name="image_encoder",
|
||||
component_model_path="unused",
|
||||
transformers_or_diffusers="transformers",
|
||||
server_args=_Args(),
|
||||
component_attn_backend=AttentionBackendEnum.FA,
|
||||
component_attn_name="image_encoder",
|
||||
)
|
||||
|
||||
def test_unexplained_mixed_backend_is_rejected(self):
|
||||
with self.assertRaisesRegex(
|
||||
ComponentAttentionBackendNotAppliedError,
|
||||
"also selected torch_sdpa without an allowed fallback",
|
||||
):
|
||||
with component_attn_backend_context_manager(
|
||||
AttentionBackendEnum.FA,
|
||||
component_name="transformer",
|
||||
require_backend_selection=True,
|
||||
):
|
||||
_record_component_attn_backend("fa", None)
|
||||
_record_component_attn_backend("torch_sdpa", None)
|
||||
_record_component_attn_backend(
|
||||
"torch_sdpa", "dense cross-attention fallback"
|
||||
)
|
||||
|
||||
def test_explicit_backend_preserves_customized_load_failure(self):
|
||||
native_load_called = False
|
||||
|
||||
class _Loader(ComponentLoader):
|
||||
def load_customized(self, *_args):
|
||||
claim_deferred_component_attn_backend()
|
||||
raise RuntimeError("customized load failed")
|
||||
|
||||
def load_native(self, *_args):
|
||||
nonlocal native_load_called
|
||||
native_load_called = True
|
||||
return object()
|
||||
|
||||
class _Args:
|
||||
component_quantizations = {}
|
||||
pipeline_config = SimpleNamespace(native_only_components=())
|
||||
|
||||
@staticmethod
|
||||
def requested_component_attention_backend(_component_name):
|
||||
return "fa"
|
||||
|
||||
@staticmethod
|
||||
def should_direct_gpu_weight_load_component(_component_name):
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def should_use_fsdp_for_component(_component_name):
|
||||
return False
|
||||
|
||||
with (
|
||||
patch.object(ComponentLoader, "for_component_type", return_value=_Loader()),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.loader.component_loaders.component_loader.current_platform.get_available_gpu_memory",
|
||||
return_value=1.0,
|
||||
),
|
||||
self.assertRaisesRegex(RuntimeError, "customized load failed"),
|
||||
):
|
||||
PipelineComponentLoader.load_component(
|
||||
component_name="text_encoder",
|
||||
component_model_path="unused",
|
||||
transformers_or_diffusers="transformers",
|
||||
server_args=_Args(),
|
||||
component_attn_backend=AttentionBackendEnum.FA,
|
||||
component_attn_name="text_encoder",
|
||||
)
|
||||
self.assertFalse(native_load_called)
|
||||
|
||||
def test_legacy_fallback_uses_a_fresh_selection_context(self):
|
||||
customized_context = None
|
||||
native_context = None
|
||||
|
||||
class _Loader(ComponentLoader):
|
||||
def load_customized(self, *_args):
|
||||
nonlocal customized_context
|
||||
customized_context = get_component_attn_backend_context()
|
||||
_record_component_attn_backend("fa", None)
|
||||
raise NativeComponentLoaderRequired("use native loader")
|
||||
|
||||
def load_native(self, *_args):
|
||||
nonlocal native_context
|
||||
native_context = get_component_attn_backend_context()
|
||||
return object()
|
||||
|
||||
class _Args:
|
||||
component_quantizations = {}
|
||||
pipeline_config = SimpleNamespace(native_only_components=())
|
||||
|
||||
@staticmethod
|
||||
def requested_component_attention_backend(_component_name):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def should_direct_gpu_weight_load_component(_component_name):
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def should_use_fsdp_for_component(_component_name):
|
||||
return False
|
||||
|
||||
with (
|
||||
patch.object(ComponentLoader, "for_component_type", return_value=_Loader()),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.loader.component_loaders.component_loader.current_platform.get_available_gpu_memory",
|
||||
return_value=1.0,
|
||||
),
|
||||
):
|
||||
PipelineComponentLoader.load_component(
|
||||
component_name="text_encoder",
|
||||
component_model_path="unused",
|
||||
transformers_or_diffusers="transformers",
|
||||
server_args=_Args(),
|
||||
component_attn_name="text_encoder",
|
||||
)
|
||||
|
||||
self.assertIsNotNone(customized_context)
|
||||
self.assertIsNotNone(native_context)
|
||||
self.assertIsNot(customized_context, native_context)
|
||||
self.assertEqual(customized_context.selected_backends, {"fa": None})
|
||||
self.assertEqual(native_context.selected_backends, {})
|
||||
|
||||
def test_diffusers_backend_rejects_component_override(self):
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "supported only by native SGLang diffusion pipelines"
|
||||
):
|
||||
DiffusersPipeline(
|
||||
"/unused",
|
||||
SimpleNamespace(
|
||||
has_requested_component_attention_backends=lambda: True
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -43,6 +43,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase):
|
||||
component_precisions={},
|
||||
encoder_parallel="replicate",
|
||||
resolve_component_attention_backend=lambda _name: (None, None),
|
||||
requested_component_attention_backend=lambda _name: None,
|
||||
should_direct_gpu_weight_load_component=lambda _name: False,
|
||||
should_use_fsdp_for_component=lambda _name: False,
|
||||
)
|
||||
@@ -252,6 +253,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase):
|
||||
native_only_components=(),
|
||||
),
|
||||
resolve_component_attention_backend=lambda _name: (None, None),
|
||||
requested_component_attention_backend=lambda _name: None,
|
||||
explicit_residency_mode=lambda _name: None,
|
||||
require_component_resident=mock.Mock(),
|
||||
should_use_fsdp_for_component=lambda _name: False,
|
||||
|
||||
@@ -216,6 +216,21 @@ class TestServerArgsPathExpansion(unittest.TestCase):
|
||||
args.component_attention_backends,
|
||||
{"text_encoder": "torch_sdpa", "transformer": "fa"},
|
||||
)
|
||||
self.assertEqual(
|
||||
args._requested_component_attention_backends,
|
||||
args.component_attention_backends,
|
||||
)
|
||||
|
||||
def test_pipeline_attention_default_is_not_an_explicit_override(self):
|
||||
args = _from_dict_without_model_resolution(
|
||||
{"model_path": "/data/my-model"},
|
||||
pipeline_config=LTX2PipelineConfig(),
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
args.component_attention_backends, {"text_encoder": "torch_sdpa"}
|
||||
)
|
||||
self.assertFalse(args.has_requested_component_attention_backends())
|
||||
|
||||
def test_component_attention_backend_lookup(self):
|
||||
args = self._from_dict_without_model_resolution(
|
||||
|
||||
@@ -40,6 +40,7 @@ class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
|
||||
"dp_size": 1,
|
||||
"use_fsdp_inference": False,
|
||||
"resolve_component_attention_backend": mock.Mock(return_value=(None, None)),
|
||||
"requested_component_attention_backend": mock.Mock(return_value=None),
|
||||
"should_direct_gpu_weight_load_component": mock.Mock(return_value=False),
|
||||
"should_use_fsdp_for_component": mock.Mock(return_value=fsdp_requested),
|
||||
}
|
||||
|
||||
@@ -56,6 +56,9 @@ class _FakeServerArgs:
|
||||
def resolve_component_attention_backend(self, _component_name):
|
||||
return None, None
|
||||
|
||||
def requested_component_attention_backend(self, _component_name):
|
||||
return None
|
||||
|
||||
def should_start_component_on_cpu(self, _component_name):
|
||||
return False
|
||||
|
||||
|
||||
Reference in New Issue
Block a user