diff --git a/python/sglang/multimodal_gen/configs/models/dits/base.py b/python/sglang/multimodal_gen/configs/models/dits/base.py index 77fecfb4d..5f21aebed 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/base.py +++ b/python/sglang/multimodal_gen/configs/models/dits/base.py @@ -31,7 +31,7 @@ class DiTArchConfig(ArchConfig): AttentionBackendEnum.TORCH_SDPA, AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN, - AttentionBackendEnum.SAGE_ATTN_THREE, + AttentionBackendEnum.SAGE_ATTN_3, } ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index d296450f3..3dae916de 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -4,7 +4,7 @@ """ DiffGenerator module for sglang-diffusion. -This module provides a consolidated interface for generating videos using +This module provides a consolidated interface for generating images/videos using diffusion models. """ diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py index 3563ddd18..84d3814c4 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 + import torch from sageattention import sageattn @@ -16,7 +17,6 @@ logger = init_logger(__name__) class SageAttentionBackend(AttentionBackend): - accept_output_buffer: bool = True @staticmethod @@ -31,10 +31,6 @@ class SageAttentionBackend(AttentionBackend): def get_impl_cls() -> type["SageAttentionImpl"]: return SageAttentionImpl - # @staticmethod - # def get_metadata_cls() -> Type["AttentionMetadata"]: - # return FlashAttentionMetadata - class SageAttentionImpl(AttentionImpl): @@ -58,6 +54,8 @@ class SageAttentionImpl(AttentionImpl): key: torch.Tensor, value: torch.Tensor, attn_metadata: AttentionMetadata, + *, + return_softmax_lse: bool = False, ) -> torch.Tensor: output = sageattn( query, @@ -66,5 +64,7 @@ class SageAttentionImpl(AttentionImpl): # since input is (batch_size, seq_len, head_num, head_dim) tensor_layout="NHD", is_causal=self.causal, + sm_scale=self.softmax_scale, + return_lse=return_softmax_lse, ) return output diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn3.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn3.py index fd5b6f2b6..c49d21525 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn3.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn3.py @@ -3,15 +3,12 @@ # SPDX-License-Identifier: Apache-2.0 import torch +from sageattn3 import sageattn3_blackwell from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( AttentionBackend, AttentionImpl, AttentionMetadata, - AttentionMetadataBuilder, -) -from sglang.multimodal_gen.runtime.layers.attention.backends.sageattn.api import ( - sageattn_blackwell, ) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -28,7 +25,7 @@ class SageAttention3Backend(AttentionBackend): @staticmethod def get_name() -> str: - return "SAGE_ATTN_THREE" + return "SAGE_ATTN_3" @staticmethod def get_impl_cls() -> type["SageAttention3Impl"]: @@ -38,14 +35,6 @@ class SageAttention3Backend(AttentionBackend): def get_metadata_cls() -> type["AttentionMetadata"]: raise NotImplementedError - @staticmethod - def get_builder_cls() -> type["AttentionMetadataBuilder"]: - raise NotImplementedError - - # @staticmethod - # def get_metadata_cls() -> Type["AttentionMetadata"]: - # return FlashAttentionMetadata - class SageAttention3Impl(AttentionImpl): @@ -73,6 +62,6 @@ class SageAttention3Impl(AttentionImpl): query = query.transpose(1, 2) key = key.transpose(1, 2) value = value.transpose(1, 2) - output = sageattn_blackwell(query, key, value, is_causal=self.causal) + output = sageattn3_blackwell(query, key, value, is_causal=self.causal) output = output.transpose(1, 2) return output diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py index 9de8eb586..311cb2985 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loader.py @@ -8,6 +8,7 @@ import importlib.util import json import os import time +import traceback from abc import ABC from collections.abc import Generator, Iterable from copy import deepcopy @@ -155,6 +156,10 @@ class ComponentLoader(ABC): ) source = "customized" except Exception as _e: + traceback.print_exc() + logger.error( + f"Error while loading customized {module_name}, falling back to native version" + ) # fallback to native version component = self.load_native( component_model_path, server_args, transformers_or_diffusers diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index 1dd0781a8..6808460bc 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -334,6 +334,7 @@ class QwenImageCrossAttention(nn.Module): supported_attention_backends={ AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA, + AttentionBackendEnum.SAGE_ATTN, }, ) 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 471b9330b..89f206d45 100755 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -134,7 +134,7 @@ class DenoisingStage(PipelineStage): AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA, - AttentionBackendEnum.SAGE_ATTN_THREE, + AttentionBackendEnum.SAGE_ATTN_3, }, # hack ) diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index 90320daed..87029ea2a 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -129,7 +129,7 @@ class CudaPlatformBase(Platform): SlidingTileAttentionBackend, ) - logger.info("Using Sliding Tile Attention backend.") + logger.info("Using Sliding Tile Attention backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend" except ImportError as e: @@ -147,30 +147,27 @@ class CudaPlatformBase(Platform): SageAttentionBackend, ) - logger.info("Using Sage Attention backend.") + logger.info("Using Sage Attention backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn.SageAttentionBackend" except ImportError as e: logger.info(e) logger.info( - "Sage Attention backend is not installed. Fall back to Flash Attention." + "Sage Attention backend is not installed (To install it, run `pip install sageattention==2.2.0 --no-build-isolation`). Falling back to Flash Attention." ) - elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE: + elif selected_backend == AttentionBackendEnum.SAGE_ATTN_3: try: from sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn3 import ( # noqa: F401 SageAttention3Backend, ) - from sglang.multimodal_gen.runtime.layers.attention.backends.sageattn.api import ( # noqa: F401 - sageattn_blackwell, - ) - logger.info("Using Sage Attention 3 backend.") + logger.info("Using Sage Attention 3 backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn3.SageAttention3Backend" except ImportError as e: logger.info(e) logger.info( - "Sage Attention 3 backend is not installed. Fall back to Flash Attention." + "Sage Attention 3 backend is not installed (To install it, see https://github.com/thu-ml/SageAttention/tree/main/sageattention3_blackwell#installation). Falling back to Flash Attention." ) elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN: try: @@ -180,7 +177,7 @@ class CudaPlatformBase(Platform): VideoSparseAttentionBackend, ) - logger.info("Using Video Sparse Attention backend.") + logger.info("Using Video Sparse Attention backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.video_sparse_attn.VideoSparseAttentionBackend" except ImportError as e: @@ -188,7 +185,7 @@ class CudaPlatformBase(Platform): "Failed to import Video Sparse Attention backend: %s", str(e) ) raise ImportError( - "Video Sparse Attention backend is not installed. " + "Video Sparse Attention backend is not installed." ) from e elif selected_backend == AttentionBackendEnum.VMOBA_ATTN: try: @@ -198,7 +195,7 @@ class CudaPlatformBase(Platform): VMOBAAttentionBackend, ) - logger.info("Using Video MOBA Attention backend.") + logger.info("Using Video MOBA Attention backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.vmoba.VMOBAAttentionBackend" except ImportError as e: @@ -209,10 +206,10 @@ class CudaPlatformBase(Platform): "Video MoBA Attention backend is not installed. " ) from e elif selected_backend == AttentionBackendEnum.AITER: - logger.info("Using AITer backend.") + logger.info("Using AITer backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.aiter.AITerBackend" elif selected_backend == AttentionBackendEnum.TORCH_SDPA: - logger.info("Using Torch SDPA backend.") + logger.info("Using Torch SDPA backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend" elif selected_backend in [ AttentionBackendEnum.FA, @@ -272,7 +269,7 @@ class CudaPlatformBase(Platform): target_backend = AttentionBackendEnum.TORCH_SDPA if target_backend == AttentionBackendEnum.TORCH_SDPA: - logger.info("Using Torch SDPA backend.") + logger.info("Using Torch SDPA backend") return "sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend" diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index 660f8e1c4..3be26231e 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -27,7 +27,7 @@ class AttentionBackendEnum(enum.Enum): SLIDING_TILE_ATTN = enum.auto() TORCH_SDPA = enum.auto() SAGE_ATTN = enum.auto() - SAGE_ATTN_THREE = enum.auto() + SAGE_ATTN_3 = enum.auto() VIDEO_SPARSE_ATTN = enum.auto() VMOBA_ATTN = enum.auto() AITER = enum.auto()