diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index 5cf41faa6..afd8e7439 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -103,6 +103,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis - `--enable-breakable-cuda-graph {true|false}`: capture supported DiT forwards as breakable CUDA graph segments to reduce launch overhead. Requires `--warmup-resolutions` for every served resolution because each resolution is captured separately. An `extra-high` or `high` request is rejected when it would mount request-scoped DiT fusions that were not present during lossless graph capture; VAE-only request-gated paths remain compatible. - `--bcg-text-buckets {N...}`: prompt-length padding buckets for breakable CUDA graph capture/replay reuse. - `--attention-backend {BACKEND}`: attention backend for native SGLang and diffusers pipelines +- `--enable-attention-backend-autotune {true|false}`: for SGLang-native pipelines, benchmark compatible attention backends on each layer's first sufficiently large input and keep a backend only when it is both numerically compatible and measurably faster. Disabled by default and currently validated on SM90 and SM12x. Explicit `--attention-backend`, component overrides, and model-required backends are never replaced. - `--component-attention-backends {MAP}`: per-component attention backend overrides, for example `text_encoder=torch_sdpa,transformer=fa` - `--attention-backend-config {CONFIG}`: attention backend configuration - `--srt-encoder-url {HTTPADDRESS}`: address of SGLang srt server with AR model for GLM-Image like models. See [Models with AR Stage](/docs/sglang-diffusion/models_with_ar). diff --git a/docs/docs/sglang-diffusion/attention_backends.mdx b/docs/docs/sglang-diffusion/attention_backends.mdx index 5e138113a..d9c6e33d6 100644 --- a/docs/docs/sglang-diffusion/attention_backends.mdx +++ b/docs/docs/sglang-diffusion/attention_backends.mdx @@ -168,6 +168,15 @@ Model paths that require one backend for correctness declare it as required; those layers keep that backend even when the surrounding component is overridden. +For an automatic, measurement-based choice, pass +`--enable-attention-backend-autotune true`. On each layer's first sufficiently +large input, the native runtime times compatible candidates and switches only +when the output remains within its numerical guard and the measured gain exceeds +the noise margin. The tuner is disabled by default and currently validated on +SM90 and SM12x. It does not replace a backend selected explicitly through +`--attention-backend` or `--component-attention-backends`, or a backend required +by the model for correctness. + ## Configuration Some backends require additional configuration. You can pass these parameters via `--attention-backend-config`. This argument accepts: diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index d4d623532..7455bcc46 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -18,6 +18,7 @@ from sglang.kernels.ops.diffusion import ( fused_pack_segmented_qkv, fused_scatter_to_padded, ) +from sglang.multimodal_gen.runtime import server_args as server_args_module from sglang.multimodal_gen.runtime.breakable_cuda_graph.replay_token import ( get_current_replay_token, ) @@ -35,6 +36,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_ulysses_parallel_rank, get_ulysses_parallel_world_size, ) +from sglang.multimodal_gen.runtime.layers.attention.autotune import ( + install as install_attention_backend_autotune, +) from sglang.multimodal_gen.runtime.layers.attention.backends import ( flash_attn as _fa_backend, ) @@ -45,7 +49,11 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i from sglang.multimodal_gen.runtime.layers.attention.backends.skip_softmax import ( get_request_skip_softmax_params, ) -from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend +from sglang.multimodal_gen.runtime.layers.attention.selector import ( + get_attn_backend, + get_component_attn_backend_context, + get_global_forced_attn_backend, +) from sglang.multimodal_gen.runtime.layers.attention.turbo_layer import ( async_a2a_communicate, ) @@ -413,7 +421,9 @@ class UlyssesAttention(nn.Module): ) self.attn_impl = impl_cls(**self._attn_impl_ctor_kwargs) wrap_attention_impl_forward(self.attn_impl) - _maybe_install_backend_autotune(self, attn_backend.get_enum()) + _maybe_install_backend_autotune( + self, attn_backend.get_enum(), required_attention_backend + ) self.num_heads = num_heads self.head_size = head_size self.num_kv_heads = num_kv_heads @@ -682,7 +692,9 @@ class LocalAttention(nn.Module): ) self.attn_impl = impl_cls(**self._attn_impl_ctor_kwargs) wrap_attention_impl_forward(self.attn_impl) - _maybe_install_backend_autotune(self, attn_backend.get_enum()) + _maybe_install_backend_autotune( + self, attn_backend.get_enum(), required_attention_backend + ) self.num_heads = num_heads self.head_size = head_size self.num_kv_heads = num_kv_heads @@ -855,7 +867,9 @@ class USPAttention(nn.Module): ) self.attn_impl = impl_cls(**self._attn_impl_ctor_kwargs) wrap_attention_impl_forward(self.attn_impl) - _maybe_install_backend_autotune(self, attn_backend.get_enum()) + _maybe_install_backend_autotune( + self, attn_backend.get_enum(), required_attention_backend + ) self.num_heads = num_heads self.head_size = head_size self.num_kv_heads = num_kv_heads @@ -2106,19 +2120,27 @@ for _attn_cls in ( del _attn_cls -def _maybe_install_backend_autotune(layer, backend) -> None: +def _maybe_install_backend_autotune( + layer, backend, required_attention_backend: AttentionBackendEnum | None +) -> None: """Opt-in: let the layer pick its backend by measurement on its first big call.""" - from sglang.multimodal_gen.runtime.server_args import get_global_server_args - try: - if not get_global_server_args().enable_attention_backend_autotune: + server_args = server_args_module.get_global_server_args() + if not server_args.enable_attention_backend_autotune: return except Exception: # no ServerArgs yet (unit tests, tooling) return - if getattr(layer, "_required_attention_backend", None) is not None: + component_context = get_component_attn_backend_context() + if ( + required_attention_backend is not None + or get_global_forced_attn_backend() is not None + or ( + component_context is not None + and component_context.require_backend_selection + ) + or server_args.is_arg_explicitly_set("attention_backend") + ): return - from sglang.multimodal_gen.runtime.layers.attention.autotune import install - layer.backend = backend layer._default_attn_backend = backend - install(layer) + install_attention_backend_autotune(layer) diff --git a/python/sglang/multimodal_gen/test/unit/test_attention_backend_autotune.py b/python/sglang/multimodal_gen/test/unit/test_attention_backend_autotune.py index 3f29fbb3b..6fe39fd7a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_attention_backend_autotune.py +++ b/python/sglang/multimodal_gen/test/unit/test_attention_backend_autotune.py @@ -9,7 +9,11 @@ from types import SimpleNamespace import pytest import torch -from sglang.multimodal_gen.runtime.layers.attention import autotune +import sglang.multimodal_gen.runtime.server_args as server_args_module +from sglang.multimodal_gen.runtime.layers.attention import ( + autotune, +) +from sglang.multimodal_gen.runtime.layers.attention import layer as attention_layer from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum QUERY = torch.zeros(4, 4) @@ -50,6 +54,23 @@ def stub(monkeypatch): return install +@pytest.fixture +def enabled_autotune(monkeypatch): + monkeypatch.setattr( + server_args_module, + "get_global_server_args", + lambda: SimpleNamespace( + enable_attention_backend_autotune=True, + is_arg_explicitly_set=lambda _name: False, + ), + ) + installed = [] + monkeypatch.setattr( + attention_layer, "install_attention_backend_autotune", installed.append + ) + return installed + + def test_keeps_the_incumbent_without_a_clear_win(stub): incumbent, rival = _Impl(), _Impl() stub( @@ -106,3 +127,73 @@ def test_small_calls_stay_on_the_default_and_leave_the_tuner_armed(monkeypatch): assert called == [], "tuning must wait for a call worth measuring" assert layer.attn_impl is incumbent + + +def test_explicit_backend_is_not_autotuned(monkeypatch, enabled_autotune): + monkeypatch.setattr( + server_args_module, + "get_global_server_args", + lambda: SimpleNamespace( + enable_attention_backend_autotune=True, + is_arg_explicitly_set=lambda name: name == "attention_backend", + ), + ) + attention_layer._maybe_install_backend_autotune( + SimpleNamespace(), + AttentionBackendEnum.TORCH_SDPA, + None, + ) + + assert enabled_autotune == [] + + +def test_required_backend_is_not_autotuned(enabled_autotune): + attention_layer._maybe_install_backend_autotune( + SimpleNamespace(), + AttentionBackendEnum.TORCH_SDPA, + AttentionBackendEnum.TORCH_SDPA, + ) + + assert enabled_autotune == [] + + +def test_globally_forced_backend_is_not_autotuned(monkeypatch, enabled_autotune): + monkeypatch.setattr( + attention_layer, + "get_global_forced_attn_backend", + lambda: AttentionBackendEnum.FA, + ) + attention_layer._maybe_install_backend_autotune( + SimpleNamespace(), + AttentionBackendEnum.FA, + None, + ) + + assert enabled_autotune == [] + + +def test_explicit_component_backend_is_not_autotuned(monkeypatch, enabled_autotune): + monkeypatch.setattr( + attention_layer, + "get_component_attn_backend_context", + lambda: SimpleNamespace(require_backend_selection=True), + ) + attention_layer._maybe_install_backend_autotune( + SimpleNamespace(), + AttentionBackendEnum.TORCH_SDPA, + None, + ) + + assert enabled_autotune == [] + + +def test_automatic_backend_is_autotuned(enabled_autotune): + layer = SimpleNamespace() + + attention_layer._maybe_install_backend_autotune( + layer, + AttentionBackendEnum.TORCH_SDPA, + None, + ) + + assert enabled_autotune == [layer]