From d269a28b47e28c5f74f3baf3401f9a7d294298c2 Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 16 Aug 2026 10:07:26 +0800 Subject: [PATCH] [diffusion] refactor: route minimax h3 vae attention through native backends (#34949) --- .../runtime/layers/attention/layer.py | 6 +- .../runtime/layers/attention/selector.py | 4 + .../vaes/minimax_h3_audio_vae/audio_vae.py | 47 ++++- .../vaes/minimax_h3_video_vae/attention.py | 108 ++++------ .../vaes/minimax_h3_video_vae/base_module.py | 3 +- .../models/vaes/minimax_h3_video_vae/flash.py | 190 ------------------ .../vaes/minimax_h3_video_vae/vae_vit.py | 30 +-- .../minimax_h3/stages/decoding.py | 4 +- .../test_minimax_h3_vae_parallel_modes.py | 66 ++++++ .../test/unit/test_ring_admission.py | 19 ++ 10 files changed, 187 insertions(+), 290 deletions(-) delete mode 100644 python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/flash.py diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index 649f233fe..a46cc6aeb 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -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 " diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py index 634f22a95..5b861f2dd 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py @@ -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" diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_audio_vae/audio_vae.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_audio_vae/audio_vae.py index 161c56269..8bd78a5a0 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_audio_vae/audio_vae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_audio_vae/audio_vae.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py index 22e880d5d..f496e4ccb 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/attention.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py index b1f92236e..695f7613b 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/flash.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/flash.py deleted file mode 100644 index 3be87d12c..000000000 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/flash.py +++ /dev/null @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/vae_vit.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/vae_vit.py index 802f4ec6a..35daf9015 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/vae_vit.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/vae_vit.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/decoding.py index 424b6033c..130c8fad4 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/decoding.py @@ -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 ) diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py index d25fb8371..4e3a31319 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_vae_parallel_modes.py @@ -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 diff --git a/python/sglang/multimodal_gen/test/unit/test_ring_admission.py b/python/sglang/multimodal_gen/test/unit/test_ring_admission.py index 0978d4654..46eae264e 100644 --- a/python/sglang/multimodal_gen/test/unit/test_ring_admission.py +++ b/python/sglang/multimodal_gen/test/unit/test_ring_admission.py @@ -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()