[diffusion] feat: make ring admission a backend capability (#33928)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-08-07 17:55:44 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 28b43bf693
commit 13938fed3f
7 changed files with 87 additions and 22 deletions
@@ -45,6 +45,12 @@ class AttentionBackend(ABC):
def supports_packed_varlen(cls) -> bool: def supports_packed_varlen(cls) -> bool:
return cls.get_impl_cls().forward_varlen is not AttentionImpl.forward_varlen return cls.get_impl_cls().forward_varlen is not AttentionImpl.forward_varlen
@classmethod
def supports_ring_rotation(cls) -> bool:
"""Whether this backend can serve as the ring-attention kernel; the
per-hop online-softmax merge needs the kernel's softmax LSE."""
return False
@classmethod @classmethod
def unsupported_requirements( def unsupported_requirements(
cls, requirements: AttentionRequirements cls, requirements: AttentionRequirements
@@ -6,6 +6,12 @@ from typing import Any, List, Optional, Tuple
import torch import torch
from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from sglang.multimodal_gen.runtime.layers.utils import register_custom_op from sglang.multimodal_gen.runtime.layers.utils import register_custom_op
from sglang.multimodal_gen.runtime.platforms import ( from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum, AttentionBackendEnum,
@@ -285,13 +291,6 @@ def flash_attn_varlen_func_op_lse(
) )
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
fa_ver = 3 fa_ver = 3
@@ -330,6 +329,11 @@ class FlashAttentionMetadataBuilder(AttentionMetadataBuilder):
class FlashAttentionBackend(AttentionBackend): class FlashAttentionBackend(AttentionBackend):
@classmethod
def supports_ring_rotation(cls) -> bool:
return True
accept_output_buffer: bool = True accept_output_buffer: bool = True
@staticmethod @staticmethod
@@ -33,6 +33,11 @@ def _trailing_padding_used_len(
class SageAttentionBackend(AttentionBackend): class SageAttentionBackend(AttentionBackend):
@classmethod
def supports_ring_rotation(cls) -> bool:
return True
accept_output_buffer: bool = True accept_output_buffer: bool = True
@staticmethod @staticmethod
@@ -689,15 +689,12 @@ class USPAttention(nn.Module):
head_size, dtype, supported_attention_backends=supported_attention_backends head_size, dtype, supported_attention_backends=supported_attention_backends
) )
if get_ring_parallel_world_size() > 1: if get_ring_parallel_world_size() > 1:
backend_enum = attn_backend.get_enum() if not attn_backend.supports_ring_rotation():
if backend_enum not in (
AttentionBackendEnum.FA,
AttentionBackendEnum.SAGE_ATTN,
):
raise RuntimeError( raise RuntimeError(
f"Ring Attention is only supported for FlashAttention or SageAttention backends, " f"Ring Attention requires a backend whose kernel exposes the "
f"but got {backend_enum.name}. " f"softmax LSE for the per-hop merge; "
f"Please ensure your platform supports these backends." f"{attn_backend.get_enum().name} does not declare support "
f"(see AttentionBackend.supports_ring_rotation)."
) )
impl_cls: Type[AttentionImpl] = attn_backend.get_impl_cls() impl_cls: Type[AttentionImpl] = attn_backend.get_impl_cls()
self.allow_cudnn_sdp = bool(extra_impl_args.get("allow_cudnn_sdp", False)) self.allow_cudnn_sdp = bool(extra_impl_args.get("allow_cudnn_sdp", False))
@@ -1578,6 +1578,14 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
use_full_unified_sequence = ( use_full_unified_sequence = (
get_sp_world_size() > 1 and get_ring_parallel_world_size() > 1 get_sp_world_size() > 1 and get_ring_parallel_world_size() > 1
) )
if use_full_unified_sequence:
# Ring support for this attention layout is not implemented; the
# full-sequence gather is correct but gives up ring's memory and
# overlap benefits.
logger.warning_once(
"zimage under ring_degree > 1 falls back to a full-sequence "
"K/V gather"
)
x_local_seq_len = x.shape[1] x_local_seq_len = x.shape[1]
if use_full_unified_sequence: if use_full_unified_sequence:
x = sequence_model_parallel_all_gather(x.contiguous(), dim=1) x = sequence_model_parallel_all_gather(x.contiguous(), dim=1)
@@ -71,6 +71,9 @@ LTX2_TWO_STAGE_PIPELINE_NAMES = ("LTX2TwoStagePipeline", "LTX2TwoStageHQPipeline
# H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident. # H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident.
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130 LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130
LORA_MERGE_MODES = ("auto", "merge", "dynamic") LORA_MERGE_MODES = ("auto", "merge", "dynamic")
# Mirrors AttentionBackend.supports_ring_rotation; the name-level check
# runs before backend classes are importable on every platform.
RING_CAPABLE_ATTENTION_BACKENDS = ("fa", "sage_attn")
def _normalize_ltx2_two_stage_device_mode(mode: str | None) -> str | None: def _normalize_ltx2_two_stage_device_mode(mode: str | None) -> str | None:
@@ -756,18 +759,21 @@ class ServerArgs(DisaggServerArgsMixin):
self.component_attention_backends["text_encoder"] = "torch_sdpa" self.component_attention_backends["text_encoder"] = "torch_sdpa"
if self.ring_degree > 1: if self.ring_degree > 1:
if self.attention_backend is not None and self.attention_backend not in ( if (
"fa", self.attention_backend is not None
"sage_attn", and self.attention_backend not in RING_CAPABLE_ATTENTION_BACKENDS
): ):
raise ValueError( raise ValueError(
"Ring Attention is only supported for flash attention or sage attention backend for now" "Ring Attention requires one of the ring-capable backends "
f"({', '.join(RING_CAPABLE_ATTENTION_BACKENDS)}), got "
f"{self.attention_backend!r}"
) )
if self.attention_backend is None: if self.attention_backend is None:
self.attention_backend = "fa" self.attention_backend = RING_CAPABLE_ATTENTION_BACKENDS[0]
logger.info( logger.info(
"Ring Attention is currently only supported for flash attention or sage attention; " "Ring Attention requires a ring-capable backend; "
"attention_backend has been automatically set to flash attention" "attention_backend has been automatically set to %s",
self.attention_backend,
) )
if self.attention_backend is None and self.backend != Backend.DIFFUSERS: if self.attention_backend is None and self.backend != Backend.DIFFUSERS:
@@ -0,0 +1,39 @@
# SPDX-License-Identifier: Apache-2.0
"""Ring admission is a backend capability, not a name whitelist."""
import unittest
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
AttentionBackend,
)
from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import (
FlashAttentionBackend,
)
from sglang.multimodal_gen.runtime.layers.attention.backends.sdpa import SDPABackend
from sglang.multimodal_gen.runtime.server_args.server_args import (
RING_CAPABLE_ATTENTION_BACKENDS,
)
class TestRingAdmission(unittest.TestCase):
def test_default_is_not_ring_capable(self):
self.assertFalse(AttentionBackend.supports_ring_rotation())
self.assertFalse(SDPABackend.supports_ring_rotation())
def test_lse_backends_declare_support(self):
self.assertTrue(FlashAttentionBackend.supports_ring_rotation())
def test_server_args_names_match_capabilities(self):
# the name-level list gates before backend classes are importable on
# every platform; keep it consistent with the classes it mirrors
self.assertIn(
FlashAttentionBackend.get_enum().name.lower(),
RING_CAPABLE_ATTENTION_BACKENDS,
)
self.assertNotIn(
SDPABackend.get_enum().name.lower(), RING_CAPABLE_ATTENTION_BACKENDS
)
if __name__ == "__main__":
unittest.main()