[diffusion] multi-platform: add Sage Attention 3 Support for sm 120 (RTX5090) (#15382)
Co-authored-by: Mengxi Li <marcyleemx@gmail.com>
This commit is contained in:
@@ -19,6 +19,7 @@ class EncoderArchConfig(ArchConfig):
|
||||
default_factory=lambda: {
|
||||
AttentionBackendEnum.FA,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN_3,
|
||||
}
|
||||
)
|
||||
output_hidden_states: bool = False
|
||||
|
||||
@@ -130,6 +130,7 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
||||
AttentionBackendEnum.AITER,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_3,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -153,6 +153,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
|
||||
AttentionBackendEnum.FA,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_3,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -286,6 +287,7 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
|
||||
AttentionBackendEnum.FA,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_3,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -302,6 +302,7 @@ class QwenImageCrossAttention(nn.Module):
|
||||
AttentionBackendEnum.AITER,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_3,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -739,6 +739,20 @@ class DenoisingStage(PipelineStage):
|
||||
torch.mps.current_allocated_memory(),
|
||||
)
|
||||
|
||||
# In offline local mode (`sglang generate`), offload transformer weights to CPU
|
||||
# after denoising to reduce peak VRAM during VAE decoding.
|
||||
if current_platform.is_cuda_alike() and server_args.is_local_mode:
|
||||
for model in (self.transformer, self.transformer_2):
|
||||
if model is not None:
|
||||
model.to("cpu")
|
||||
logger.info(
|
||||
"Offloaded denoiser transformer weights to CPU after denoising to reduce peak VRAM during VAE decoding."
|
||||
)
|
||||
try:
|
||||
torch.cuda.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs):
|
||||
"""Shard latents for Sequence Parallelism if applicable."""
|
||||
if get_sp_world_size() <= 1:
|
||||
|
||||
@@ -20,7 +20,7 @@ from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.common import is_blackwell
|
||||
from sglang.multimodal_gen.runtime.utils.common import is_blackwell, is_sm120
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import import_pynvml
|
||||
|
||||
@@ -162,7 +162,6 @@ class CudaPlatformBase(Platform):
|
||||
)
|
||||
|
||||
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)
|
||||
@@ -224,6 +223,7 @@ class CudaPlatformBase(Platform):
|
||||
elif selected_backend:
|
||||
raise ValueError(f"Invalid attention backend for {cls.device_name}")
|
||||
else:
|
||||
|
||||
if is_blackwell():
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import (
|
||||
set_fa_ver,
|
||||
@@ -231,6 +231,20 @@ class CudaPlatformBase(Platform):
|
||||
|
||||
set_fa_ver(4)
|
||||
target_backend = AttentionBackendEnum.FA
|
||||
if is_sm120():
|
||||
try:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend,
|
||||
)
|
||||
|
||||
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, Falling back to Torch SDPA (To install it, see https://github.com/thu-ml/SageAttention/tree/main/sageattention3_blackwell#installation)"
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if not cls.has_device_capability(80):
|
||||
logger.info(
|
||||
@@ -243,7 +257,6 @@ class CudaPlatformBase(Platform):
|
||||
"torch.float16 or torch.bfloat16."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
# FlashAttn is valid for the model, checking if the package is
|
||||
# installed.
|
||||
if target_backend == AttentionBackendEnum.FA:
|
||||
|
||||
@@ -269,6 +269,13 @@ def is_blackwell():
|
||||
return torch.cuda.get_device_capability()[0] == 10
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_sm120():
|
||||
if not is_cuda():
|
||||
return False
|
||||
return torch.cuda.get_device_capability()[0] == 12
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_hpu() -> bool:
|
||||
return hasattr(torch, "hpu") and torch.hpu.is_available()
|
||||
|
||||
Reference in New Issue
Block a user