[diffusion] fix: fall back to a component's default attention backend (#35796)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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():
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user