diff --git a/docs/docs/sglang-diffusion/attention_backends.mdx b/docs/docs/sglang-diffusion/attention_backends.mdx index 09f0e3423..e402a5b90 100644 --- a/docs/docs/sglang-diffusion/attention_backends.mdx +++ b/docs/docs/sglang-diffusion/attention_backends.mdx @@ -57,7 +57,7 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa `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. + Requires sageattention. On Hopper (SM90), PyPI sageattention==2.2.0 is unsupported because it lacks the upstream SM90 binding fix. Install pip install --force-reinstall git+https://github.com/thu-ml/SageAttention.git@d9704247a5139ab4c03bf7fc6b35cc0e2cbb5ea4 --no-build-isolation. Upstream SageAttention CUDA extensions target SM80/SM86/SM89/SM90/SM120; see upstream setup.py. `sage_attn_3` @@ -401,7 +401,7 @@ Some backends require additional configuration. You can pass these parameters vi Yes ❌ ❌ - Optional dependency on CUDA and MUSA. Falls back to FlashAttention if sageattention is not installed. + Optional dependency on CUDA and MUSA. On Hopper, also falls back to FlashAttention when the installed package lacks the SM90 binding fix. `sage_attn_3` diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index 04f015cdc..9000c22e6 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -151,7 +151,26 @@ class _SageAttentionBackendResolver(_CudaAttentionBackendResolver): def resolve(cls, platform) -> str | AttentionBackendEnum: try: from sageattention import sageattn # noqa: F401 + except ImportError as e: + logger.info(e) + logger.info( + "Sage Attention backend is not installed (To install it, run `pip install git+https://github.com/thu-ml/SageAttention.git@d9704247a5139ab4c03bf7fc6b35cc0e2cbb5ea4 --no-build-isolation`). Falling back to Flash Attention." + ) + return AttentionBackendEnum.FA + if platform.is_hopper(): + try: + # fixed SM90 bindings retain the fake implementation under its own name + from sageattention.sm90_compile import ( # noqa: F401 + qk_int8_sv_f8_accum_f32_fuse_v_scale_attn_inst_buf_fake_impl, + ) + except ImportError: + logger.warning( + "Installed Sage Attention is missing the SM90 binding fix. Falling back to Flash Attention. Reinstall with `pip install --force-reinstall git+https://github.com/thu-ml/SageAttention.git@d9704247a5139ab4c03bf7fc6b35cc0e2cbb5ea4 --no-build-isolation`." + ) + return AttentionBackendEnum.FA + + try: from sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn import ( # noqa: F401 SageAttentionBackend, ) @@ -160,7 +179,7 @@ class _SageAttentionBackendResolver(_CudaAttentionBackendResolver): except ImportError as e: logger.info(e) logger.info( - "Sage Attention backend is not installed (To install it, run `pip install sageattention==2.2.0 --no-build-isolation`). Falling back to Flash Attention." + "Sage Attention backend failed to import. Falling back to Flash Attention." ) return AttentionBackendEnum.FA diff --git a/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py b/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py index 6eea8552a..3109175b9 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py +++ b/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py @@ -1,3 +1,5 @@ +import sys +import types import unittest from unittest.mock import patch @@ -6,7 +8,10 @@ import torch from sglang.multimodal_gen.runtime.layers.attention.selector import ( _cached_get_attn_backend, ) -from sglang.multimodal_gen.runtime.platforms.cuda import CudaPlatformBase +from sglang.multimodal_gen.runtime.platforms.cuda import ( + CudaPlatformBase, + _SageAttentionBackendResolver, +) from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum SDPA_BACKEND_CLS_STR = ( @@ -17,6 +22,7 @@ SDPA_BACKEND_CLS_STR = ( class FakeCudaPlatform(CudaPlatformBase): is_sm120_device = False is_blackwell_device = False + is_hopper_device = False supports_flash_attention = True @classmethod @@ -27,6 +33,10 @@ class FakeCudaPlatform(CudaPlatformBase): def is_blackwell(cls): return cls.is_blackwell_device + @classmethod + def is_hopper(cls): + return cls.is_hopper_device + @classmethod def has_device_capability( cls, @@ -40,6 +50,7 @@ class TestCudaAttentionBackendSelection(unittest.TestCase): def setUp(self): FakeCudaPlatform.is_sm120_device = False FakeCudaPlatform.is_blackwell_device = False + FakeCudaPlatform.is_hopper_device = False FakeCudaPlatform.supports_flash_attention = True _cached_get_attn_backend.cache_clear() @@ -112,6 +123,25 @@ class TestCudaAttentionBackendSelection(unittest.TestCase): with self.assertRaisesRegex(ValueError, "Invalid attention backend"): self.resolve(AttentionBackendEnum.AITER_SAGE) + def test_hopper_sage_attention_without_sm90_fix_falls_back(self): + FakeCudaPlatform.is_hopper_device = True + sageattention = types.ModuleType("sageattention") + sageattention.__path__ = [] + sageattention.sageattn = object() + sm90_compile = types.ModuleType("sageattention.sm90_compile") + + with patch.dict( + sys.modules, + { + "sageattention": sageattention, + "sageattention.sm90_compile": sm90_compile, + }, + ): + self.assertEqual( + _SageAttentionBackendResolver.resolve(FakeCudaPlatform), + AttentionBackendEnum.FA, + ) + def test_explicit_backend_rejected_by_a_model_fails_closed(self): with self.assertRaisesRegex( ValueError, "not supported by this attention layer"