[diffusion] refactor: route minimax h3 vae attention through native backends (#34949)

This commit is contained in:
Mick
2026-08-16 10:07:26 +08:00
committed by GitHub
parent 4f9da62547
commit d269a28b47
10 changed files with 187 additions and 290 deletions
@@ -681,6 +681,7 @@ class USPAttention(nn.Module):
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
default_attention_backend: AttentionBackendEnum | None = None,
prefix: str = "",
dropout_rate: float = 0.0,
skip_sequence_parallel: bool = False,
@@ -695,6 +696,8 @@ class USPAttention(nn.Module):
text/image encoder outputs), the full USP pipeline is redundant:
each rank's local Q shard can attend directly to the locally-held
full KV without any collective communication.
default_attention_backend:
fallback used only when no global or component override is active.
is_cross_attention:
sparse backend preferences may select a compatible dense backend
for cross-attention while remaining strict for self-attention.
@@ -713,9 +716,10 @@ class USPAttention(nn.Module):
head_size,
dtype,
supported_attention_backends=supported_attention_backends,
default_attention_backend=default_attention_backend,
is_cross_attention=is_cross_attention,
)
if get_ring_parallel_world_size() > 1:
if not skip_sequence_parallel and get_ring_parallel_world_size() > 1:
if not attn_backend.supports_ring_rotation():
raise RuntimeError(
f"Ring Attention requires a backend whose kernel exposes the "
@@ -153,6 +153,7 @@ def get_attn_backend(
supported_attention_backends: set[AttentionBackendEnum] | None = None,
selected_attention_backend: AttentionBackendEnum | None = None,
attention_requirements: AttentionRequirements | None = None,
default_attention_backend: AttentionBackendEnum | None = None,
is_cross_attention: bool = False,
) -> type[AttentionBackend]:
requirements = attention_requirements or AttentionRequirements()
@@ -188,6 +189,9 @@ def get_attn_backend(
server_args, ServerArgs
) and server_args.is_arg_explicitly_set("attention_backend")
if selected_backend is None:
selected_backend = default_attention_backend
allowed_fallback_reason = None
if selected_backend is None:
allowed_fallback_reason = "platform default fallback"
@@ -7,9 +7,14 @@ import numpy as np
import torch
import torch.nn.functional as F
from torch import nn
from torch.nn.functional import scaled_dot_product_attention
from torch.nn.utils.parametrizations import weight_norm
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
)
from .bigvgan import AttrDict, BigVGAN
@@ -52,6 +57,21 @@ class CausalAttention(nn.Module):
self.num_heads = num_heads
self.scale = self.head_dim**-0.5
self.proj = nn.Linear(out_dim, out_dim)
self.attn = (
USPAttention(
num_heads=num_heads,
head_size=self.head_dim,
causal=True,
supported_attention_backends={
AttentionBackendEnum.FA,
AttentionBackendEnum.TORCH_SDPA,
},
default_attention_backend=AttentionBackendEnum.TORCH_SDPA,
skip_sequence_parallel=True,
)
if current_platform.is_cuda()
else None
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, N, C = x.shape
@@ -62,20 +82,33 @@ class CausalAttention(nn.Module):
)
q, k, v = (
qkv.reshape(B, N, 3, self.num_heads, self.head_dim)
.permute(2, 0, 3, 1, 4)
.permute(2, 0, 1, 3, 4)
.unbind(0)
)
x = scaled_dot_product_attention(
q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True
)
if self.attn is None:
x = F.scaled_dot_product_attention(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
attn_mask=None,
dropout_p=0.0,
is_causal=True,
).transpose(1, 2)
else:
input_dtype = q.dtype
if self.attn.backend != AttentionBackendEnum.TORCH_SDPA:
# released audio VAE stays FP32; an explicit fused backend
# owns only the attention compute precision
q, k, v = (tensor.to(self.attn.dtype) for tensor in (q, k, v))
x = self.attn(q, k, v).to(input_dtype)
if self.in_dim > self.out_dim:
x = torch.mean(x, dim=1)
x = torch.mean(x, dim=2)
if self.in_dim // self.num_heads != self.out_dim:
x = nn.functional.adaptive_avg_pool1d(x, self.out_dim)
else:
x = x.transpose(1, 2).reshape(B, N, -1)
x = x.reshape(B, N, -1)
x = self.proj(x)
return x
@@ -1,31 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
# Attention module for the MiniMax H3 visual VAE (inference-only bundle).
from contextlib import nullcontext
from typing import Optional
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from diffusers.utils import logging
from torch.nn.attention import SDPBackend, sdpa_kernel
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.platforms import current_platform
from .flash import flash_attn
from .vit_utils import _env_flag, apply_rotary_pos_emb_qk
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
_FORCE_ROCM_MATH_SDPA = current_platform.is_rocm() and "gfx95" in str(
torch.cuda.get_device_properties(0).gcnArchName
)
def _sdpa_attention(query, key, value):
context = sdpa_kernel([SDPBackend.MATH]) if _FORCE_ROCM_MATH_SDPA else nullcontext()
with context:
return F.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
dropout_p=0.0,
).transpose(1, 2)
def _vit_norm_input(module, hidden_states):
if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_FP32_NORM", "1"):
return hidden_states.float()
weight = getattr(module, "weight", None)
return hidden_states.to(getattr(weight, "dtype", hidden_states.dtype))
weight = module.weight
return hidden_states.to(weight.dtype if weight is not None else hidden_states.dtype)
def _apply_qk_norm(module, hidden_states):
if (
_env_flag("MINIMAX_H3_VAE_DECODER_VIT_FP32_NORM", "1")
and isinstance(module, (nn.LayerNorm, nn.RMSNorm))
and getattr(module, "weight", None) is None
and getattr(module, "bias", None) is None
and module.weight is None
and (not isinstance(module, nn.LayerNorm) or module.bias is None)
and hidden_states.is_cuda
and hidden_states.dtype in (torch.float16, torch.bfloat16)
and not torch.is_grad_enabled()
@@ -83,74 +102,27 @@ class Attention(nn.Module):
)
self.to_qkv = nn.Linear(self.embed_dim, self.attn_inner_dim * 3, bias=bias)
self.to_out = nn.Linear(self.attn_inner_dim, self.embed_dim, bias=out_bias)
# Decode ranks process independent complete tiles. Reuse USPAttention's
# backend dispatch, while deliberately bypassing its sequence collectives.
self.attn = (
USPAttention(
num_heads=heads,
head_size=dim_head,
causal=False,
skip_sequence_parallel=True,
)
if current_platform.is_cuda()
else None
)
if len(kwargs) > 0 and (not dist.is_initialized() or dist.get_rank() == 0):
logger.warning(f"Unused kwargs: {kwargs}")
def _perform_attention(self, query, key, value, pack_info):
cu_seqlens = pack_info.get("cu_seqlens", None)
mask_mod = pack_info.get("mask_mod", None)
block_sparse = pack_info.get("block_sparse", None)
valid_seq_len = pack_info.get("valid_seq_len", None)
if cu_seqlens is not None:
raise NotImplementedError(
"varlen attention is not supported in this inference-only bundle"
)
padded_seq_len = query.shape[1]
if valid_seq_len is not None:
valid_seq_len = int(valid_seq_len)
if not 0 < valid_seq_len <= padded_seq_len:
raise ValueError(
"valid_seq_len must be in (0, padded_seq_len], got "
f"{valid_seq_len} for padded_seq_len={padded_seq_len}"
)
query = query[:, :valid_seq_len]
key = key[:, :valid_seq_len]
value = value[:, :valid_seq_len]
if mask_mod is not None:
hidden_states = flash_attn(
query,
key,
value,
mask_mod=mask_mod,
block_sparse=block_sparse,
)
else:
hidden_states = flash_attn(
query,
key,
value,
)
if valid_seq_len is not None and valid_seq_len < padded_seq_len:
hidden_states = torch.cat(
[
hidden_states,
hidden_states.new_zeros(
hidden_states.shape[0],
padded_seq_len - valid_seq_len,
hidden_states.shape[2],
hidden_states.shape[3],
),
],
dim=1,
)
return hidden_states
def perform_attention(self, query, key, value, pack_info={}):
return self._perform_attention(query, key, value, pack_info)
def forward(
self,
hidden_states: torch.Tensor,
rotary_pos_emb: Optional[torch.Tensor] = None,
pack_info: dict = {},
) -> torch.Tensor:
batch_size, seq_len, _ = hidden_states.shape
@@ -166,7 +138,13 @@ class Attention(nn.Module):
if rotary_pos_emb is not None:
query, key = apply_rotary_pos_emb_qk(query, key, rotary_pos_emb)
hidden_states = self.perform_attention(query, key, value, pack_info)
if self.attn is not None and query.dtype in (torch.float16, torch.bfloat16):
hidden_states = self.attn(query, key, value)
else:
# FlashAttention kernels do not accept FP32. Preserve the explicit
# no-autocast and MPS paths instead of making backend selection
# change H3's supported precision contract.
hidden_states = _sdpa_attention(query, key, value)
hidden_states = hidden_states.reshape(batch_size, seq_len, -1)
hidden_states = self.to_out(hidden_states)
@@ -256,12 +256,11 @@ class TransformerBlock(nn.Module):
self,
hidden_states: torch.FloatTensor,
rotary_pos_emb: Optional[torch.FloatTensor] = None,
pack_info: dict = {},
):
norm_hidden_states = self.norm1(_vit_norm_input(self.norm1, hidden_states)).to(
hidden_states.dtype
)
attn_output = self.attn(norm_hidden_states, rotary_pos_emb, pack_info)
attn_output = self.attn(norm_hidden_states, rotary_pos_emb)
if self.use_scale:
hidden_states = _scaled_residual_add(
hidden_states, attn_output, self.scale1
@@ -1,190 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Torch-native attention implemented with PyTorch SDPA instead of FA4/CUTLASS.
import os
from contextlib import nullcontext
import torch
import torch.nn.functional as F
_BLOCK_CAUSAL_MASK_MOD_CACHE = {}
def _auto_sdpa_backend_name() -> str | None:
"""Return the ROCm-only correctness fallback for H3 video-VAE SDPA."""
if torch.version.hip is None:
return None
from sglang.srt.utils import is_gfx95_supported
# Fused ROCm SDPA corrupts the dense ViT decode on gfx950. Keep every
# non-gfx950 platform, including CUDA, on PyTorch's unchanged auto path.
return "math" if is_gfx95_supported() else None
_AUTO_SDPA_BACKEND = _auto_sdpa_backend_name()
def _as_bool_mask(mask, *, device):
if not isinstance(mask, torch.Tensor):
mask = torch.as_tensor(mask, device=device)
return mask.to(device=device, dtype=torch.bool)
def _ensure_nonempty_rows(mask):
if mask.numel() == 0 or mask.shape[-1] == 0:
return mask
empty = ~mask.any(dim=-1)
mask[..., 0] |= empty
return mask
def _sdpa_kernel_context():
backend_name = os.environ.get("MINIMAX_H3_TORCH_SDPA_BACKEND", "auto").lower()
if backend_name in {"", "auto", "default"}:
backend_name = _AUTO_SDPA_BACKEND
if backend_name is None:
return nullcontext()
from torch.nn.attention import SDPBackend, sdpa_kernel
backends = {
"math": SDPBackend.MATH,
"flash": SDPBackend.FLASH_ATTENTION,
"flash_attention": SDPBackend.FLASH_ATTENTION,
"efficient": SDPBackend.EFFICIENT_ATTENTION,
"mem_efficient": SDPBackend.EFFICIENT_ATTENTION,
"cudnn": SDPBackend.CUDNN_ATTENTION,
"cudnn_attention": SDPBackend.CUDNN_ATTENTION,
}
if backend_name not in backends:
raise ValueError(
"MINIMAX_H3_TORCH_SDPA_BACKEND must be one of "
f"{sorted([*backends, 'auto', 'default'])}, got {backend_name!r}"
)
return sdpa_kernel(backends=[backends[backend_name]])
def _sdpa_attention(query, key, value, causal=False, attn_mask=None):
# query/key/value arrive as [B, S, H, D]; PyTorch SDPA expects
# [B, H, S, D].
q = query.transpose(1, 2)
k = key.transpose(1, 2)
v = value.transpose(1, 2)
if attn_mask is not None and attn_mask.dim() == 3:
attn_mask = attn_mask.unsqueeze(0)
with _sdpa_kernel_context():
out = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=attn_mask,
dropout_p=0.0,
is_causal=causal,
)
return out.transpose(1, 2).nan_to_num(0.0)
def _mask_mod_to_dense(mask_mod, batch, heads, q_len, kv_len, device, aux_tensors=None):
q_idx = torch.arange(q_len, device=device).view(q_len, 1)
kv_idx = torch.arange(kv_len, device=device).view(1, kv_len)
dense = torch.empty((batch, heads, q_len, kv_len), dtype=torch.bool, device=device)
for b in range(batch):
b_idx = torch.tensor(b, device=device)
for h in range(heads):
h_idx = torch.tensor(h, device=device)
mask = mask_mod(b_idx, h_idx, q_idx, kv_idx, None, aux_tensors)
dense[b, h] = _as_bool_mask(mask, device=device)
return _ensure_nonempty_rows(dense)
#########################################################
# Block causal attention
#########################################################
def make_block_causal_mask_mod(num_tokens, block_size, num_special=0, suffix=False):
if num_tokens < 0:
raise ValueError(f"num_tokens must be non-negative, got {num_tokens}")
if block_size <= 0:
raise ValueError(f"block_size must be positive, got {block_size}")
if num_special < 0:
raise ValueError(f"num_special must be non-negative, got {num_special}")
cache_key = (num_tokens, block_size, num_special, suffix)
if cache_key in _BLOCK_CAUSAL_MASK_MOD_CACHE:
return _BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key]
if suffix:
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
del b, h, seqlen_info, aux_tensors
q_is_special = q_idx >= num_tokens
kv_is_special = kv_idx >= num_tokens
return (
q_is_special
| kv_is_special
| (q_idx // block_size >= kv_idx // block_size)
)
else:
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
del b, h, seqlen_info, aux_tensors
q_is_special = q_idx < num_special
kv_is_special = kv_idx < num_special
q_block_idx = (q_idx - num_special) // block_size
kv_block_idx = (kv_idx - num_special) // block_size
return q_is_special | kv_is_special | (q_block_idx >= kv_block_idx)
mask_mod.block_sparse_cache_key = (
"block_causal",
num_tokens,
block_size,
num_special,
suffix,
)
_BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key] = mask_mod
return mask_mod
#########################################################
# Public entry point
#########################################################
@torch.compiler.disable
def flash_attn(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
causal: bool = False,
mask_mod=None,
block_sparse=None,
aux_tensors=None,
) -> torch.Tensor:
use_masked = mask_mod is not None or block_sparse is not None
if block_sparse is not None and mask_mod is None:
raise ValueError("block_sparse requires mask_mod")
if causal and mask_mod is not None:
raise ValueError(
"causal must be encoded in mask_mod when using masked attention"
)
if aux_tensors is not None and not use_masked:
raise ValueError("aux_tensors is only supported with masked attention")
if use_masked:
batch, q_len, heads, _ = query.shape
kv_len = key.shape[1]
dense_mask = _mask_mod_to_dense(
mask_mod,
batch,
heads,
q_len,
kv_len,
query.device,
aux_tensors=aux_tensors,
)
return _sdpa_attention(query, key, value, attn_mask=dense_mask)
return _sdpa_attention(query, key, value, causal=causal)
@@ -8,7 +8,6 @@ from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils import logging
from .base_module import RotaryEmbeddingND, TransformerBlock
from .flash import make_block_causal_mask_mod
from .vit_utils import create_token_ids, prepare_rotary_pos_emb
logger = logging.get_logger(__name__)
@@ -106,12 +105,6 @@ class ViTBase(ModelMixin, ConfigMixin):
self.max_mask_ratio = mask_config.get("max_mask_ratio", 0.75)
self.aspect_ratio_range = mask_config.get("aspect_ratio_range", (0.75, 1.5))
self.max_retries = mask_config.get("max_retries", 100)
if (
self.mask_enabled
and self.mask_style == "drop"
and getattr(self, "t_causal", False)
):
logger.warning("mask_style='drop' with t_causal may cause issues")
if self.mask_enabled and "mask_token" in self._buffers:
del self._buffers["mask_token"]
self.mask_token = nn.Parameter(torch.randn(1, 1, self._mask_dim) * 0.02)
@@ -134,11 +127,9 @@ class ViTBase(ModelMixin, ConfigMixin):
)
return hidden_states, img_ids
def forward_transformer_blocks(self, hidden_states, rotary_pos_emb, pack_info=None):
if pack_info is None:
pack_info = {}
def forward_transformer_blocks(self, hidden_states, rotary_pos_emb):
for block in self.transformer_blocks:
hidden_states = block(hidden_states, rotary_pos_emb, pack_info)
hidden_states = block(hidden_states, rotary_pos_emb)
return hidden_states
def apply_mask_postprocess(self, hidden_states, num_patches):
@@ -179,6 +170,9 @@ class ViT3DDecoder(ViTBase):
):
super().__init__()
if t_causal:
raise ValueError("MiniMax H3's released ViT decoder is non-causal")
dim = heads * dim_head
rope_apply_dim = int(dim_head * rope_dim_ratio)
@@ -190,8 +184,6 @@ class ViT3DDecoder(ViTBase):
self.init_suffix_tokens(dim, num_register_tokens, has_cls_token=False)
self.t_causal = t_causal
self.transformer_blocks = nn.ModuleList(
[
TransformerBlock(
@@ -326,16 +318,6 @@ class ViT3DDecoder(ViTBase):
)
cache_img_ids = img_ids
pack_info = {}
if self.t_causal:
spatial_size = latent_H * latent_W
mask_mod = make_block_causal_mask_mod(
num_tokens=num_patches,
block_size=spatial_size,
suffix=True,
)
pack_info["mask_mod"] = mask_mod
if cache_hit:
rotary_pos_emb = cache_record[2]
else:
@@ -351,7 +333,7 @@ class ViT3DDecoder(ViTBase):
)
for block in self.transformer_blocks:
hidden_states = block(hidden_states, rotary_pos_emb, pack_info)
hidden_states = block(hidden_states, rotary_pos_emb)
hidden_states = self.norm_out(hidden_states)
@@ -11,6 +11,7 @@ from sglang.multimodal_gen.runtime.distributed import (
get_world_group,
model_parallel_is_initialized,
)
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentUse,
)
@@ -361,7 +362,8 @@ class MiniMaxH3DecodingStage(DecodingStage):
server_args,
decode_fn=selected_video_vae.decode_base,
)
visual_frames = video_decode(visual_decode_latent)
with set_forward_context(current_timestep=0, attn_metadata=None):
visual_frames = video_decode(visual_decode_latent)
visual_frames = selected_video_vae.processor.revert_tensor(
visual_frames
)
@@ -4,14 +4,24 @@
from unittest import mock
import pytest
import torch
import torch.nn as nn
from sglang.multimodal_gen.configs.models.vaes.minimax_h3_video import (
MiniMaxH3VideoVAEConfig,
)
from sglang.multimodal_gen.runtime.models.vaes.minimax_h3 import MiniMaxH3VideoVAE
from sglang.multimodal_gen.runtime.models.vaes.minimax_h3_audio_vae.audio_vae import (
CausalAttention,
)
from sglang.multimodal_gen.runtime.models.vaes.minimax_h3_video_vae import (
AutoencoderKLLegacy,
)
from sglang.multimodal_gen.runtime.models.vaes.minimax_h3_video_vae.attention import (
Attention,
_apply_qk_norm,
)
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
def _init_kwargs(config: MiniMaxH3VideoVAEConfig):
@@ -49,3 +59,59 @@ def test_unvalidated_decode_modes_are_rejected(mode):
config = MiniMaxH3VideoVAEConfig(parallel_decode_mode=mode)
with pytest.raises(ValueError, match="use tiled"):
config.resolved_parallel_decode_mode()
def test_vit_attention_uses_local_usp_backend_dispatch():
module = (
"sglang.multimodal_gen.runtime.models.vaes." "minimax_h3_video_vae.attention"
)
with (
mock.patch(f"{module}.current_platform.is_cuda", return_value=True),
mock.patch(f"{module}.USPAttention", autospec=True) as usp_attention,
):
Attention(heads=2, dim_head=64)
assert usp_attention.call_args.kwargs["skip_sequence_parallel"] is True
def test_vit_qk_norm_supports_affine_free_rmsnorm():
norm = nn.RMSNorm(64, elementwise_affine=False)
hidden_states = torch.randn(1, 2, 2, 64)
output = _apply_qk_norm(norm, hidden_states)
assert output.shape == hidden_states.shape
def test_audio_vae_attention_defaults_to_local_sdpa_and_allows_fa():
class RecordingFA(nn.Module):
backend = AttentionBackendEnum.FA
dtype = torch.bfloat16
def forward(self, query, key, value):
self.input_dtype = query.dtype
return query
module = (
"sglang.multimodal_gen.runtime.models.vaes." "minimax_h3_audio_vae.audio_vae"
)
recording_fa = RecordingFA()
with (
mock.patch(f"{module}.current_platform.is_cuda", return_value=True),
mock.patch(
f"{module}.USPAttention", autospec=True, return_value=recording_fa
) as usp_attention,
):
attention = CausalAttention(in_dim=64, out_dim=32, num_heads=2)
output = attention(torch.randn(1, 4, 64))
kwargs = usp_attention.call_args.kwargs
assert kwargs["causal"] is True
assert kwargs["skip_sequence_parallel"] is True
assert kwargs["default_attention_backend"] == AttentionBackendEnum.TORCH_SDPA
assert kwargs["supported_attention_backends"] == {
AttentionBackendEnum.FA,
AttentionBackendEnum.TORCH_SDPA,
}
assert recording_fa.input_dtype == torch.bfloat16
assert output.dtype == torch.float32
@@ -2,6 +2,9 @@
"""Ring admission is a backend capability, not a name whitelist."""
import unittest
from unittest.mock import patch
import torch
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
AttentionBackend,
@@ -10,6 +13,7 @@ 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.layers.attention.layer import USPAttention
from sglang.multimodal_gen.runtime.server_args.server_args import (
RING_CAPABLE_ATTENTION_BACKENDS,
)
@@ -34,6 +38,21 @@ class TestRingAdmission(unittest.TestCase):
SDPABackend.get_enum().name.lower(), RING_CAPABLE_ATTENTION_BACKENDS
)
def test_local_usp_backend_does_not_require_ring_capability(self):
layer_module = "sglang.multimodal_gen.runtime.layers.attention.layer"
with (
patch(f"{layer_module}.get_compute_dtype", return_value=torch.float16),
patch(f"{layer_module}.get_attn_backend", return_value=SDPABackend),
patch(f"{layer_module}.get_ring_parallel_world_size", return_value=2),
):
attention = USPAttention(
num_heads=2,
head_size=64,
skip_sequence_parallel=True,
)
self.assertEqual(attention.backend, SDPABackend.get_enum())
if __name__ == "__main__":
unittest.main()