From 1fb4bf3558c9b92f3b0813ac155b29e6f06bfec9 Mon Sep 17 00:00:00 2001 From: R0CKSTAR Date: Sat, 4 Apr 2026 16:24:02 +0800 Subject: [PATCH] [diffusion] fix: validate attention backend for Ring Attention in USPAttention (#21828) Signed-off-by: Xiaodong Ye --- .../multimodal_gen/runtime/layers/attention/layer.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index fac06a8d6..c98604be7 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -361,6 +361,17 @@ class USPAttention(nn.Module): attn_backend = get_attn_backend( head_size, dtype, supported_attention_backends=supported_attention_backends ) + if get_ring_parallel_world_size() > 1: + backend_enum = attn_backend.get_enum() + if backend_enum not in ( + AttentionBackendEnum.FA, + AttentionBackendEnum.SAGE_ATTN, + ): + raise RuntimeError( + f"Ring Attention is only supported for FlashAttention or SageAttention backends, " + f"but got {backend_enum.name}. " + f"Please ensure your platform supports these backends." + ) impl_cls: Type["AttentionImpl"] = attn_backend.get_impl_cls() self.attn_impl = impl_cls( num_heads=num_heads,