From 62c470697e243050b48104c9d1d3334cf91f6ca7 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 31 Aug 2026 14:00:02 +0800 Subject: [PATCH] [diffusion] chore: enforce component attention backend application (#36907) --- docs/docs/sglang-diffusion/api/cli.mdx | 2 +- .../sglang-diffusion/attention_backends.mdx | 4 + .../runtime/layers/attention/selector.py | 95 +++++-- .../component_loaders/component_loader.py | 89 ++++--- .../runtime/models/dits/minimax_h3.py | 12 +- .../runtime/pipelines/diffusers_pipeline.py | 6 + .../runtime/server_args/server_args.py | 20 ++ .../unit/test_attention_backend_selector.py | 242 +++++++++++++++++- .../test/unit/test_image_encoder_loader.py | 2 + .../test/unit/test_server_args.py | 15 ++ .../unit/test_transformer_loader_fallback.py | 1 + .../test/unit/test_vae_loader.py | 3 + 12 files changed, 408 insertions(+), 83 deletions(-) diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index 56120ad6f..f5d7fe842 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -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: diff --git a/docs/docs/sglang-diffusion/attention_backends.mdx b/docs/docs/sglang-diffusion/attention_backends.mdx index 4b9cb952c..5cc57575e 100644 --- a/docs/docs/sglang-diffusion/attention_backends.mdx +++ b/docs/docs/sglang-diffusion/attention_backends.mdx @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py index ff5a57cb3..49fe1ca10 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index c0b45afb8..b116cc971 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -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,25 +776,14 @@ 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, - ) + 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( f"Error while loading component: {component_name}, {component_model_path=}" diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py index 58e588845..03ce2b2cc 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py @@ -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,12 +2010,9 @@ 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() diff --git a/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py index 27c24b093..4774252dd 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py @@ -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] = [] diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index 0e7d92179..c94286221 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -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: diff --git a/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py b/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py index fd1427343..5ab486125 100644 --- a/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py +++ b/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py @@ -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() diff --git a/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py index f7c97e090..675f8224f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py @@ -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, diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 3a76697a2..2c1f29036 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -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( diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_loader_fallback.py b/python/sglang/multimodal_gen/test/unit/test_transformer_loader_fallback.py index b412a3ed9..6b7d8c8db 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_loader_fallback.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_loader_fallback.py @@ -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), } diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py index ee73b768b..47fc32558 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -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