From 977096ae03ac4d4148842ad7a9a052c8258bf790 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 2 Feb 2026 09:47:40 +0800 Subject: [PATCH] [diffusion] cli: introduce generic attention backend configuration in ServerArgs (#18036) --- python/pyproject.toml | 3 +- python/pyproject_other.toml | 6 +- .../multimodal_gen/docs/attention_backends.md | 60 ++++++++-- python/sglang/multimodal_gen/docs/cli.md | 1 + .../runtime/entrypoints/utils.py | 5 +- .../attention/backends/sliding_tile_attn.py | 10 +- .../pipelines_core/stages/causal_denoising.py | 2 +- .../pipelines_core/stages/denoising.py | 14 ++- .../multimodal_gen/runtime/server_args.py | 110 +++++++++--------- 9 files changed, 136 insertions(+), 75 deletions(-) diff --git a/python/pyproject.toml b/python/pyproject.toml index 5ed29db44..fd8cdb2d9 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -112,7 +112,8 @@ diffusion = [ "st_attn==0.0.7 ; platform_machine != 'aarch64' and platform_machine != 'arm64'", "vsa==0.0.4 ; platform_machine != 'aarch64' and platform_machine != 'arm64'", "runai_model_streamer>=0.15.5", - "cache-dit==1.2.0" + "cache-dit==1.2.0", + "addict" ] tracing = [ diff --git a/python/pyproject_other.toml b/python/pyproject_other.toml index e8acb6d94..82922ab19 100755 --- a/python/pyproject_other.toml +++ b/python/pyproject_other.toml @@ -96,7 +96,8 @@ diffusion_hip = [ "st_attn==0.0.7", "vsa==0.0.4", "runai_model_streamer>=0.15.5", - "cache-dit==1.1.8" + "cache-dit==1.1.8", + "addict" ] # For Intel Gaudi(device : hpu) follow the installation guide @@ -125,7 +126,8 @@ diffusion_musa = [ "st_attn==0.0.7", "vsa==0.0.4", "runai_model_streamer>=0.15.5", - "cache-dit==1.1.8" + "cache-dit==1.1.8", + "addict" ] test = [ diff --git a/python/sglang/multimodal_gen/docs/attention_backends.md b/python/sglang/multimodal_gen/docs/attention_backends.md index 6e191cfca..6b2f85c07 100644 --- a/python/sglang/multimodal_gen/docs/attention_backends.md +++ b/python/sglang/multimodal_gen/docs/attention_backends.md @@ -23,11 +23,11 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa |---|---|---| | `fa` / `fa3` / `fa4` | `FA` | FlashAttention. `fa3/fa4` are normalized to `fa` during argument parsing (`ServerArgs.__post_init__`). | | `torch_sdpa` | `TORCH_SDPA` | PyTorch `scaled_dot_product_attention`. | -| `sliding_tile_attn` | `SLIDING_TILE_ATTN` | Sliding Tile Attention (STA). Requires `st_attn` and a mask-strategy config file set via the `SGLANG_DIFFUSION_ATTENTION_CONFIG` environment variable. | +| `sliding_tile_attn` | `SLIDING_TILE_ATTN` | Sliding Tile Attention (STA). Requires `st_attn`. Configure via `--attention-backend-config`. | | `sage_attn` | `SAGE_ATTN` | Requires `sageattention`. Upstream SageAttention CUDA extensions target SM80/SM86/SM89/SM90/SM120 (compute capability 8.0/8.6/8.9/9.0/12.0); see upstream `setup.py`: https://github.com/thu-ml/SageAttention/blob/main/setup.py. | | `sage_attn_3` | `SAGE_ATTN_3` | Requires SageAttention3 installed per upstream instructions. | -| `video_sparse_attn` | `VIDEO_SPARSE_ATTN` | Requires `vsa`. | -| `vmoba_attn` | `VMOBA_ATTN` | Requires `kernel.attn.vmoba_attn.vmoba`. | +| `video_sparse_attn` | `VIDEO_SPARSE_ATTN` | Requires `vsa`. Configure `sparsity` via `--attention-backend-config`. | +| `vmoba_attn` | `VMOBA_ATTN` | Requires `kernel.attn.vmoba_attn.vmoba`. Configure via `--attention-backend-config`. | | `aiter` | `AITER` | Requires `aiter`. | ## Selection priority @@ -38,17 +38,59 @@ The selection order in `runtime/layers/attention/selector.py` is: 2. CLI `--attention-backend` (`ServerArgs.attention_backend`) 3. Auto selection (platform capability, dtype, and installed packages) +## Configuration + +Some backends require additional configuration. You can pass these parameters via `--attention-backend-config`. This argument accepts: +- A path to a JSON or YAML configuration file. +- A JSON string (e.g., `'{"sparsity": 0.5}'`). +- Key-value pairs (e.g., `"sparsity=0.5,enable_x=true"`). + +### Supported Configuration Parameters + +#### Sliding Tile Attention (`sliding_tile_attn`) + +| Parameter | Type | Description | Default | +| :--- | :--- | :--- | :--- | +| `mask_strategy_file_path` | `str` | **Required.** Path to the mask strategy JSON file. | - | +| `sta_mode` | `str` | Mode of STA. | `STA_inference` | +| `skip_time_steps` | `int` | Number of steps to use full attention before switching to sparse attention. | `15` | + +#### Video Sparse Attention (`video_sparse_attn`) + +| Parameter | Type | Description | Default | +| :--- | :--- | :--- | :--- | +| `sparsity` | `float` | Validation sparsity (0.0 - 1.0). | `0.0` | + +#### V-MoBA (`vmoba_attn`) + +| Parameter | Type | Description | Default | +| :--- | :--- | :--- | :--- | +| `temporal_chunk_size` | `int` | Chunk size for temporal dimension. | - | +| `temporal_topk` | `int` | Top-K tokens to select in temporal dimension. | - | +| `spatial_chunk_size` | `list[int]` | Chunk size for spatial dimension (H, W). | - | +| `spatial_topk` | `int` | Top-K tokens to select in spatial dimension. | - | +| `st_chunk_size` | `list[int]` | Chunk size for spatiotemporal dimension (T, H, W). | - | +| `st_topk` | `int` | Top-K tokens to select in spatiotemporal dimension. | - | +| `moba_select_mode` | `str` | Selection mode (e.g., `threshold`). | `threshold` | +| `moba_threshold` | `float` | Threshold value for selection. | `0.25` | +| `moba_threshold_type` | `str` | Type of thresholding (e.g., `query_head`). | `query_head` | +| `first_full_step` | `int` | Number of initial steps to use full attention. | `12` | +| `first_full_layer` | `int` | Number of initial layers to use full attention. | `0` | +| `temporal_layer` | `int` | Number of temporal layers. | `1` | +| `spatial_layer` | `int` | Number of spatial layers. | `1` | +| `st_layer` | `int` | Number of spatiotemporal layers. | `1` | + ## Platform support matrix | Backend | CUDA | ROCm | MPS | Notes | |---|---:|---:|---:|---| | `fa` | ✅ | ✅ | ❌ | CUDA requires SM80+ and fp16/bf16. FlashAttention is only used when the required runtime is installed; otherwise it falls back to `torch_sdpa`. | | `torch_sdpa` | ✅ | ✅ | ✅ | Most compatible option across platforms. | -| `sliding_tile_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `st_attn` and `SGLANG_DIFFUSION_ATTENTION_CONFIG`. | +| `sliding_tile_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `st_attn`. Configure via `--attention-backend-config`. | | `sage_attn` | ✅ | ❌ | ❌ | CUDA-only (optional dependency). | | `sage_attn_3` | ✅ | ❌ | ❌ | CUDA-only (optional dependency). | -| `video_sparse_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `vsa`. | -| `vmoba_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `kernel.attn.vmoba_attn.vmoba`. | +| `video_sparse_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `vsa`. Configure `sparsity` via `--attention-backend-config`. | +| `vmoba_attn` | ✅ | ❌ | ❌ | CUDA-only. Requires `kernel.attn.vmoba_attn.vmoba`. Configure via `--attention-backend-config`. | | `aiter` | ✅ | ❌ | ❌ | Requires `aiter`. | ## Usage @@ -72,12 +114,12 @@ sglang generate \ ### Using Sliding Tile Attention (STA) ```bash -export SGLANG_DIFFUSION_ATTENTION_CONFIG=/abs/path/to/mask_strategy.json - +# Pass the mask strategy file path via config sglang generate \ --model-path \ --prompt "..." \ - --attention-backend sliding_tile_attn + --attention-backend sliding_tile_attn \ + --attention-backend-config "mask_strategy_file_path=/abs/path/to/mask_strategy.json" ``` ### Notes for ROCm / MPS diff --git a/python/sglang/multimodal_gen/docs/cli.md b/python/sglang/multimodal_gen/docs/cli.md index 7fa177b42..ae294e76f 100644 --- a/python/sglang/multimodal_gen/docs/cli.md +++ b/python/sglang/multimodal_gen/docs/cli.md @@ -22,6 +22,7 @@ The SGLang-diffusion CLI provides a quick way to access the inference pipeline f - `--ulysses-degree {ULYSSES_DEGREE}`: The degree of DeepSpeed-Ulysses-style SP in USP - `--ring-degree {RING_DEGREE}`: The degree of ring attention-style SP in USP - `--attention-backend {BACKEND}`: Attention backend to use. For SGLang-native pipelines use `fa`, `torch_sdpa`, `sage_attn`, etc. For diffusers pipelines use diffusers backend names like `flash`, `_flash_3_hub`, `sage`, `xformers`. +- `--attention-backend-config {CONFIG}`: Configuration for the attention backend. Can be a JSON string (e.g., '{"k": "v"}'), a path to a JSON/YAML file, or key=value pairs (e.g., "k=v,k2=v2"). - `--cache-dit-config {PATH}`: Path to a Cache-DiT YAML/JSON config (diffusers backend only) - `--dit-precision {DTYPE}`: Precision for the DiT model (currently supports fp32, fp16, and bf16). diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 3815b5577..5edd3f97a 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -211,7 +211,10 @@ def prepare_request( """ Create a Req object with sampling_params as a parameter. """ - req = Req(sampling_params=sampling_params, VSA_sparsity=server_args.VSA_sparsity) + req = Req( + sampling_params=sampling_params, + VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, + ) try: diffusers_kwargs = sampling_params.diffusers_kwargs except AttributeError: diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py index e89392ae3..37a1acf30 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py @@ -8,7 +8,6 @@ from typing import Any import torch from einops import rearrange -import sglang.multimodal_gen.envs as envs from sglang.multimodal_gen.runtime.distributed import get_sp_group from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( AttentionBackend, @@ -21,6 +20,7 @@ from sglang.multimodal_gen.runtime.managers.forward_context import ( get_forward_context, ) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.server_args import get_global_server_args from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.utils import dict_to_3d_list @@ -120,12 +120,14 @@ class SlidingTileAttentionImpl(AttentionImpl): raise ValueError("st attn not supported") # TODO(will-refactor): for now this is the mask strategy, but maybe we should # have a more general config for STA? - config_file = envs.SGLANG_DIFFUSION_ATTENTION_CONFIG - if config_file is None: + mask_strategy_file_path = ( + get_global_server_args().attention_backend_config.mask_strategy_file_path + ) + if mask_strategy_file_path is None: raise ValueError("SGLANG_DIFFUSION_ATTENTION_CONFIG is not set") # TODO(kevin): get mask strategy for different STA modes - with open(config_file) as f: + with open(mask_strategy_file_path) as f: mask_strategy = json.load(f) self.mask_strategy = dict_to_3d_list(mask_strategy) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py index 3c6bde541..f8640186e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py @@ -270,7 +270,7 @@ class CausalDMDDenoisingStage(DenoisingStage): ), # type: ignore patch_size=server_args.pipeline_config.dit_config.patch_size, # type: ignore STA_param=batch.STA_param, # type: ignore - VSA_sparsity=server_args.VSA_sparsity, # type: ignore + VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, # type: ignore device=get_local_torch_device(), # type: ignore ) # type: ignore assert ( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 1880ee6c9..57bc2b02d 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -719,7 +719,7 @@ class DenoisingStage(PipelineStage): if ( not is_warmup and self.attn_backend.get_enum() == AttentionBackendEnum.SLIDING_TILE_ATTN - and server_args.STA_mode == STA_Mode.STA_SEARCHING + and server_args.attention_backend_config.STA_mode == "STA_SEARCHING" ): self.save_sta_search_results(batch) @@ -1178,11 +1178,11 @@ class DenoisingStage(PipelineStage): raw_latent_shape=batch.raw_latent_shape[2:5], patch_size=server_args.pipeline_config.dit_config.patch_size, STA_param=batch.STA_param, - VSA_sparsity=server_args.VSA_sparsity, + VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, device=get_local_torch_device(), ) elif self.attn_backend.get_enum() == AttentionBackendEnum.VMOBA_ATTN: - moba_params = server_args.moba_config.copy() + moba_params = server_args.attention_backend_config.moba_config.copy() moba_params.update( { "current_timestep": i, @@ -1382,8 +1382,12 @@ class DenoisingStage(PipelineStage): Prepare Sliding Tile Attention (STA) parameters and settings. """ # TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280 - STA_mode = server_args.STA_mode - skip_time_steps = server_args.skip_time_steps + try: + STA_mode = STA_Mode[server_args.attention_backend_config.STA_mode] + except Exception as e: + logger.error(f"Passed STA_mode: {STA_mode} doesn't exist") + raise e + skip_time_steps = server_args.attention_backend_config.skip_time_steps if batch.timesteps is None: raise ValueError("Timesteps must be provided") timesteps_num = batch.timesteps.shape[0] diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 122e5d274..4aed80cea 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -3,7 +3,6 @@ # SPDX-License-Identifier: Apache-2.0 # Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py """The arguments of sglang-diffusion Inference.""" - import argparse import dataclasses import inspect @@ -17,8 +16,11 @@ from dataclasses import field from enum import Enum from typing import Any, Optional +import addict +import yaml + from sglang.multimodal_gen import envs -from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig, STA_Mode +from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig from sglang.multimodal_gen.runtime.platforms import ( AttentionBackendEnum, current_platform, @@ -235,6 +237,7 @@ class ServerArgs: # Attention attention_backend: str = None + attention_backend_config: addict.Dict | None = None cache_dit_config: str | dict[str, Any] | None = ( None # cache-dit config for diffusers ) @@ -297,11 +300,6 @@ class ServerArgs: # ComfyUI integration comfyui_mode: bool = False - # STA (Sliding Tile Attention) parameters - mask_strategy_file_path: str | None = None - STA_mode: STA_Mode = STA_Mode.STA_INFERENCE - skip_time_steps: int = 15 - # Compilation enable_torch_compile: bool = False @@ -311,13 +309,6 @@ class ServerArgs: disable_autocast: bool | None = None - # VSA parameters - VSA_sparsity: float = 0.0 # inference/validation sparsity - - # V-MoBA parameters - moba_config_path: str | None = None - moba_config: dict[str, Any] = field(default_factory=dict) - # Master port for distributed inference # TODO: do not hard code master_port: int | None = None @@ -394,6 +385,45 @@ class ServerArgs: if self.vae_cpu_offload is None: self.vae_cpu_offload = True + def _parse_attention_backend_config(self, config_str: str) -> dict[str, Any]: + """parse attention backend config from string.""" + if not config_str: + return {} + + # 1. treat as file path + if os.path.exists(config_str): + if config_str.endswith((".yaml", ".yml")): + with open(config_str, "r") as f: + return yaml.safe_load(f) + elif config_str.endswith(".json"): + with open(config_str, "r") as f: + return json.load(f) + + # 2. treat as JSON string + try: + return json.loads(config_str) + except json.JSONDecodeError: + pass + + # 3. treat as k=v pairs (simple implementation). e.g., "sparsity=0.5,enable_x=true" + try: + config = {} + pairs = config_str.split(",") + for pair in pairs: + k, v = pair.split("=", 1) + k = k.strip() + v = v.strip() + if v.lower() == "true": + v = True + elif v.lower() == "false": + v = False + elif v.replace(".", "", 1).isdigit(): + v = float(v) if "." in v else int(v) + config[k] = v + return config + except Exception: + raise ValueError(f"Could not parse attention backend config: {config_str}") + def __post_init__(self): # configure logger before use configure_logger(server_args=self) @@ -403,6 +433,14 @@ class ServerArgs: if self.attention_backend in ["fa3", "fa4"]: self.attention_backend = "fa" + # normalize attention_backend_config + if self.attention_backend_config is None: + self.attention_backend_config = addict.Dict() + elif isinstance(self.attention_backend_config, str): + self.attention_backend_config = addict.Dict( + self._parse_attention_backend_config(self.attention_backend_config) + ) + # handle warmup if self.warmup_resolutions is not None: self.warmup = True @@ -420,16 +458,6 @@ class ServerArgs: # TODO: remove hard code initial_master_port = (self.master_port or 30005) + random.randint(0, 100) self.master_port = self.settle_port(initial_master_port, 37) - if self.moba_config_path: - try: - with open(self.moba_config_path) as f: - self.moba_config = json.load(f) - logger.info("Loaded V-MoBA config from %s", self.moba_config_path) - except (FileNotFoundError, json.JSONDecodeError) as e: - logger.error( - "Failed to load V-MoBA config from %s: %s", self.moba_config_path, e - ) - raise self.check_server_args() @@ -468,6 +496,12 @@ class ServerArgs: "sage, or xformers." ), ) + parser.add_argument( + "--attention-backend-config", + type=str, + default=None, + 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( "--diffusers-attention-backend", type=str, @@ -569,20 +603,6 @@ class ServerArgs: help="Path to a text file containing prompts (one per line) for batch processing", ) - # STA (Sliding Tile Attention) parameters - parser.add_argument( - "--STA-mode", - type=str, - default=ServerArgs.STA_mode.value, - choices=[mode.value for mode in STA_Mode], - help="STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None", - ) - parser.add_argument( - "--skip-time-steps", - type=int, - default=ServerArgs.skip_time_steps, - help="Number of time steps to warmup (full attention) for STA", - ) parser.add_argument( "--mask-strategy-file-path", type=str, @@ -663,14 +683,6 @@ class ServerArgs: help="Disable autocast for denoising loop and vae decoding in pipeline sampling", ) - # VSA parameters - parser.add_argument( - "--VSA-sparsity", - type=float, - default=ServerArgs.VSA_sparsity, - help="Validation sparsity for VSA", - ) - # Master port for distributed inference parser.add_argument( "--master-port", @@ -1157,9 +1169,3 @@ def get_global_server_args() -> ServerArgs: # TODO(will): may need to handle this for CI. raise ValueError("Global sgl_diffusion args is not set.") return _global_server_args - - -def parse_int_list(value: str) -> list[int]: - if not value: - return [] - return [int(x.strip()) for x in value.split(",")]