From 0447ade32696e6ee84b3671feac43c3d5d070472 Mon Sep 17 00:00:00 2001 From: li_max Date: Fri, 21 Aug 2026 22:56:44 +0800 Subject: [PATCH] [diffusion] fix: fall back to a component's default attention backend (#35796) Co-authored-by: Mick --- docs/docs/sglang-diffusion/api/cli.mdx | 2 +- .../sglang-diffusion/attention_backends.mdx | 15 +- .../runtime/layers/attention/layer.py | 11 +- .../runtime/layers/attention/selector.py | 52 ++++++- .../component_loaders/component_loader.py | 47 +++++-- .../component_loaders/transformer_loader.py | 2 + .../runtime/models/adapter/ltx_2_connector.py | 18 --- .../runtime/models/bridges/mova_dual_tower.py | 2 +- .../runtime/models/dits/cosmos3video.py | 1 + .../runtime/models/dits/helios.py | 1 + .../runtime/models/dits/ltx_2.py | 3 + .../runtime/models/dits/mova_video_dit.py | 1 + .../runtime/models/dits/qwen_image.py | 3 +- .../runtime/models/dits/sana_wm_components.py | 1 + .../vaes/minimax_h3_video_vae/attention.py | 10 +- .../pipelines_core/composed_pipeline_base.py | 22 ++- .../runtime/server_args/server_args.py | 9 +- .../unit/test_attention_backend_selector.py | 133 +++++++++++++++++- .../test_cross_attention_backend_roles.py | 94 +++++++++++++ .../test_minimax_h3_vae_parallel_modes.py | 8 +- .../test/unit/test_qwen3vl_text.py | 90 ++++++++++++ 21 files changed, 462 insertions(+), 63 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/unit/test_cross_attention_backend_roles.py diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index f0808f2f0..b77552986 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -344,7 +344,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. +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. 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 1433feb30..26e2c123a 100644 --- a/docs/docs/sglang-diffusion/attention_backends.mdx +++ b/docs/docs/sglang-diffusion/attention_backends.mdx @@ -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`. -Backend selection is performed by the shared attention layers (e.g. `LocalAttention` / `USPAttention` / `UlyssesAttention` in `sglang.multimodal_gen.runtime.layers.attention.layer`) and therefore applies to any model component using these layers (e.g. diffusion transformer / DiT and encoders). +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. When using the diffusers backend, `--attention-backend` is passed through to diffusers' `set_attention_backend` (e.g., `flash`, `_flash_3_hub`, `sage`, `xformers`, `native`). @@ -129,7 +129,14 @@ The selection order in `runtime/layers/attention/selector.py` is: 1. `global_force_attn_backend(...)` / `global_force_attn_backend_context_manager(...)` 2. Component override from `--component-attention-backends` while that component is being constructed 3. CLI `--attention-backend` (`ServerArgs.attention_backend`) -4. Auto selection (platform capability, dtype, and installed packages) +4. Layer or component default, when declared +5. Auto selection (platform capability, dtype, and installed packages) + +An explicit global backend mismatch fails for DiT self-attention. Auxiliary +components may fall back to their declared default or another compatible backend. +Sparse backends selected for self-attention similarly fall back to a compatible +dense backend for cross-attention. Explicit component overrides are otherwise +strict. ## Configuration @@ -631,6 +638,10 @@ sglang generate \ ``` Component keys match pipeline module names from `model_index.json`, such as `text_encoder`, `text_encoder_2`, `transformer`, `transformer_2`, or `connectors`. +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. ### Per-request override (denoise loop) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index 74f03edf0..b721a5899 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -597,6 +597,8 @@ class LocalAttention(nn.Module): softmax_scale: float | None = None, causal: bool = False, supported_attention_backends: set[AttentionBackendEnum] | None = None, + default_attention_backend: AttentionBackendEnum | None = None, + is_cross_attention: bool = False, compute_dtype: torch.dtype | None = None, **extra_impl_args, ) -> None: @@ -610,7 +612,11 @@ class LocalAttention(nn.Module): dtype = compute_dtype or get_compute_dtype() attn_backend = get_attn_backend( - head_size, dtype, supported_attention_backends=supported_attention_backends + head_size, + dtype, + supported_attention_backends=supported_attention_backends, + default_attention_backend=default_attention_backend, + is_cross_attention=is_cross_attention, ) impl_cls = attn_backend.get_impl_cls() self.allow_cudnn_sdp = bool(extra_impl_args.get("allow_cudnn_sdp", False)) @@ -740,7 +746,8 @@ class USPAttention(nn.Module): each rank's local Q shard can attend directly to the locally-held full KV without any collective communication. default_attention_backend: - fallback used only when no global or component override is active. + preferred fallback when the global backend is incompatible with + this layer. Explicit component overrides otherwise remain strict. is_cross_attention: sparse backend preferences may select a compatible dense backend for cross-attention while remaining strict for self-attention. diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py index 5b861f2dd..dfeb857ab 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py @@ -69,6 +69,7 @@ class ComponentAttnBackendContext(NamedTuple): backend: AttentionBackendEnum | None component_name: str | None selected_backends: dict[str, str | None] + allow_global_backend_fallback: bool = False component_attn_backend_context: ContextVar[ComponentAttnBackendContext | None] = ( @@ -113,6 +114,11 @@ def get_component_attn_backend_name() -> str | None: return context.component_name if context is not None else None +def _component_allows_global_backend_fallback() -> bool: + context = get_component_attn_backend_context() + return context is not None and context.allow_global_backend_fallback + + def _record_component_attn_backend(backend_name: str, reason: str | None) -> bool: context = get_component_attn_backend_context() if context is None or context.component_name is None: @@ -166,6 +172,7 @@ def get_attn_backend( ) selected_backend = selected_attention_backend + selected_from_global_cli = False selection_is_explicit = selected_backend is not None if selected_backend is None: selected_backend = get_global_forced_attn_backend() @@ -188,6 +195,7 @@ def get_attn_backend( selection_is_explicit = isinstance( server_args, ServerArgs ) and server_args.is_arg_explicitly_set("attention_backend") + selected_from_global_cli = selection_is_explicit if selected_backend is None: selected_backend = default_attention_backend @@ -197,6 +205,14 @@ def get_attn_backend( allowed_fallback_reason = "platform default fallback" elif is_cross_attention and selected_backend.is_sparse: allowed_fallback_reason = "dense cross-attention fallback" + elif selected_from_global_cli and ( + default_attention_backend is not None + or _component_allows_global_backend_fallback() + ): + # The global CLI backend is strict for DiT components. Auxiliary + # components may instead use a declared default or platform-compatible + # backend. A component-specific CLI override otherwise remains strict. + allowed_fallback_reason = "global backend fallback" elif not selection_is_explicit: allowed_fallback_reason = "platform default fallback" @@ -206,7 +222,7 @@ def get_attn_backend( candidate_backends = [selected_backend] if allowed_fallback_reason is not None: - for candidate in (None, *be_tuple): + for candidate in (default_attention_backend, None, *be_tuple): if candidate not in candidate_backends: candidate_backends.append(candidate) @@ -229,9 +245,17 @@ def get_attn_backend( selection_error = error continue - candidate_name = candidate_cls.get_enum().name.lower() + candidate_backend = candidate_cls.get_enum() + candidate_name = candidate_backend.name.lower() + if is_cross_attention and candidate_backend.is_sparse: + if selection_error is None: + selection_error = ValueError( + f"Sparse attention backend '{candidate_name}' cannot serve " + "cross-attention" + ) + continue if supported_backends and not _is_backend_supported( - candidate_cls.get_enum(), supported_backends + candidate_backend, supported_backends ): if selection_error is None: selection_error = ValueError( @@ -254,14 +278,22 @@ def get_attn_backend( break if attention_backend_cls is None: + component_name = get_component_attn_backend_name() + component_suffix = ( + f" for component '{component_name}'" if component_name is not None else "" + ) if unsupported_requirements: raise ValueError( f"Attention backend '{unsupported_backend_name}' does not implement " - f"{', '.join(unsupported_requirements)}" + f"{', '.join(unsupported_requirements)}{component_suffix}" ) if selection_error is not None: - raise selection_error - raise ValueError("No compatible attention backend is available") + raise ValueError( + f"{selection_error}{component_suffix}" + ) from selection_error + raise ValueError( + f"No compatible attention backend is available{component_suffix}" + ) backend_name = attention_backend_cls.get_enum().name.lower() reason = fallback_reason @@ -332,13 +364,19 @@ def _is_backend_supported( def component_attn_backend_context_manager( attn_backend: AttentionBackendEnum | None, component_name: str | None = None, + allow_global_backend_fallback: bool = False, ) -> Generator[None, None, None]: if attn_backend is None and component_name is None: yield return token = component_attn_backend_context.set( - ComponentAttnBackendContext(attn_backend, component_name, {}) + ComponentAttnBackendContext( + attn_backend, + component_name, + {}, + allow_global_backend_fallback, + ) ) try: yield 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 f548edf61..6e6ce828b 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 @@ -81,6 +81,11 @@ class ComponentLoader(ABC): # diffusers or transformers expected_library: str = "" + # --attention-backend primarily selects the DiT backend. Auxiliary + # components may fall back when that global choice is incompatible; an + # explicit --component-attention-backends entry remains strict. + allow_global_attention_backend_fallback = True + _loaders_registered = False def __init_subclass__(cls, **kwargs): @@ -125,9 +130,12 @@ class ComponentLoader(ABC): component_name: str, attn_backend: Any, component_attn_name: str | None, + allow_global_backend_fallback: bool, ) -> AutoModel: with component_attn_backend_context_manager( - attn_backend, component_name=component_attn_name + attn_backend, + component_name=component_attn_name, + allow_global_backend_fallback=allow_global_backend_fallback, ): load_kwargs = self.customized_load_kwargs_for_component( server_args, component_name @@ -144,9 +152,12 @@ class ComponentLoader(ABC): transformers_or_diffusers: str, attn_backend: Any, component_attn_name: str | None, + allow_global_backend_fallback: bool, ) -> AutoModel: with component_attn_backend_context_manager( - attn_backend, component_name=component_attn_name + attn_backend, + component_name=component_attn_name, + allow_global_backend_fallback=allow_global_backend_fallback, ): component = self.load_native( component_model_path, @@ -198,6 +209,7 @@ class ComponentLoader(ABC): component_name, attn_backend, component_attn_name, + self.allow_global_attention_backend_fallback, ) source = "sgl-diffusion" except (ComponentCheckpointUnsupportedError, ComponentResidencyError): @@ -231,6 +243,7 @@ class ComponentLoader(ABC): transformers_or_diffusers, attn_backend, component_attn_name, + self.allow_global_attention_backend_fallback, ) source = "native" logger.warning( @@ -501,6 +514,10 @@ class TokenizerLoader(ComponentLoader): class GenericComponentLoader(ComponentLoader): """Generic loader for components that don't have a specific loader.""" + # An unknown out-of-tree component may itself be the primary transformer. + # Require it to opt into fallback through a registered component loader. + allow_global_attention_backend_fallback = False + def __init__( self, library="transformers", component_architecture: str | None = None ) -> None: @@ -521,6 +538,8 @@ class PipelineComponentLoader: transformers_or_diffusers: str, server_args: ServerArgs, component_architecture: str | None = None, + component_attn_backend: Any = None, + component_attn_name: str | None = None, ): """ Load a pipeline component. @@ -538,15 +557,21 @@ class PipelineComponentLoader: ) try: - # Load the component - return loader.load( - component_model_path, - server_args, - component_name, - transformers_or_diffusers, - ) - except Exception as e: + with component_attn_backend_context_manager( + component_attn_backend, + component_name=component_attn_name, + allow_global_backend_fallback=( + loader.allow_global_attention_backend_fallback + ), + ): + return loader.load( + component_model_path, + server_args, + component_name, + transformers_or_diffusers, + ) + except Exception: logger.error( f"Error while loading component: {component_name}, {component_model_path=}" ) - raise e + raise diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 72a2a9653..6f2e18949 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -139,6 +139,8 @@ def _server_args_for_transformer_component( class TransformerLoader(ComponentLoader): """Shared loader for (video/audio) DiT transformers.""" + allow_global_attention_backend_fallback = False + component_names = [ "transformer", "unconditional_transformer", diff --git a/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_connector.py b/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_connector.py index c92a10fd1..f33c3ab06 100644 --- a/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_connector.py +++ b/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_connector.py @@ -11,11 +11,9 @@ from diffusers.models.attention import FeedForward from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import ( LTX2ConnectorConfig, ) -from sglang.multimodal_gen.runtime.layers.attention import USPAttention from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) -from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum def apply_interleaved_rotary_emb( @@ -151,22 +149,6 @@ class LTX2Attention(torch.nn.Module): self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) self.to_out.append(torch.nn.Dropout(dropout)) - # Scaled dot product attention - self.attn = USPAttention( - num_heads=heads, - head_size=self.head_dim, - dropout_rate=0, - softmax_scale=None, - causal=False, - supported_attention_backends={ - AttentionBackendEnum.FA, - AttentionBackendEnum.AITER, - AttentionBackendEnum.TORCH_SDPA, - AttentionBackendEnum.SAGE_ATTN, - AttentionBackendEnum.SAGE_ATTN_3, - }, - ) - def forward( self, hidden_states: torch.Tensor, diff --git a/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py b/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py index 55ff13e99..c94191b29 100644 --- a/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py +++ b/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py @@ -222,7 +222,7 @@ class ConditionalCrossAttention(nn.Module): head_size=self.head_dim, causal=False, softmax_scale=None, - # is_cross_attention=True, + is_cross_attention=True, ) def forward( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py b/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py index 52bfc4d64..c0482ca59 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py @@ -743,6 +743,7 @@ class Cosmos3CrossAttention(nn.Module): causal=False, supported_attention_backends=supported_attention_backends, prefix=add_prefix("attn", prefix), + is_cross_attention=True, ) def forward( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/helios.py b/python/sglang/multimodal_gen/runtime/models/dits/helios.py index 4efd0518a..4be5d334f 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/helios.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/helios.py @@ -372,6 +372,7 @@ class HeliosCrossAttention(nn.Module): head_size=self.head_dim, causal=False, skip_sequence_parallel=True, + is_cross_attention=True, ) def project_kv(self, encoder_hidden_states): diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py index 5850cc0ff..bc40d0901 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py @@ -740,6 +740,7 @@ class LTX2Attention(nn.Module): ) -> None: super().__init__() + is_cross_attention = context_dim is not None self.query_dim = int(query_dim) self.context_dim = int(query_dim if context_dim is None else context_dim) self.heads = int(heads) @@ -836,6 +837,7 @@ class LTX2Attention(nn.Module): softmax_scale=None, causal=False, supported_attention_backends=supported_attention_backends, + is_cross_attention=is_cross_attention, prefix=f"{prefix}.attn", enable_packed_qkv_input_a2a=self.enable_packed_qkv_input_a2a, # official LTX2 torch_sdpa uses cuDNN; cuda setup disables it @@ -850,6 +852,7 @@ class LTX2Attention(nn.Module): softmax_scale=None, causal=False, supported_attention_backends=supported_attention_backends, + is_cross_attention=is_cross_attention, prefix=f"{prefix}.attn", # official LTX2 torch_sdpa uses cuDNN; cuda setup disables it allow_cudnn_sdp=True, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py b/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py index 616f9a447..cc13c42a2 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py @@ -236,6 +236,7 @@ class CrossAttention(nn.Module): head_size=self.head_dim, causal=False, softmax_scale=None, + is_cross_attention=True, ) def forward(self, x: torch.Tensor, y: torch.Tensor): diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index 593e6fa3a..612a2843e 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -641,7 +641,8 @@ class QwenImageCrossAttention(nn.Module): self.norm_added_q = RMSNorm(head_dim, eps=eps) self.norm_added_k = RMSNorm(head_dim, eps=eps) - # Scaled dot product attention + # Despite the historical class name, this is joint text-image + # self-attention: Q/K/V are concatenated before the kernel call. self.attn = USPAttention( num_heads=self.local_num_heads, head_size=self.head_dim, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py index 92526ba0b..8dcfcf987 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py @@ -3009,6 +3009,7 @@ class MultiHeadCrossAttention(nn.Module): self.attn = LocalAttention( num_heads=num_heads, head_size=self.head_dim, + is_cross_attention=True, ) def forward( diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py index f496e4ccb..ec683a0da 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py @@ -11,7 +11,10 @@ from diffusers.utils import logging from torch.nn.attention import SDPBackend, sdpa_kernel from sglang.multimodal_gen.runtime.layers.attention import USPAttention -from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.platforms import ( + AttentionBackendEnum, + current_platform, +) from .vit_utils import _env_flag, apply_rotary_pos_emb_qk @@ -110,6 +113,11 @@ class Attention(nn.Module): num_heads=heads, head_size=dim_head, causal=False, + supported_attention_backends={ + AttentionBackendEnum.FA, + AttentionBackendEnum.TORCH_SDPA, + }, + default_attention_backend=AttentionBackendEnum.TORCH_SDPA, skip_sequence_parallel=True, ) if current_platform.is_cuda() diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py index fd03627ae..9641f5795 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py @@ -18,9 +18,6 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import ( RoleType, filter_modules_for_role, ) -from sglang.multimodal_gen.runtime.layers.attention.selector import ( - component_attn_backend_context_manager, -) from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( PipelineComponentLoader, ) @@ -557,16 +554,15 @@ class ComposedPipelineBase(ABC): attn_backend.name.lower(), matched_backend_key, ) - with component_attn_backend_context_manager( - attn_backend, component_name=matched_backend_key or module_name - ): - module, memory_usage = PipelineComponentLoader.load_component( - component_name=load_module_name, - component_model_path=component_model_path, - transformers_or_diffusers=transformers_or_diffusers, - server_args=server_args, - component_architecture=architecture, - ) + module, memory_usage = PipelineComponentLoader.load_component( + component_name=load_module_name, + component_model_path=component_model_path, + transformers_or_diffusers=transformers_or_diffusers, + server_args=server_args, + component_architecture=architecture, + component_attn_backend=attn_backend, + component_attn_name=matched_backend_key or module_name, + ) self.memory_usages[load_module_name] = memory_usage 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 b6709d27c..451062f19 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -1823,10 +1823,11 @@ class ServerArgs(DisaggServerArgsMixin): type=str, default=None, help=( - "The attention backend to use. For SGLang-native pipelines, use " - "values like fa, torch_sdpa, sage_attn, etc. For diffusers pipelines, " - "use diffusers attention backend names such as flash, _flash_3_hub, " - "sage, or xformers." + "The global attention backend. Native DiT components treat it as " + "strict; auxiliary native components use a compatible fallback when " + "needed. Use --component-attention-backends for a component-scoped " + "choice. For diffusers pipelines, use names such as flash, " + "_flash_3_hub, sage, or xformers." ), ) parser.add_argument( 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 eec9bd285..bf27cc968 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 @@ -9,8 +9,22 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i ) from sglang.multimodal_gen.runtime.layers.attention.selector import ( _cached_get_attn_backend, + component_attn_backend_context_manager, get_attn_backend, + get_component_attn_backend_context, ) +from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( + ComponentLoader, + GenericComponentLoader, + PipelineComponentLoader, +) +from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( + TextEncoderLoader, +) +from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import ( + TransformerLoader, +) +from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import VAELoader from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum from sglang.multimodal_gen.runtime.server_args import ServerArgs @@ -55,6 +69,16 @@ class _FakeAITERBackend: return ("packed varlen attention",) if requirements.packed_varlen else () +class _FakeSparseBackend: + @classmethod + def get_enum(cls) -> AttentionBackendEnum: + return AttentionBackendEnum.LASER_ATTN + + @classmethod + def unsupported_requirements(cls, _requirements) -> tuple[str, ...]: + return () + + class _FakePlatform: device_name = "test" selected_backend = None @@ -64,6 +88,8 @@ class _FakePlatform: cls.selected_backend = selected_backend if selected_backend == AttentionBackendEnum.AITER: return "fake.AITERBackend" + if selected_backend == AttentionBackendEnum.LASER_ATTN: + return "fake.SparseBackend" if selected_backend in (None, AttentionBackendEnum.FA): return "fake.FABackend" return "fake.SDPABackend" @@ -72,6 +98,7 @@ class _FakePlatform: _FAKE_BACKENDS = { "fake.AITERBackend": _FakeAITERBackend, "fake.FABackend": _FakeFABackend, + "fake.SparseBackend": _FakeSparseBackend, "fake.SDPABackend": _FakeSDPABackend, } @@ -89,13 +116,19 @@ class TestAttentionBackendFallback(unittest.TestCase): is_cross_attention: bool, supported: set[AttentionBackendEnum], attention_requirements: AttentionRequirements | None = None, + default_attention_backend: AttentionBackendEnum | None = None, + component_backend: AttentionBackendEnum | None = None, + allow_global_backend_fallback: bool = False, server_args: object | None = None, ): if server_args is None: server_args = _ServerArgs(backend.name.lower(), explicit=explicit) with ( patch(f"{_SELECTOR}.get_global_forced_attn_backend", return_value=None), - patch(f"{_SELECTOR}.get_component_forced_attn_backend", return_value=None), + patch( + f"{_SELECTOR}.get_component_forced_attn_backend", + return_value=component_backend, + ), patch(f"{_SELECTOR}.get_global_server_args", return_value=server_args), patch( "sglang.multimodal_gen.runtime.platforms.current_platform", @@ -105,12 +138,18 @@ class TestAttentionBackendFallback(unittest.TestCase): f"{_SELECTOR}.resolve_obj_by_qualname", side_effect=_FAKE_BACKENDS.__getitem__, ), + component_attn_backend_context_manager( + component_backend, + component_name="text_encoder", + allow_global_backend_fallback=allow_global_backend_fallback, + ), ): return get_attn_backend( 128, torch.bfloat16, supported_attention_backends=supported, attention_requirements=attention_requirements, + default_attention_backend=default_attention_backend, is_cross_attention=is_cross_attention, ) @@ -161,6 +200,45 @@ class TestAttentionBackendFallback(unittest.TestCase): supported={AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA}, ) + def test_explicit_global_backend_uses_component_default(self): + backend = self._resolve( + AttentionBackendEnum.AITER, + explicit=True, + is_cross_attention=False, + supported={AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA}, + default_attention_backend=AttentionBackendEnum.TORCH_SDPA, + ) + + self.assertIs(backend, _FakeSDPABackend) + self.assertEqual( + _FakePlatform.selected_backend, AttentionBackendEnum.TORCH_SDPA + ) + + def test_explicit_global_backend_falls_back_for_auxiliary_component(self): + backend = self._resolve( + AttentionBackendEnum.AITER, + explicit=True, + is_cross_attention=False, + supported={AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA}, + allow_global_backend_fallback=True, + ) + + self.assertIs(backend, _FakeFABackend) + self.assertIsNone(_FakePlatform.selected_backend) + + def test_explicit_component_backend_remains_strict(self): + with self.assertRaisesRegex( + ValueError, "not supported by this attention layer" + ): + self._resolve( + AttentionBackendEnum.FA, + explicit=True, + is_cross_attention=False, + supported={AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA}, + component_backend=AttentionBackendEnum.AITER, + allow_global_backend_fallback=True, + ) + def test_sparse_backend_falls_back_for_cross_attention(self): backend = self._resolve( AttentionBackendEnum.LASER_ATTN, @@ -172,6 +250,17 @@ class TestAttentionBackendFallback(unittest.TestCase): self.assertIs(backend, _FakeFABackend) self.assertIsNone(_FakePlatform.selected_backend) + def test_sparse_backend_falls_back_for_unconstrained_cross_attention(self): + backend = self._resolve( + AttentionBackendEnum.LASER_ATTN, + explicit=True, + is_cross_attention=True, + supported=set(), + ) + + self.assertIs(backend, _FakeFABackend) + self.assertIsNone(_FakePlatform.selected_backend) + def test_sparse_backend_mismatch_fails_for_self_attention(self): with self.assertRaisesRegex( ValueError, "not supported by this attention layer" @@ -184,5 +273,47 @@ class TestAttentionBackendFallback(unittest.TestCase): ) +class TestComponentAttentionBackendScope(unittest.TestCase): + def _load_with_policy(self, allow_global_backend_fallback: bool): + captured_context = None + + class _Loader: + def load(self, *_args): + nonlocal captured_context + captured_context = get_component_attn_backend_context() + return object(), 0.0 + + _Loader.allow_global_attention_backend_fallback = allow_global_backend_fallback + with patch.object( + ComponentLoader, "for_component_type", return_value=_Loader() + ): + PipelineComponentLoader.load_component( + component_name="text_encoder", + component_model_path="unused", + transformers_or_diffusers="transformers", + server_args=object(), + component_attn_name="text_encoder", + ) + return captured_context + + def test_auxiliary_loader_enables_global_fallback(self): + context = self._load_with_policy(True) + + self.assertIsNotNone(context) + self.assertTrue(context.allow_global_backend_fallback) + + def test_dit_loader_keeps_global_backend_strict(self): + context = self._load_with_policy(False) + + self.assertIsNotNone(context) + self.assertFalse(context.allow_global_backend_fallback) + + def test_builtin_loader_scopes(self): + self.assertFalse(TransformerLoader.allow_global_attention_backend_fallback) + self.assertFalse(GenericComponentLoader.allow_global_attention_backend_fallback) + self.assertTrue(TextEncoderLoader.allow_global_attention_backend_fallback) + self.assertTrue(VAELoader.allow_global_attention_backend_fallback) + + if __name__ == "__main__": unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_cross_attention_backend_roles.py b/python/sglang/multimodal_gen/test/unit/test_cross_attention_backend_roles.py new file mode 100644 index 000000000..84d4e0243 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_cross_attention_backend_roles.py @@ -0,0 +1,94 @@ +from unittest import mock + +import pytest +import torch +from torch import nn + +from sglang.multimodal_gen.runtime.layers.attention import layer as attention_layer +from sglang.multimodal_gen.runtime.models.bridges import mova_dual_tower +from sglang.multimodal_gen.runtime.models.dits import ltx_2 +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum + + +class _FakeAttentionImpl(nn.Module): + def __init__(self, **_kwargs) -> None: + super().__init__() + + +class _FakeAttentionBackend: + @classmethod + def get_enum(cls) -> AttentionBackendEnum: + return AttentionBackendEnum.FA + + @classmethod + def get_impl_cls(cls): + return _FakeAttentionImpl + + +def test_local_attention_forwards_cross_attention_role(): + with ( + mock.patch.object( + attention_layer, "get_compute_dtype", return_value=torch.bfloat16 + ), + mock.patch.object( + attention_layer, "get_attn_backend", return_value=_FakeAttentionBackend + ) as get_backend, + mock.patch.object(attention_layer, "wrap_attention_impl_forward"), + ): + attention_layer.LocalAttention( + num_heads=1, + head_size=64, + is_cross_attention=True, + ) + + assert get_backend.call_args.kwargs["is_cross_attention"] is True + + +@pytest.mark.parametrize("use_local_attention", [False, True]) +def test_ltx2_derives_cross_attention_role_from_context(use_local_attention): + selected_layer = "LocalAttention" if use_local_attention else "USPAttention" + with ( + mock.patch.object(ltx_2, "get_tp_world_size", return_value=1), + mock.patch.object(ltx_2, "ColumnParallelLinear", return_value=nn.Identity()), + mock.patch.object(ltx_2, "RowParallelLinear", return_value=nn.Identity()), + mock.patch.object(ltx_2, selected_layer) as attention, + ): + ltx_2.LTX2Attention( + query_dim=8, + context_dim=8, + heads=1, + dim_head=8, + use_local_attention=use_local_attention, + ) + cross_attention_kwargs = attention.call_args.kwargs + attention.reset_mock() + ltx_2.LTX2Attention( + query_dim=8, + heads=1, + dim_head=8, + use_local_attention=use_local_attention, + ) + self_attention_kwargs = attention.call_args.kwargs + + assert cross_attention_kwargs["is_cross_attention"] is True + assert self_attention_kwargs["is_cross_attention"] is False + + +def test_mova_bridge_marks_conditional_attention_as_cross_attention(): + with ( + mock.patch.object(mova_dual_tower, "get_tp_world_size", return_value=1), + mock.patch.object( + mova_dual_tower, "ColumnParallelLinear", return_value=nn.Identity() + ), + mock.patch.object( + mova_dual_tower, "RowParallelLinear", return_value=nn.Identity() + ), + mock.patch.object(mova_dual_tower, "USPAttention") as attention, + ): + mova_dual_tower.ConditionalCrossAttention( + dim=8, + kv_dim=8, + num_heads=1, + ) + + assert attention.call_args.kwargs["is_cross_attention"] is True diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py index 4e3a31319..3f54206de 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py @@ -71,7 +71,13 @@ def test_vit_attention_uses_local_usp_backend_dispatch(): ): Attention(heads=2, dim_head=64) - assert usp_attention.call_args.kwargs["skip_sequence_parallel"] is True + kwargs = usp_attention.call_args.kwargs + assert kwargs["skip_sequence_parallel"] is True + assert kwargs["default_attention_backend"] == AttentionBackendEnum.TORCH_SDPA + assert kwargs["supported_attention_backends"] == { + AttentionBackendEnum.FA, + AttentionBackendEnum.TORCH_SDPA, + } def test_vit_qk_norm_supports_affine_free_rmsnorm(): diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py index 29bf13e1f..d7cf6fb6d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py @@ -4,6 +4,15 @@ import torch from torch import nn import sglang.multimodal_gen.runtime.models.encoders.qwen3vl as qwen3vl +from sglang.multimodal_gen.runtime.layers.attention.selector import ( + _cached_get_attn_backend, + component_attn_backend_context_manager, +) +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.server_args import ServerArgs + +_SELECTOR = "sglang.multimodal_gen.runtime.layers.attention.selector" +_LAYER = "sglang.multimodal_gen.runtime.layers.attention.layer" class _IdentityAttention(nn.Module): @@ -11,6 +20,41 @@ class _IdentityAttention(nn.Module): return query +class _ExplicitServerArgs(ServerArgs): + def __init__(self) -> None: + self.attention_backend = "aiter" + self._explicit_arg_names = {"attention_backend"} + + +class _FakeAttentionImpl(nn.Module): + def __init__(self, **_kwargs) -> None: + super().__init__() + + +class _FakeFABackend: + @classmethod + def get_enum(cls) -> AttentionBackendEnum: + return AttentionBackendEnum.FA + + @classmethod + def get_impl_cls(cls): + return _FakeAttentionImpl + + @classmethod + def unsupported_requirements(cls, _requirements) -> tuple[str, ...]: + return () + + +class _FakePlatform: + device_name = "test" + + @staticmethod + def get_attn_backend_cls_str(selected_backend, _head_size, _dtype): + if selected_backend not in (None, AttentionBackendEnum.FA): + raise AssertionError(f"Unexpected backend: {selected_backend}") + return "fake.FABackend" + + def test_qwen3vl_attention_uses_interleaved_mrope(monkeypatch): captured_kwargs = {} @@ -44,6 +88,52 @@ def test_qwen3vl_attention_uses_interleaved_mrope(monkeypatch): assert captured_kwargs == {"mrope_interleaved": True} +def test_qwen3vl_auxiliary_component_falls_back_from_global_backend(monkeypatch): + monkeypatch.setattr( + qwen3vl, "_make_text_linear", lambda *args, **kwargs: nn.Identity() + ) + monkeypatch.setattr( + qwen3vl, "_make_text_row_linear", lambda *args, **kwargs: nn.Identity() + ) + monkeypatch.setattr( + qwen3vl, "_make_text_rms_norm", lambda *args, **kwargs: nn.Identity() + ) + monkeypatch.setattr( + qwen3vl, "build_qwen_vl_text_rope", lambda *args, **kwargs: object() + ) + monkeypatch.setattr(f"{_LAYER}.get_compute_dtype", lambda: torch.bfloat16) + monkeypatch.setattr(f"{_LAYER}.wrap_attention_impl_forward", lambda _impl: None) + monkeypatch.setattr(f"{_SELECTOR}.get_global_forced_attn_backend", lambda: None) + monkeypatch.setattr( + f"{_SELECTOR}.get_global_server_args", lambda: _ExplicitServerArgs() + ) + monkeypatch.setattr( + "sglang.multimodal_gen.runtime.platforms.current_platform", _FakePlatform + ) + monkeypatch.setattr( + f"{_SELECTOR}.resolve_obj_by_qualname", lambda _name: _FakeFABackend + ) + _cached_get_attn_backend.cache_clear() + config = SimpleNamespace( + head_dim=8, + hidden_size=8, + num_attention_heads=1, + num_key_value_heads=1, + attention_dropout=0.0, + attention_bias=False, + rms_norm_eps=1e-6, + ) + + with component_attn_backend_context_manager( + None, + component_name="text_encoder", + allow_global_backend_fallback=True, + ): + attention = qwen3vl.Qwen3VLTextAttention(config, layer_idx=0) + + assert attention.attn.backend == AttentionBackendEnum.FA + + def test_qwen3vl_attention_passes_three_axis_positions_to_srt_rope(monkeypatch): attention = qwen3vl.Qwen3VLTextAttention.__new__(qwen3vl.Qwen3VLTextAttention) nn.Module.__init__(attention)