[diffusion] cli: support component attention backend overrides (#24320)

This commit is contained in:
Mick
2026-05-05 08:39:27 +08:00
committed by GitHub
parent 078f84d80d
commit 2f7d99b7f7
10 changed files with 444 additions and 43 deletions
+23
View File
@@ -78,6 +78,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
- `--sp-degree {N}`: sequence parallelism size - `--sp-degree {N}`: sequence parallelism size
- `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls - `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls
- `--attention-backend {BACKEND}`: attention backend for native SGLang pipelines - `--attention-backend {BACKEND}`: attention backend for native SGLang pipelines
- `--component-attention-backends {MAP}`: per-component attention backend overrides, for example `text_encoder=torch_sdpa,transformer=fa`
- `--attention-backend-config {CONFIG}`: attention backend configuration - `--attention-backend-config {CONFIG}`: attention backend configuration
### Sampling and output ### Sampling and output
@@ -195,6 +196,28 @@ sglang serve \
The component key must match the key in the model's `model_index.json`, and the path must be either a Hugging Face repo ID or a complete component directory. The component key must match the key in the model's `model_index.json`, and the path must be either a Hugging Face repo ID or a complete component directory.
## Component Attention Backend Overrides
Use `--component-attention-backends` when one pipeline component needs a different native attention backend from the global `--attention-backend`.
```bash
sglang generate \
--model-path Lightricks/LTX-2.3 \
--attention-backend fa \
--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.
You can also pass dotted CLI entries:
```bash
sglang generate \
--model-path <MODEL_PATH_OR_ID> \
--component-attention-backends.text_encoder torch_sdpa \
--component-attention-backends.transformer fa
```
## Diffusers Backend ## Diffusers Backend
Use `--backend diffusers` to force vanilla diffusers pipelines when no native SGLang implementation exists or when a model requires a custom pipeline class. Use `--backend diffusers` to force vanilla diffusers pipelines when no native SGLang implementation exists or when a model requires a custom pipeline class.
@@ -42,8 +42,9 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa
The selection order in `runtime/layers/attention/selector.py` is: The selection order in `runtime/layers/attention/selector.py` is:
1. `global_force_attn_backend(...)` / `global_force_attn_backend_context_manager(...)` 1. `global_force_attn_backend(...)` / `global_force_attn_backend_context_manager(...)`
2. CLI `--attention-backend` (`ServerArgs.attention_backend`) 2. Component override from `--component-attention-backends` while that component is being constructed
3. Auto selection (platform capability, dtype, and installed packages) 3. CLI `--attention-backend` (`ServerArgs.attention_backend`)
4. Auto selection (platform capability, dtype, and installed packages)
## Configuration ## Configuration
@@ -122,6 +123,20 @@ sglang generate \
--attention-backend torch_sdpa --attention-backend torch_sdpa
``` ```
### Override one component
Use component overrides when a specific module needs different attention semantics from the main transformer:
```bash
sglang generate \
--model-path <MODEL_PATH_OR_ID> \
--prompt "..." \
--attention-backend fa \
--component-attention-backends text_encoder=torch_sdpa
```
Component keys match pipeline module names from `model_index.json`, such as `text_encoder`, `text_encoder_2`, `transformer`, `transformer_2`, or `connectors`.
### Using Sliding Tile Attention (STA) ### Using Sliding Tile Attention (STA)
```bash ```bash
@@ -81,6 +81,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
- `--sp-degree &#123;N&#125;`: sequence parallelism size - `--sp-degree &#123;N&#125;`: sequence parallelism size
- `--ulysses-degree &#123;N&#125;` and `--ring-degree &#123;N&#125;`: USP parallelism controls - `--ulysses-degree &#123;N&#125;` and `--ring-degree &#123;N&#125;`: USP parallelism controls
- `--attention-backend &#123;BACKEND&#125;`: attention backend for native SGLang pipelines - `--attention-backend &#123;BACKEND&#125;`: attention backend for native SGLang pipelines
- `--component-attention-backends &#123;MAP&#125;`: per-component attention backend overrides, for example `text_encoder=torch_sdpa,transformer=fa`
- `--attention-backend-config &#123;CONFIG&#125;`: attention backend configuration - `--attention-backend-config &#123;CONFIG&#125;`: attention backend configuration
### Sampling and output ### Sampling and output
@@ -198,6 +199,28 @@ sglang serve \
The component key must match the key in the model's `model_index.json`, and the path must be either a Hugging Face repo ID or a complete component directory. The component key must match the key in the model's `model_index.json`, and the path must be either a Hugging Face repo ID or a complete component directory.
## Component Attention Backend Overrides
Use `--component-attention-backends` when one pipeline component needs a different native attention backend from the global `--attention-backend`.
```bash Command
sglang generate \
--model-path Lightricks/LTX-2.3 \
--attention-backend fa \
--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.
You can also pass dotted CLI entries:
```bash Command
sglang generate \
--model-path <MODEL_PATH_OR_ID> \
--component-attention-backends.text_encoder torch_sdpa \
--component-attention-backends.transformer fa
```
## Diffusers Backend ## Diffusers Backend
Use `--backend diffusers` to force vanilla diffusers pipelines when no native SGLang implementation exists or when a model requires a custom pipeline class. Use `--backend diffusers` to force vanilla diffusers pipelines when no native SGLang implementation exists or when a model requires a custom pipeline class.
@@ -106,8 +106,9 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa
The selection order in `runtime/layers/attention/selector.py` is: The selection order in `runtime/layers/attention/selector.py` is:
1. `global_force_attn_backend(...)` / `global_force_attn_backend_context_manager(...)` 1. `global_force_attn_backend(...)` / `global_force_attn_backend_context_manager(...)`
2. CLI `--attention-backend` (`ServerArgs.attention_backend`) 2. Component override from `--component-attention-backends` while that component is being constructed
3. Auto selection (platform capability, dtype, and installed packages) 3. CLI `--attention-backend` (`ServerArgs.attention_backend`)
4. Auto selection (platform capability, dtype, and installed packages)
## Configuration ## Configuration
@@ -454,6 +455,20 @@ sglang generate \
--attention-backend torch_sdpa --attention-backend torch_sdpa
``` ```
### Override one component
Use component overrides when a specific module needs different attention semantics from the main transformer:
```bash
sglang generate \
--model-path <MODEL_PATH_OR_ID> \
--prompt "..." \
--attention-backend fa \
--component-attention-backends text_encoder=torch_sdpa
```
Component keys match pipeline module names from `model_index.json`, such as `text_encoder`, `text_encoder_2`, `transformer`, `transformer_2`, or `connectors`.
### Using Sliding Tile Attention (STA) ### Using Sliding Tile Attention (STA)
```bash ```bash
@@ -6,8 +6,9 @@
import os import os
from collections.abc import Generator from collections.abc import Generator
from contextlib import contextmanager from contextlib import contextmanager
from contextvars import ContextVar
from functools import cache from functools import cache
from typing import cast from typing import NamedTuple, cast
import torch import torch
@@ -63,6 +64,16 @@ def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
forced_attn_backend: AttentionBackendEnum | None = None forced_attn_backend: AttentionBackendEnum | None = None
class ComponentAttnBackendContext(NamedTuple):
backend: AttentionBackendEnum | None
component_name: str | None
component_attn_backend_context: ContextVar[ComponentAttnBackendContext | None] = (
ContextVar("component_attn_backend_context", default=None)
)
def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None: def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None:
""" """
Force all attention operations to use a specified backend. Force all attention operations to use a specified backend.
@@ -86,10 +97,25 @@ def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
return forced_attn_backend return forced_attn_backend
def get_component_attn_backend_context() -> ComponentAttnBackendContext | None:
return component_attn_backend_context.get()
def get_component_forced_attn_backend() -> AttentionBackendEnum | None:
context = get_component_attn_backend_context()
return context.backend if context is not None else None
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
def get_attn_backend( def get_attn_backend(
head_size: int, head_size: int,
dtype: torch.dtype, dtype: torch.dtype,
supported_attention_backends: set[AttentionBackendEnum] | None = None, supported_attention_backends: set[AttentionBackendEnum] | None = None,
selected_attention_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]: ) -> type[AttentionBackend]:
if supported_attention_backends is None: if supported_attention_backends is None:
be_tuple = tuple() be_tuple = tuple()
@@ -98,7 +124,41 @@ def get_attn_backend(
be_tuple = tuple( be_tuple = tuple(
sorted(list(supported_attention_backends), key=lambda b: b.name) sorted(list(supported_attention_backends), key=lambda b: b.name)
) )
return _cached_get_attn_backend(head_size, dtype, be_tuple)
selected_backend = selected_attention_backend or get_global_forced_attn_backend()
if selected_backend is None:
selected_backend = get_component_forced_attn_backend()
if selected_backend is None:
server_args = get_global_server_args()
if server_args.attention_backend is not None:
try:
selected_backend = AttentionBackendEnum[
server_args.attention_backend.upper()
]
except KeyError:
raise ValueError(
f"Invalid attention backend '{server_args.attention_backend}' specified via command line. "
f"Available options are: {[e.name.lower() for e in AttentionBackendEnum]}"
)
component_name = get_component_attn_backend_name()
backend_not_specified = selected_backend is None
attention_backend_cls = _cached_get_attn_backend(
head_size,
dtype,
be_tuple,
selected_backend,
)
if component_name:
backend_name = attention_backend_cls.get_enum().name.lower()
if backend_not_specified:
logger.info_once(
f"Attention backend not specified for {component_name}, "
f"using {backend_name} backend for {component_name}"
)
else:
logger.info_once(f"Using {backend_name} backend for {component_name}")
return attention_backend_cls
@cache @cache
@@ -106,32 +166,11 @@ def _cached_get_attn_backend(
head_size: int, head_size: int,
dtype: torch.dtype, dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum], supported_attention_backends: tuple[AttentionBackendEnum],
selected_backend: AttentionBackendEnum | None,
) -> type[AttentionBackend]: ) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced via global_force_attn_backend() or --attention-backend CLI arg.
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
supported_attention_backends = set(supported_attention_backends) supported_attention_backends = set(supported_attention_backends)
selected_backend = None
backend_by_global_setting: AttentionBackendEnum | None = (
get_global_forced_attn_backend()
)
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
else:
# Check the server arguments for a backend override
server_args = get_global_server_args()
if server_args.attention_backend is not None:
try:
selected_backend = AttentionBackendEnum[
server_args.attention_backend.upper()
]
except KeyError:
raise ValueError(
f"Invalid attention backend '{server_args.attention_backend}' specified via command line. "
f"Available options are: {[e.name.lower() for e in AttentionBackendEnum]}"
)
# get device-specific attn_backend # get device-specific attn_backend
if len(supported_attention_backends) == 0: if len(supported_attention_backends) == 0:
@@ -140,14 +179,16 @@ def _cached_get_attn_backend(
elif selected_backend is None and len(supported_attention_backends) == 1: elif selected_backend is None and len(supported_attention_backends) == 1:
selected_backend = next(iter(supported_attention_backends)) selected_backend = next(iter(supported_attention_backends))
elif selected_backend is None: elif selected_backend is None:
logger.debug(f"Attention backend not specified") logger.debug("Attention backend not specified")
elif selected_backend not in supported_attention_backends: elif selected_backend not in supported_attention_backends:
supported_attention_backends_str = [ supported_attention_backends_str = [
supported_attention_backend.__str__() supported_attention_backend.__str__()
for supported_attention_backend in supported_attention_backends for supported_attention_backend in supported_attention_backends
] ]
logger.debug( logger.debug(
f"Selected attention backend: '{selected_backend}' not in supported attention backends: {supported_attention_backends_str}" "Selected attention backend: '%s' not in supported attention backends: %s",
selected_backend,
supported_attention_backends_str,
) )
selected_backend = None selected_backend = None
@@ -161,6 +202,24 @@ def _cached_get_attn_backend(
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls)) return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
@contextmanager
def component_attn_backend_context_manager(
attn_backend: AttentionBackendEnum | None,
component_name: str | None = None,
) -> 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)
)
try:
yield
finally:
component_attn_backend_context.reset(token)
@contextmanager @contextmanager
def global_force_attn_backend_context_manager( def global_force_attn_backend_context_manager(
attn_backend: AttentionBackendEnum, attn_backend: AttentionBackendEnum,
@@ -16,6 +16,10 @@ from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.configs.models import ModelConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.layers.attention.selector import (
component_attn_backend_context_manager,
get_component_attn_backend_context,
)
from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.loader.utils import (
_normalize_component_type, _normalize_component_type,
component_name_to_loader_cls, component_name_to_loader_cls,
@@ -114,7 +118,23 @@ class ComponentLoader(ABC):
component_model_path, component_model_path,
gpu_mem_before_loading, gpu_mem_before_loading,
) )
attn_backend = None
component_attn_name = None
if get_component_attn_backend_context() is None:
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:
logger.info(
"Using %s backend for component: %s",
attn_backend.name.lower(),
matched_backend_key,
)
try: try:
with component_attn_backend_context_manager(
attn_backend, component_name=component_attn_name
):
component = self.load_customized( component = self.load_customized(
component_model_path, server_args, component_name component_model_path, server_args, component_name
) )
@@ -130,6 +150,9 @@ class ComponentLoader(ABC):
f"Error while loading customized {component_name}, falling back to native version" f"Error while loading customized {component_name}, falling back to native version"
) )
# fallback to native version # fallback to native version
with component_attn_backend_context_manager(
attn_backend, component_name=component_attn_name
):
component = self.load_native( component = self.load_native(
component_model_path, server_args, transformers_or_diffusers component_model_path, server_args, transformers_or_diffusers
) )
@@ -18,6 +18,9 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import (
RoleType, RoleType,
filter_modules_for_role, 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 ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader, PipelineComponentLoader,
) )
@@ -411,6 +414,20 @@ class ComposedPipelineBase(ABC):
component_model_path = self._resolve_component_path( component_model_path = self._resolve_component_path(
server_args, module_name, load_module_name server_args, module_name, load_module_name
) )
attn_backend, matched_backend_key = (
server_args.resolve_component_attention_backend(
module_name, load_module_name
)
)
if attn_backend is not None:
logger.info(
"Using %s backend for component: %s",
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( module, memory_usage = PipelineComponentLoader.load_component(
component_name=load_module_name, component_name=load_module_name,
component_model_path=component_model_path, component_model_path=component_model_path,
@@ -179,10 +179,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
self.vae = vae self.vae = vae
self.pipeline = weakref.ref(pipeline) if pipeline else None self.pipeline = weakref.ref(pipeline) if pipeline else None
# TODO(will): hack, should use the actual one in dit selected_attention_backend = self._infer_transformer_attention_backend()
self.attn_backend = get_attn_backend( self.attn_backend = get_attn_backend(
head_size=attn_head_size, head_size=attn_head_size,
dtype=torch.float16, dtype=torch.float16,
selected_attention_backend=selected_attention_backend,
) )
# cfg # cfg
@@ -195,6 +196,26 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
self._cached_num_steps = None self._cached_num_steps = None
self._is_warmed_up = False self._is_warmed_up = False
def _infer_transformer_attention_backend(self) -> AttentionBackendEnum | None:
backends = {
backend
for transformer in (self.transformer, self.transformer_2)
if transformer is not None
for module in transformer.modules()
if isinstance(
(backend := getattr(module, "backend", None)), AttentionBackendEnum
)
}
if not backends:
return None
if len(backends) > 1:
logger.warning(
"Multiple transformer attention backends detected: %s. "
"Using one backend for denoising metadata.",
sorted(backend.name.lower() for backend in backends),
)
return sorted(backends, key=lambda backend: backend.name)[0]
def component_uses( def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None self, server_args: ServerArgs, stage_name: str | None = None
) -> list[ComponentUse]: ) -> list[ComponentUse]:
@@ -131,6 +131,9 @@ class ServerArgs(DisaggArgsMixin):
# Attention # Attention
attention_backend: str = None attention_backend: str = None
attention_backend_config: addict.Dict | None = None attention_backend_config: addict.Dict | None = None
component_attention_backends: dict[str, str] | str | None = field(
default_factory=dict
)
cache_dit_config: str | dict[str, Any] | None = ( cache_dit_config: str | dict[str, Any] | None = (
None # cache-dit config for diffusers None # cache-dit config for diffusers
) )
@@ -470,6 +473,11 @@ class ServerArgs(DisaggArgsMixin):
def _adjust_attention_backend(self): def _adjust_attention_backend(self):
if self.attention_backend in ["fa3", "fa4"]: if self.attention_backend in ["fa3", "fa4"]:
self.attention_backend = "fa" self.attention_backend = "fa"
self.component_attention_backends = (
self._normalize_component_attention_backends(
self.component_attention_backends
)
)
# attention_backend_config # attention_backend_config
if self.attention_backend_config is None: if self.attention_backend_config is None:
@@ -512,6 +520,82 @@ class ServerArgs(DisaggArgsMixin):
return return
self._set_default_attention_backend() self._set_default_attention_backend()
@staticmethod
def _normalize_attention_backend_name(backend: str) -> str:
if not isinstance(backend, str):
raise ValueError("Attention backend name must be a string")
normalized = backend.strip().lower()
if normalized in ("fa3", "fa4"):
normalized = "fa"
try:
return AttentionBackendEnum[normalized.upper()].name.lower()
except KeyError:
raise ValueError(
f"Invalid attention backend '{backend}'. "
f"Available options are: {[e.name.lower() for e in AttentionBackendEnum]}"
) from None
@staticmethod
def _parse_component_attention_backend_map(
value: dict[str, str] | str | None,
) -> dict[str, str]:
if value is None or value == "":
return {}
if isinstance(value, dict):
return dict(value)
if not isinstance(value, str):
raise ValueError(
"component_attention_backends must be a dict or a comma-separated component=backend string"
)
try:
parsed = json.loads(value)
if not isinstance(parsed, dict):
raise ValueError
return parsed
except (json.JSONDecodeError, ValueError):
pass
result: dict[str, str] = {}
for pair in value.split(","):
pair = pair.strip()
if not pair:
continue
if "=" not in pair:
raise ValueError(
"component_attention_backends must use component=backend entries"
)
component, backend = pair.split("=", 1)
result[component.strip()] = backend.strip()
return result
@classmethod
def _normalize_component_attention_backends(
cls, value: dict[str, str] | str | None
) -> dict[str, str]:
raw = cls._parse_component_attention_backend_map(value)
normalized: dict[str, str] = {}
for component, backend in raw.items():
if not isinstance(component, str):
raise ValueError("Component attention backend key must be a string")
component_name = component.strip().replace("-", "_")
if not component_name:
raise ValueError("Component attention backend key must not be empty")
normalized[component_name] = cls._normalize_attention_backend_name(backend)
return normalized
def resolve_component_attention_backend(
self, *component_names: str | None
) -> tuple[AttentionBackendEnum | None, str | None]:
for component_name in component_names:
if component_name is None:
continue
key = component_name.replace("-", "_")
backend = self.component_attention_backends.get(key)
if backend is not None:
return AttentionBackendEnum[backend.upper()], key
return None, None
def _adjust_warmup(self): def _adjust_warmup(self):
if self.warmup_resolutions is not None: if self.warmup_resolutions is not None:
self.warmup = True self.warmup = True
@@ -808,6 +892,16 @@ class ServerArgs(DisaggArgsMixin):
default=None, default=None,
help="Configuration for the attention backend. Can be a JSON string, a path to a JSON/YAML file, or key=value pairs.", help="Configuration for the attention backend. Can be a JSON string, a path to a JSON/YAML file, or key=value pairs.",
) )
parser.add_argument(
"--component-attention-backends",
type=str,
default=None,
help=(
"Per-component attention backend overrides for native pipelines. "
"Use component names from model_index.json, e.g. "
"'text_encoder=torch_sdpa,transformer=fa'."
),
)
parser.add_argument( parser.add_argument(
"--cache-dit-config", "--cache-dit-config",
type=str, type=str,
@@ -1267,6 +1361,43 @@ class ServerArgs(DisaggArgsMixin):
component_paths[component] = path component_paths[component] = path
return component_paths, remaining return component_paths, remaining
@staticmethod
def _extract_component_attention_backends(
unknown_args: list[str],
) -> tuple[dict[str, str], list[str]]:
component_attention_backends: dict[str, str] = {}
remaining: list[str] = []
i = 0
while i < len(unknown_args):
arg = unknown_args[i]
key_part = arg.split("=", 1)[0] if "=" in arg else arg
component = None
if key_part.startswith("--component-attention-backends."):
component = key_part[len("--component-attention-backends.") :].replace(
"-", "_"
)
elif key_part.startswith("--component_attention_backends."):
component = key_part[len("--component_attention_backends.") :].replace(
"-", "_"
)
if component is not None:
if "=" in arg:
component_attention_backends[component] = arg.split("=", 1)[1]
elif i + 1 < len(unknown_args) and not unknown_args[i + 1].startswith(
"-"
):
i += 1
component_attention_backends[component] = unknown_args[i]
else:
remaining.append(arg)
i += 1
continue
else:
remaining.append(arg)
i += 1
return component_attention_backends, remaining
@classmethod @classmethod
def from_cli_args( def from_cli_args(
cls, args: argparse.Namespace, unknown_args: list[str] | None = None cls, args: argparse.Namespace, unknown_args: list[str] | None = None
@@ -1276,6 +1407,9 @@ class ServerArgs(DisaggArgsMixin):
# extract dynamic --<component>-path from unknown args # extract dynamic --<component>-path from unknown args
dynamic_paths, remaining = cls._extract_component_paths(unknown_args) dynamic_paths, remaining = cls._extract_component_paths(unknown_args)
dynamic_attention_backends, remaining = (
cls._extract_component_attention_backends(remaining)
)
if remaining: if remaining:
raise SystemExit(f"error: unrecognized arguments: {' '.join(remaining)}") raise SystemExit(f"error: unrecognized arguments: {' '.join(remaining)}")
@@ -1291,6 +1425,12 @@ class ServerArgs(DisaggArgsMixin):
existing = dict(provided_args.get("component_paths") or {}) existing = dict(provided_args.get("component_paths") or {})
existing.update(dynamic_paths) existing.update(dynamic_paths)
provided_args["component_paths"] = existing provided_args["component_paths"] = existing
if dynamic_attention_backends:
existing = cls._parse_component_attention_backend_map(
provided_args.get("component_attention_backends")
)
existing.update(dynamic_attention_backends)
provided_args["component_attention_backends"] = existing
return cls.from_dict(provided_args) return cls.from_dict(provided_args)
@@ -48,6 +48,71 @@ class TestServerArgsPathExpansion(unittest.TestCase):
args.component_paths["vae"], os.path.expanduser("~/fake/local/vae") args.component_paths["vae"], os.path.expanduser("~/fake/local/vae")
) )
def test_component_attention_backends_are_normalized(self):
args = self._from_dict_without_model_resolution(
{
"model_path": "/data/my-model",
"component_attention_backends": "text-encoder=torch_sdpa,transformer=fa3",
}
)
self.assertEqual(
args.component_attention_backends,
{"text_encoder": "torch_sdpa", "transformer": "fa"},
)
def test_component_attention_backend_lookup(self):
args = self._from_dict_without_model_resolution(
{
"model_path": "/data/my-model",
"component_attention_backends": {"text_encoder": "torch_sdpa"},
}
)
backend, matched_key = args.resolve_component_attention_backend(
"text_encoder", "transformer"
)
self.assertEqual(backend.name, "TORCH_SDPA")
self.assertEqual(matched_key, "text_encoder")
def test_invalid_component_attention_backend_raises(self):
with self.assertRaises(ValueError):
self._from_dict_without_model_resolution(
{
"model_path": "/data/my-model",
"component_attention_backends": {"text_encoder": "bad_backend"},
}
)
with self.assertRaises(ValueError):
self._from_dict_without_model_resolution(
{
"model_path": "/data/my-model",
"component_attention_backends": "text_encoder",
}
)
def test_dynamic_component_attention_backend_cli_args(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
argv = [
"--model-path",
"/fake",
"--component-attention-backends.text-encoder",
"torch_sdpa",
]
with patch.object(sys, "argv", ["sglang"] + argv):
args, unknown_args = parser.parse_known_args(argv)
with patch.object(
PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig()
):
server_args = ServerArgs.from_cli_args(args, unknown_args)
self.assertEqual(
server_args.component_attention_backends, {"text_encoder": "torch_sdpa"}
)
class TestOffloadDefaults(unittest.TestCase): class TestOffloadDefaults(unittest.TestCase):
def _from_dict_with_task_type( def _from_dict_with_task_type(