[Diffusion] Fix SANA VAE dtype and TurboWan backend selection (#28769)

This commit is contained in:
Xiaoyu Zhang
2026-06-23 22:57:08 +08:00
committed by GitHub
parent 12b08e620b
commit 31d71c47af
2 changed files with 69 additions and 18 deletions
@@ -52,7 +52,7 @@ class SanaPipelineConfig(SpatialImagePipelineConfig):
# DC-AE does not support tiling or SP VAE decode yet.
vae_tiling: bool = False
vae_sp: bool = False
vae_precision: str = "bf16"
vae_precision: str = "fp32"
dit_config: DiTConfig = field(default_factory=SanaConfig)
vae_config: VAEConfig = field(default_factory=SanaVAEConfig)
@@ -12,21 +12,23 @@ from torch.nn import Module
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
AttentionImpl,
)
from sglang.multimodal_gen.runtime.layers.attention.backends.sparse_linear_attn import (
SageSparseLinearAttentionBackend,
SparseLinearAttentionBackend,
)
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
from sglang.multimodal_gen.runtime.managers.forward_context import (
ForwardContext,
get_forward_context,
)
from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import get_compute_dtype
logger = init_logger(__name__)
_TURBO_WAN_SPARSE_BACKENDS = {
AttentionBackendEnum.SLA_ATTN,
AttentionBackendEnum.SAGE_SLA_ATTN,
}
def post_all2all(local_seq_2_local_head, seq_world_size):
def post_func(input):
@@ -74,6 +76,52 @@ def single_all_to_all(input, local_seq_2_local_head, group, async_op=False):
return res
def _attention_backend_from_name(
backend_name: str | None,
) -> AttentionBackendEnum | None:
if backend_name is None:
return None
try:
return AttentionBackendEnum[backend_name.upper()]
except KeyError:
return None
def _resolve_turbo_wan_sparse_backend(
attention_type: str,
requested_attention_backend: str | None = None,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
) -> tuple[AttentionBackendEnum, str | None]:
available_backends = _TURBO_WAN_SPARSE_BACKENDS
if supported_attention_backends is not None:
available_backends = _TURBO_WAN_SPARSE_BACKENDS & supported_attention_backends
if not available_backends:
available_backends = _TURBO_WAN_SPARSE_BACKENDS
preferred_backend = (
AttentionBackendEnum.SAGE_SLA_ATTN
if attention_type == "sagesla"
else AttentionBackendEnum.SLA_ATTN
)
if preferred_backend not in available_backends:
preferred_backend = sorted(available_backends, key=lambda b: b.name)[0]
requested_backend = _attention_backend_from_name(requested_attention_backend)
if requested_backend in available_backends:
return requested_backend, None
if requested_attention_backend is None:
return preferred_backend, None
return (
preferred_backend,
"TurboWan only supports `sla_attn` or `sage_sla_attn`; "
f"got attention_backend={requested_attention_backend!r}. "
f"Using `{preferred_backend.name.lower()}` from "
f"attention_type={attention_type!r}.",
)
def async_a2a_communicate(
a2a_inputs: Union[torch.Tensor, List[torch.Tensor]],
cp_size: int,
@@ -237,21 +285,24 @@ class MinimalA2AAttnOp(DistributedAttention):
prefix: str = "",
):
dtype = get_compute_dtype()
try:
requested_attention_backend = get_global_server_args().attention_backend
except ValueError:
requested_attention_backend = None
selected_attention_backend, warning_message = _resolve_turbo_wan_sparse_backend(
attention_type,
requested_attention_backend,
supported_attention_backends,
)
if warning_message is not None:
logger.warning_once(warning_message)
attn_backend = get_attn_backend(
head_size, dtype, supported_attention_backends=supported_attention_backends
head_size,
dtype,
supported_attention_backends={selected_attention_backend},
selected_attention_backend=selected_attention_backend,
)
# Maintained for compatibility purposes; can be removed when CI allows setting Attention_backend or when TurboWan supports FA.
if attn_backend not in (
SparseLinearAttentionBackend,
SageSparseLinearAttentionBackend,
):
logger.warning_once(
"TurboWan now only supports `sla_attn` or `sage_sla_attn` and has been automatically set to attention_type. Please set --attention-backend to `sla_attn` or `sage_sla_attn`."
)
if attention_type == "sagesla":
attn_backend = SageSparseLinearAttentionBackend
else:
attn_backend = SparseLinearAttentionBackend
impl_cls: Type[AttentionImpl] = attn_backend.get_impl_cls()
local_attn = impl_cls(
num_heads=num_heads,