diff --git a/docs/docs/sglang-diffusion/compatibility_matrix.mdx b/docs/docs/sglang-diffusion/compatibility_matrix.mdx index 9cf49963a..8d7ef4c2d 100644 --- a/docs/docs/sglang-diffusion/compatibility_matrix.mdx +++ b/docs/docs/sglang-diffusion/compatibility_matrix.mdx @@ -79,6 +79,12 @@ Rows are grouped when a family shares the same runtime path or optimization supp 480p / 720p VSA + + SANA-Video +
Efficient-Large-Model/SANA-Video_2B_480p_diffusers
+ T2V, 480p + No dedicated optimization listed + Wan2.2
Wan-AI/Wan2.2-TI2V-5B-DiffusersWan-AI/Wan2.2-T2V-A14B-Diffusersnvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4Wan-AI/Wan2.2-I2V-A14B-Diffusers
@@ -254,6 +260,21 @@ Optimization columns are abbreviated to keep the matrix readable: ❌ ❌ + + SANA-Video 2B + Efficient-Large-Model/SANA-Video_2B_480p_diffusers + 480p + ❌ + ❌ + ❌ + ❌ + ❌ + ❌ + ❌ + ❌ + ❌ + ❌ + FastWan2.2 TI2V 5B FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers
FastVideo/FastWan2.2-TI2V-5B-Diffusers diff --git a/python/sglang/multimodal_gen/configs/models/dits/__init__.py b/python/sglang/multimodal_gen/configs/models/dits/__init__.py index 10b7b5d8e..7b4d0fe3f 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/__init__.py +++ b/python/sglang/multimodal_gen/configs/models/dits/__init__.py @@ -18,6 +18,7 @@ from sglang.multimodal_gen.configs.models.dits.longlive2 import LongLive2VideoCo from sglang.multimodal_gen.configs.models.dits.minimax_h3 import MiniMaxH3DiTConfig from sglang.multimodal_gen.configs.models.dits.mova_audio import MOVAAudioConfig from sglang.multimodal_gen.configs.models.dits.mova_video import MOVAVideoConfig +from sglang.multimodal_gen.configs.models.dits.sana_video import SanaVideoConfig from sglang.multimodal_gen.configs.models.dits.stablediffusion3 import ( StableDiffusion3TransformerConfig, ) @@ -37,5 +38,6 @@ __all__ = [ "Hunyuan3DDiTConfig", "MOVAAudioConfig", "MOVAVideoConfig", + "SanaVideoConfig", "StableDiffusion3TransformerConfig", ] diff --git a/python/sglang/multimodal_gen/configs/models/dits/sana_video.py b/python/sglang/multimodal_gen/configs/models/dits/sana_video.py new file mode 100644 index 000000000..59e527fd4 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/dits/sana_video.py @@ -0,0 +1,75 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Architecture configuration for the SANA-Video 3D transformer.""" + +from dataclasses import dataclass, field + +from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig + + +@dataclass +class SanaVideoArchConfig(DiTArchConfig): + patch_size: tuple[int, int, int] = (1, 2, 2) + in_channels: int = 16 + out_channels: int = 16 + num_layers: int = 20 + attention_head_dim: int = 112 + num_attention_heads: int = 20 + num_cross_attention_heads: int = 20 + cross_attention_head_dim: int = 112 + cross_attention_dim: int = 2240 + caption_channels: int = 2304 + mlp_ratio: float = 3.0 + dropout: float = 0.0 + attention_bias: bool = False + sample_size: int = 30 + norm_elementwise_affine: bool = False + norm_eps: float = 1e-6 + guidance_embeds: bool = False + guidance_embeds_scale: float = 0.1 + qk_norm: str = "rms_norm_across_heads" + rope_max_seq_len: int = 1024 + + param_names_mapping: dict = field( + default_factory=lambda: { + # Self-attention q/k/v share the same input. + r"^(transformer_blocks\.\d+\.attn1)\.to_q\.(.*)$": ( + r"\1.to_qkv.\2", + 0, + 3, + ), + r"^(transformer_blocks\.\d+\.attn1)\.to_k\.(.*)$": ( + r"\1.to_qkv.\2", + 1, + 3, + ), + r"^(transformer_blocks\.\d+\.attn1)\.to_v\.(.*)$": ( + r"\1.to_qkv.\2", + 2, + 3, + ), + # Cross-attention k/v share the text input. + r"^(transformer_blocks\.\d+\.attn2)\.to_k\.(.*)$": ( + r"\1.to_kv.\2", + 0, + 2, + ), + r"^(transformer_blocks\.\d+\.attn2)\.to_v\.(.*)$": ( + r"\1.to_kv.\2", + 1, + 2, + ), + r"^transformer\.(.*)$": r"\1", + } + ) + + def __post_init__(self) -> None: + super().__post_init__() + self.patch_size = tuple(self.patch_size) + self.hidden_size = self.num_attention_heads * self.attention_head_dim + self.num_channels_latents = self.out_channels + + +@dataclass +class SanaVideoConfig(DiTConfig): + arch_config: DiTArchConfig = field(default_factory=SanaVideoArchConfig) + prefix: str = "SanaVideo" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py index 0404d57f4..5621aff2e 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py @@ -49,6 +49,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import ( from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.sana_video import ( + SanaVideoPipelineConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.stablediffusion3 import ( StableDiffusion3PipelineConfig, ) @@ -78,6 +81,7 @@ __all__ = [ "Flux2FinetunedPipelineConfig", "PipelineConfig", "SanaPipelineConfig", + "SanaVideoPipelineConfig", "SlidingTileAttnConfig", "MOVAPipelineConfig", "Pi05PipelineConfig", diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py new file mode 100644 index 000000000..a1d35266f --- /dev/null +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Pipeline configuration for SANA-Video text-to-video generation.""" + +from collections.abc import Callable +from dataclasses import dataclass, field + +import torch + +from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig +from sglang.multimodal_gen.configs.models.dits.sana_video import SanaVideoConfig +from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput +from sglang.multimodal_gen.configs.models.encoders.base import EncoderConfig +from sglang.multimodal_gen.configs.models.encoders.gemma2 import Gemma2Config +from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig +from sglang.multimodal_gen.configs.pipeline_configs.base import ( + ModelTaskType, + PipelineConfig, +) + + +def sana_video_postprocess_text( + outputs: BaseEncoderOutput, _text_inputs +) -> torch.Tensor: + return outputs.last_hidden_state + + +@dataclass +class SanaVideoPipelineConfig(PipelineConfig): + task_type: ModelTaskType = ModelTaskType.T2V + should_use_guidance: bool = False + flow_shift: float | None = 8.0 + # Linear attention deliberately accumulates its score products in FP32. + enable_autocast: bool = False + + dit_config: DiTConfig = field(default_factory=SanaVideoConfig) + vae_config: VAEConfig = field(default_factory=WanVAEConfig) + vae_tiling: bool = False + vae_sp: bool = False + vae_precision: str = "fp32" + vae_decode_precision: str = "fp32" + + text_encoder_configs: tuple[EncoderConfig, ...] = field( + default_factory=lambda: (Gemma2Config(),) + ) + text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",)) + text_encoder_extra_args: list[dict] = field( + default_factory=lambda: [ + { + "padding": "max_length", + "return_attention_mask": True, + "add_special_tokens": True, + } + ] + ) + preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field( + default_factory=lambda: (None,) + ) + postprocess_text_funcs: tuple[Callable, ...] = field( + default_factory=lambda: (sana_video_postprocess_text,) + ) + + def __post_init__(self) -> None: + self.vae_config.load_encoder = False + self.vae_config.load_decoder = True + + def adjust_num_frames(self, num_frames: int) -> int: + temporal_scale = self.vae_config.arch_config.temporal_compression_ratio + if num_frames < 1: + raise ValueError("num_frames must be positive") + return ((num_frames - 1) // temporal_scale) * temporal_scale + 1 + + def prepare_latent_shape(self, batch, batch_size, num_frames): + spatial_scale = self.vae_config.arch_config.spatial_compression_ratio + return ( + batch_size, + self.dit_config.arch_config.num_channels_latents, + num_frames, + batch.height // spatial_scale, + batch.width // spatial_scale, + ) + + def get_latent_dtype(self, prompt_dtype: torch.dtype) -> torch.dtype: + return torch.float32 + + def get_pos_prompt_embeds(self, batch): + return batch.prompt_embeds[0] + + def get_neg_prompt_embeds(self, batch): + return batch.negative_prompt_embeds[0] + + @staticmethod + def _unwrap_attention_mask(mask): + if isinstance(mask, (list, tuple)): + return mask[0] if mask else None + return mask + + def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype): + return { + "encoder_attention_mask": self._unwrap_attention_mask( + batch.prompt_attention_mask + ) + } + + def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): + return { + "encoder_attention_mask": self._unwrap_attention_mask( + batch.negative_attention_mask + ) + } + + def post_denoising_loop(self, latents, batch): + return latents + + def shard_latents_for_sp(self, batch, latents): + return latents, False + + def gather_latents_for_sp(self, latents, batch=None): + return latents diff --git a/python/sglang/multimodal_gen/configs/sample/sana_video.py b/python/sglang/multimodal_gen/configs/sample/sana_video.py new file mode 100644 index 000000000..36ba9fc7d --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/sana_video.py @@ -0,0 +1,27 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Sampling defaults for SANA-Video 2B 480p.""" + +from dataclasses import dataclass + +from sglang.multimodal_gen.configs.sample.sampling_params import ( + DataType, + SamplingParams, +) + + +@dataclass +class SanaVideoSamplingParams(SamplingParams): + data_type: DataType = DataType.VIDEO + num_frames: int = 81 + fps: int = 16 + guidance_scale: float = 6.0 + num_inference_steps: int = 50 + height: int = 480 + width: int = 832 + max_sequence_length: int | None = 300 + negative_prompt: str = ( + "A chaotic sequence with misshapen, deformed limbs in heavy motion blur, " + "sudden disappearance, jump cuts, jerky movements, rapid shot changes, " + "frames out of sync, inconsistent character shapes, temporal artifacts, " + "jitter, and ghosting effects, creating a disorienting visual experience." + ) diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index 27e574f4d..2929f3286 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -91,6 +91,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( QwenImagePipelineConfig, ) from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.sana_video import ( + SanaVideoPipelineConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.sana_wm import SanaWMPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.stablediffusion3 import ( StableDiffusion3PipelineConfig, @@ -161,6 +164,7 @@ from sglang.multimodal_gen.configs.sample.qwenimage import ( QwenImageSamplingParams, ) from sglang.multimodal_gen.configs.sample.sana import SanaSamplingParams +from sglang.multimodal_gen.configs.sample.sana_video import SanaVideoSamplingParams from sglang.multimodal_gen.configs.sample.sana_wm import SanaWMSamplingParams from sglang.multimodal_gen.configs.sample.stablediffusion3 import ( StableDiffusion3SamplingParams, @@ -1068,6 +1072,20 @@ def _register_configs(): ], ) + # SANA-Video (register before generic SANA to avoid detector overlap). + register_configs( + sampling_param_cls=SanaVideoSamplingParams, + pipeline_config_cls=SanaVideoPipelineConfig, + hf_model_paths=[ + "Efficient-Large-Model/SANA-Video_2B_480p_diffusers", + ], + model_detectors=[ + lambda hf_id: ( + "sana-video" in hf_id.lower() or "sana_video" in hf_id.lower() + ) + ], + ) + # Cosmos3 — single checkpoint serves T2V, I2V, and T2I. Mode is dispatched # per-request inside the pipeline from ``num_frames`` and ``image_path``. # Both Nano (16B) and Super (64B) share the same pipeline; arch dimensions @@ -1102,6 +1120,8 @@ def _register_configs(): "sana" in hf_id.lower() and "sana-wm" not in hf_id.lower() and "sana_wm" not in hf_id.lower() + and "sana-video" not in hf_id.lower() + and "sana_video" not in hf_id.lower() ) ], ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/sana_video.py b/python/sglang/multimodal_gen/runtime/models/dits/sana_video.py new file mode 100644 index 000000000..f8f1765a3 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/dits/sana_video.py @@ -0,0 +1,513 @@ +# Copyright 2025 The HuggingFace Team and SANA-Video Team. +# SPDX-License-Identifier: Apache-2.0 +"""Native SGLang implementation of the SANA-Video 3D transformer.""" + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F +from diffusers.models.embeddings import PixArtAlphaTextProjection + +from sglang.multimodal_gen.configs.models.dits.sana_video import SanaVideoConfig +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm +from sglang.multimodal_gen.runtime.layers.linear import MergedColumnParallelLinear +from sglang.multimodal_gen.runtime.layers.rotary_embedding.mrope import ( + get_1d_rotary_pos_embed, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) +from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT +from sglang.multimodal_gen.runtime.models.dits.sana import SanaAdaLayerNormSingle + + +def apply_interleaved_rotary_emb( + hidden_states: torch.Tensor, + freqs_cos: torch.Tensor, + freqs_sin: torch.Tensor, +) -> torch.Tensor: + """Apply Diffusers-compatible interleaved real RoPE to ``[B, N, H, D]``.""" + x1, x2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) + cos = freqs_cos[..., 0::2].to(device=hidden_states.device) + sin = freqs_sin[..., 1::2].to(device=hidden_states.device) + output = torch.empty_like(hidden_states) + output[..., 0::2] = x1 * cos - x2 * sin + output[..., 1::2] = x1 * sin + x2 * cos + return output + + +class SanaVideoRotaryPosEmbed(nn.Module): + """3D RoPE split across temporal, height, and width head dimensions.""" + + def __init__( + self, + attention_head_dim: int, + patch_size: tuple[int, int, int], + max_seq_len: int, + theta: float = 10000.0, + ) -> None: + super().__init__() + self.attention_head_dim = attention_head_dim + self.patch_size = patch_size + self.max_seq_len = max_seq_len + self.theta = theta + self._init_freqs_buffers() + + def _init_freqs_buffers(self) -> None: + h_dim = w_dim = 2 * (self.attention_head_dim // 6) + t_dim = self.attention_head_dim - h_dim - w_dim + self.split_sizes = (t_dim, h_dim, w_dim) + + freqs_cos = [] + freqs_sin = [] + for dim in self.split_sizes: + cos, sin = get_1d_rotary_pos_embed( + dim, + self.max_seq_len, + theta=self.theta, + dtype=torch.float64, + ) + freqs_cos.append(cos.repeat_interleave(2, dim=-1)) + freqs_sin.append(sin.repeat_interleave(2, dim=-1)) + self.register_buffer( + "freqs_cos", torch.cat(freqs_cos, dim=-1), persistent=False + ) + self.register_buffer( + "freqs_sin", torch.cat(freqs_sin, dim=-1), persistent=False + ) + + def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + _, _, num_frames, height, width = hidden_states.shape + patch_t, patch_h, patch_w = self.patch_size + frames = num_frames // patch_t + height = height // patch_h + width = width // patch_w + + cos_t, cos_h, cos_w = self.freqs_cos.split(self.split_sizes, dim=-1) + sin_t, sin_h, sin_w = self.freqs_sin.split(self.split_sizes, dim=-1) + + def expand_axis(table, axis): + if axis == 0: + return ( + table[:frames] + .view(frames, 1, 1, -1) + .expand(frames, height, width, -1) + ) + if axis == 1: + return ( + table[:height] + .view(1, height, 1, -1) + .expand(frames, height, width, -1) + ) + return table[:width].view(1, 1, width, -1).expand(frames, height, width, -1) + + cos = torch.cat( + [ + expand_axis(cos_t, 0), + expand_axis(cos_h, 1), + expand_axis(cos_w, 2), + ], + dim=-1, + ).reshape(1, frames * height * width, 1, -1) + sin = torch.cat( + [ + expand_axis(sin_t, 0), + expand_axis(sin_h, 1), + expand_axis(sin_w, 2), + ], + dim=-1, + ).reshape(1, frames * height * width, 1, -1) + return cos, sin + + +class GLUMBTempConv(nn.Module): + """SANA-Video gated spatial MLP with temporal aggregation.""" + + def __init__(self, channels: int, expand_ratio: float) -> None: + super().__init__() + hidden_channels = int(expand_ratio * channels) + self.nonlinearity = nn.SiLU() + self.conv_inverted = nn.Conv2d(channels, hidden_channels * 2, 1) + self.conv_depth = nn.Conv2d( + hidden_channels * 2, + hidden_channels * 2, + 3, + padding=1, + groups=hidden_channels * 2, + ) + self.conv_point = nn.Conv2d(hidden_channels, channels, 1, bias=False) + self.conv_temp = nn.Conv2d( + channels, + channels, + kernel_size=(3, 1), + padding=(1, 0), + bias=False, + ) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + batch_size, num_frames, height, width, channels = hidden_states.shape + hidden_states = hidden_states.reshape( + batch_size * num_frames, height, width, channels + ).permute(0, 3, 1, 2) + hidden_states = self.nonlinearity(self.conv_inverted(hidden_states)) + hidden_states = self.conv_depth(hidden_states) + hidden_states, gate = hidden_states.chunk(2, dim=1) + hidden_states = hidden_states * self.nonlinearity(gate) + hidden_states = self.conv_point(hidden_states) + + temporal = hidden_states.reshape( + batch_size, num_frames, channels, height * width + ).permute(0, 2, 1, 3) + hidden_states = temporal + self.conv_temp(temporal) + return hidden_states.permute(0, 2, 3, 1).reshape( + batch_size, num_frames, height, width, channels + ) + + +class SanaVideoLinearAttention(nn.Module): + """Diffusers-compatible ReLU linear attention with packed QKV.""" + + def __init__( + self, query_dim: int, num_heads: int, head_dim: int, bias: bool + ) -> None: + super().__init__() + self.num_heads = num_heads + self.head_dim = head_dim + self.inner_dim = num_heads * head_dim + self.to_qkv = MergedColumnParallelLinear( + query_dim, + [self.inner_dim, self.inner_dim, self.inner_dim], + bias=bias, + gather_output=True, + ) + self.norm_q = RMSNorm(self.inner_dim, eps=1e-5) + self.norm_k = RMSNorm(self.inner_dim, eps=1e-5) + self.to_out = nn.ModuleList( + [nn.Linear(self.inner_dim, query_dim, bias=True), nn.Identity()] + ) + + def forward( + self, + hidden_states: torch.Tensor, + rotary_emb: tuple[torch.Tensor, torch.Tensor], + ) -> torch.Tensor: + original_dtype = hidden_states.dtype + batch_size, sequence_length, _ = hidden_states.shape + qkv, _ = self.to_qkv(hidden_states) + query, key, value = qkv.split(self.inner_dim, dim=-1) + query = self.norm_q(query).view( + batch_size, sequence_length, self.num_heads, self.head_dim + ) + key = self.norm_k(key).view( + batch_size, sequence_length, self.num_heads, self.head_dim + ) + value = value.view(batch_size, sequence_length, self.num_heads, self.head_dim) + + query = F.relu(query) + key = F.relu(key) + query_rotate = apply_interleaved_rotary_emb(query, *rotary_emb) + key_rotate = apply_interleaved_rotary_emb(key, *rotary_emb) + + query = query.permute(0, 2, 3, 1) + key = key.permute(0, 2, 3, 1) + query_rotate = query_rotate.permute(0, 2, 3, 1).float() + key_rotate = key_rotate.permute(0, 2, 3, 1).float() + value = value.permute(0, 2, 3, 1).float() + + normalizer = 1.0 / ( + key.sum(dim=-1, keepdim=True).transpose(-2, -1) @ query + 1e-15 + ) + scores = value @ key_rotate.transpose(-1, -2) + hidden_states = (scores @ query_rotate) * normalizer + hidden_states = hidden_states.flatten(1, 2).transpose(1, 2) + hidden_states = hidden_states.to(original_dtype) + return self.to_out[0](hidden_states) + + +class SanaVideoCrossAttention(nn.Module): + """Text cross-attention with packed K/V projections.""" + + def __init__( + self, + query_dim: int, + cross_attention_dim: int, + num_heads: int, + head_dim: int, + ) -> None: + super().__init__() + self.num_heads = num_heads + self.head_dim = head_dim + self.inner_dim = num_heads * head_dim + self.to_q = nn.Linear(query_dim, self.inner_dim, bias=True) + self.to_kv = MergedColumnParallelLinear( + cross_attention_dim, + [self.inner_dim, self.inner_dim], + bias=True, + gather_output=True, + ) + self.norm_q = RMSNorm(self.inner_dim, eps=1e-5) + self.norm_k = RMSNorm(self.inner_dim, eps=1e-5) + self.to_out = nn.ModuleList( + [nn.Linear(self.inner_dim, query_dim, bias=True), nn.Identity()] + ) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + encoder_attention_mask: torch.Tensor | None, + ) -> torch.Tensor: + batch_size, query_length, _ = hidden_states.shape + key_length = encoder_hidden_states.shape[1] + query = self.norm_q(self.to_q(hidden_states)) + key_value, _ = self.to_kv(encoder_hidden_states) + key, value = key_value.split(self.inner_dim, dim=-1) + key = self.norm_k(key) + + query = query.view( + batch_size, query_length, self.num_heads, self.head_dim + ).transpose(1, 2) + key = key.view(batch_size, key_length, self.num_heads, self.head_dim).transpose( + 1, 2 + ) + value = value.view( + batch_size, key_length, self.num_heads, self.head_dim + ).transpose(1, 2) + + attention_mask = None + if encoder_attention_mask is not None: + attention_mask = encoder_attention_mask.to(torch.bool)[:, None, None, :] + + hidden_states = F.scaled_dot_product_attention( + query, + key, + value, + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, query_length, self.inner_dim + ) + return self.to_out[0](hidden_states) + + +class SanaVideoTransformerBlock(nn.Module): + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + num_cross_attention_heads: int, + cross_attention_head_dim: int, + cross_attention_dim: int, + mlp_ratio: float, + norm_eps: float, + attention_bias: bool, + ) -> None: + super().__init__() + self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=norm_eps) + self.attn1 = SanaVideoLinearAttention( + dim, num_attention_heads, attention_head_dim, attention_bias + ) + self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=norm_eps) + self.attn2 = SanaVideoCrossAttention( + dim, + cross_attention_dim, + num_cross_attention_heads, + cross_attention_head_dim, + ) + self.ff = GLUMBTempConv(dim, mlp_ratio) + self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + encoder_attention_mask: torch.Tensor | None, + timestep: torch.Tensor, + frames: int, + height: int, + width: int, + rotary_emb: tuple[torch.Tensor, torch.Tensor], + ) -> torch.Tensor: + batch_size = hidden_states.shape[0] + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.scale_shift_table[None, None] + + timestep.reshape(batch_size, timestep.shape[1], 6, -1) + ).unbind(dim=2) + + norm_hidden_states = self.norm1(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa + hidden_states = hidden_states + gate_msa * self.attn1( + norm_hidden_states.to(hidden_states.dtype), rotary_emb + ) + hidden_states = hidden_states + self.attn2( + hidden_states, encoder_hidden_states, encoder_attention_mask + ) + + norm_hidden_states = self.norm2(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp + norm_hidden_states = norm_hidden_states.unflatten(1, (frames, height, width)) + ff_output = self.ff(norm_hidden_states).flatten(1, 3) + return hidden_states + gate_mlp * ff_output + + +class SanaVideoModulatedNorm(nn.Module): + def __init__(self, dim: int, eps: float) -> None: + super().__init__() + self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=eps) + + def forward( + self, + hidden_states: torch.Tensor, + embedded_timestep: torch.Tensor, + scale_shift_table: torch.Tensor, + ) -> torch.Tensor: + shift, scale = ( + scale_shift_table[None, None] + embedded_timestep[:, :, None] + ).unbind(dim=2) + return self.norm(hidden_states) * (1 + scale) + shift + + +class SanaVideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): + _fsdp_shard_conditions = [ + lambda _name, module: isinstance(module, SanaVideoTransformerBlock) + ] + _compile_conditions = [ + lambda _name, module: isinstance(module, SanaVideoTransformerBlock) + ] + param_names_mapping = SanaVideoConfig().arch_config.param_names_mapping + reverse_param_names_mapping = {} + + def __init__(self, config: SanaVideoConfig, hf_config=None, **kwargs) -> None: + super().__init__(config, hf_config=hf_config or {}, **kwargs) + arch = config.arch_config + self.out_channels = arch.out_channels + self.patch_size = tuple(arch.patch_size) + self.inner_dim = arch.num_attention_heads * arch.attention_head_dim + self.hidden_size = self.inner_dim + self.num_attention_heads = arch.num_attention_heads + self.num_channels_latents = arch.num_channels_latents + self.caption_channels = arch.caption_channels + self.cross_attention_dim = arch.cross_attention_dim + + self.rope = SanaVideoRotaryPosEmbed( + arch.attention_head_dim, self.patch_size, arch.rope_max_seq_len + ) + self.patch_embedding = nn.Conv3d( + arch.in_channels, + self.inner_dim, + kernel_size=self.patch_size, + stride=self.patch_size, + ) + if arch.guidance_embeds: + raise NotImplementedError( + "SANA-Video checkpoints with embedded guidance are not supported" + ) + self.time_embed = SanaAdaLayerNormSingle(self.inner_dim) + self.caption_projection = PixArtAlphaTextProjection( + in_features=arch.caption_channels, hidden_size=self.inner_dim + ) + self.caption_norm = RMSNorm(self.inner_dim, eps=1e-5) + self.transformer_blocks = nn.ModuleList( + [ + SanaVideoTransformerBlock( + dim=self.inner_dim, + num_attention_heads=arch.num_attention_heads, + attention_head_dim=arch.attention_head_dim, + num_cross_attention_heads=arch.num_cross_attention_heads, + cross_attention_head_dim=arch.cross_attention_head_dim, + cross_attention_dim=arch.cross_attention_dim, + mlp_ratio=arch.mlp_ratio, + norm_eps=arch.norm_eps, + attention_bias=arch.attention_bias, + ) + for _ in range(arch.num_layers) + ] + ) + self.scale_shift_table = nn.Parameter( + torch.randn(2, self.inner_dim) / self.inner_dim**0.5 + ) + self.norm_out = SanaVideoModulatedNorm(self.inner_dim, arch.norm_eps) + self.proj_out = nn.Linear( + self.inner_dim, math.prod(self.patch_size) * self.out_channels + ) + self.layer_names = ["transformer_blocks"] + + def post_load_weights(self) -> None: + if self.rope.freqs_cos.is_meta or self.rope.freqs_sin.is_meta: + self.rope._init_freqs_buffers() + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + timestep: torch.Tensor, + guidance: torch.Tensor | None = None, + encoder_attention_mask: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + del guidance, kwargs + if encoder_hidden_states is None: + raise ValueError("SANA-Video requires encoder_hidden_states") + if isinstance(encoder_attention_mask, (list, tuple)): + encoder_attention_mask = ( + encoder_attention_mask[0] if encoder_attention_mask else None + ) + + batch_size, _, num_frames, height, width = hidden_states.shape + patch_t, patch_h, patch_w = self.patch_size + post_patch_frames = num_frames // patch_t + post_patch_height = height // patch_h + post_patch_width = width // patch_w + rotary_emb = self.rope(hidden_states) + + hidden_states = self.patch_embedding(hidden_states) + hidden_states = hidden_states.flatten(2).transpose(1, 2) + timestep, embedded_timestep = self.time_embed( + timestep.flatten(), hidden_dtype=hidden_states.dtype + ) + timestep = timestep.view(batch_size, -1, timestep.shape[-1]) + embedded_timestep = embedded_timestep.view( + batch_size, -1, embedded_timestep.shape[-1] + ) + + encoder_hidden_states = self.caption_projection(encoder_hidden_states) + encoder_hidden_states = encoder_hidden_states.view( + batch_size, -1, hidden_states.shape[-1] + ) + encoder_hidden_states = self.caption_norm(encoder_hidden_states) + + for block in self.transformer_blocks: + hidden_states = block( + hidden_states, + encoder_hidden_states, + encoder_attention_mask, + timestep, + post_patch_frames, + post_patch_height, + post_patch_width, + rotary_emb, + ) + + hidden_states = self.norm_out( + hidden_states, embedded_timestep, self.scale_shift_table + ) + hidden_states = self.proj_out(hidden_states) + hidden_states = hidden_states.reshape( + batch_size, + post_patch_frames, + post_patch_height, + post_patch_width, + patch_t, + patch_h, + patch_w, + self.out_channels, + ) + hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) + return hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3).float() + + +EntryClass = SanaVideoTransformer3DModel diff --git a/python/sglang/multimodal_gen/runtime/pipelines/sana_video.py b/python/sglang/multimodal_gen/runtime/pipelines/sana_video.py new file mode 100644 index 000000000..54f8564d5 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines/sana_video.py @@ -0,0 +1,152 @@ +# SPDX-License-Identifier: Apache-2.0 +"""SANA-Video text-to-video pipeline.""" + +import torch + +from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline +from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( + ComposedPipelineBase, +) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages import ( + InputValidationStage, + TextEncodingStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs + +SANA_VIDEO_COMPLEX_HUMAN_INSTRUCTION = ( + "Given a user prompt, generate an 'Enhanced prompt' that provides detailed " + "visual descriptions suitable for video generation. Evaluate the level of " + "detail in the user prompt:\n" + "- If the prompt is simple, focus on adding specifics about colors, shapes, " + "sizes, textures, motion, and temporal relationships to create vivid and " + "dynamic scenes.\n" + "- If the prompt is already detailed, refine and enhance the existing details " + "slightly without overcomplicating.\n" + "Here are examples of how to transform or refine prompts:\n" + "- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat slowly " + "settling into a curled position, peacefully falling asleep on a warm sunny " + "windowsill, with gentle sunlight filtering through surrounding pots of " + "blooming red flowers.\n" + "- User Prompt: A busy city street -> Enhanced: A bustling city street scene " + "at dusk, featuring glowing street lamps gradually lighting up, a diverse " + "crowd of people in colorful clothing walking past, and a double-decker bus " + "smoothly passing by towering glass skyscrapers.\n" + "Please generate only the enhanced description for the prompt below and avoid " + "including any additional commentary or evaluations:\n" + "User Prompt: " +) + + +def select_sana_video_prompt_window( + tensor: torch.Tensor, max_sequence_length: int +) -> torch.Tensor: + """Keep the BOS token and the final prompt window, matching Diffusers.""" + if tensor.shape[1] < max_sequence_length: + raise ValueError( + f"Encoded prompt has {tensor.shape[1]} tokens, expected at least " + f"{max_sequence_length}" + ) + if max_sequence_length == 1: + return tensor[:, :1] + return torch.cat([tensor[:, :1], tensor[:, -(max_sequence_length - 1) :]], dim=1) + + +class SanaVideoTextEncodingStage(TextEncodingStage): + """Apply SANA-Video's asymmetric positive/negative prompt encoding.""" + + @staticmethod + def _normalize_text(text: str | list[str]) -> str | list[str]: + if isinstance(text, str): + return text.lower().strip() + return [item.lower().strip() for item in text] + + def _encode_negative_text(self, batch, server_args, all_indices): + cache_key = self._build_negative_text_cache_key(batch, server_args, all_indices) + cached = self._get_cached_negative_text_embedding(cache_key) + if cached is not None: + return cached + outputs = self.encode_text( + self._normalize_text(batch.negative_prompt), + server_args, + encoder_index=all_indices, + return_attention_mask=True, + max_length=300, + ) + self._maybe_cache_negative_text_embedding(cache_key, outputs) + return outputs + + @torch.no_grad() + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + assert batch.prompt is not None + self.tokenizers[0].padding_side = "right" + all_indices = list(range(len(self.text_encoders))) + max_sequence_length = batch.max_sequence_length or 300 + prompt = self._normalize_text(batch.prompt) + prompt_list = [prompt] if isinstance(prompt, str) else prompt + enhanced_prompt = [ + SANA_VIDEO_COMPLEX_HUMAN_INSTRUCTION + item for item in prompt_list + ] + instruction_tokens = len( + self.tokenizers[0].encode(SANA_VIDEO_COMPLEX_HUMAN_INSTRUCTION) + ) + encoded_length = instruction_tokens + max_sequence_length - 2 + positive_outputs = list( + self.encode_text( + enhanced_prompt, + server_args, + encoder_index=all_indices, + return_attention_mask=True, + max_length=encoded_length, + ) + ) + + for output_index in (0, 1, 3): + positive_outputs[output_index] = [ + select_sana_video_prompt_window(tensor, max_sequence_length) + for tensor in positive_outputs[output_index] + ] + positive_outputs[4] = [ + [int(value) for value in mask.sum(dim=1).tolist()] + for mask in positive_outputs[1] + ] + + self._append_positive_text_outputs(batch, *positive_outputs) + if batch.do_classifier_free_guidance: + negative_outputs = self._encode_negative_text( + batch, server_args, all_indices + ) + self._append_negative_text_outputs( + batch, + positive_outputs[0], + *negative_outputs, + ) + return batch + + +class SanaVideoPipeline(LoRAPipeline, ComposedPipelineBase): + pipeline_name = "SanaVideoPipeline" + _required_config_modules = [ + "text_encoder", + "tokenizer", + "vae", + "transformer", + "scheduler", + ] + + def create_pipeline_stages(self, server_args: ServerArgs): + self.add_stage(InputValidationStage()) + self.add_stage( + SanaVideoTextEncodingStage( + text_encoders=[self.get_module("text_encoder")], + tokenizers=[self.get_module("tokenizer")], + ), + "prompt_encoding_stage_primary", + ) + self.add_standard_timestep_preparation_stage() + self.add_standard_latent_preparation_stage() + self.add_standard_denoising_stage() + self.add_standard_decoding_stage() + + +EntryClass = SanaVideoPipeline diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index a5373340a..1bd33881c 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -34,6 +34,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( MULTI_IMAGE_TI2I_UPLOAD_sampling_params, PI05_ACTION_CI_sampling_params, REALTIME_MODEL_sampling_params, + SANA_VIDEO_T2V_CI_sampling_params, SANA_WM_TI2V_CI_sampling_params, T2I_sampling_params, T2V_sampling_params, @@ -53,6 +54,7 @@ from sglang.multimodal_gen.test.test_utils import ( DEFAULT_QWEN_IMAGE_EDIT_MODEL_NAME_FOR_TEST, DEFAULT_QWEN_IMAGE_LAYERED_MODEL_NAME_FOR_TEST, DEFAULT_QWEN_IMAGE_MODEL_NAME_FOR_TEST, + DEFAULT_SANA_VIDEO_MODEL_NAME_FOR_TEST, DEFAULT_SANA_WM_STREAMING_MODEL_NAME_FOR_TEST, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_WAN_2_1_I2V_14B_480P_MODEL_NAME_FOR_TEST, @@ -236,6 +238,17 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [ modality="video", ), ), + DiffusionTestCase( + "sana_video_2b_t2v", + DiffusionServerArgs( + model_path=DEFAULT_SANA_VIDEO_MODEL_NAME_FOR_TEST, + modality="video", + ), + SANA_VIDEO_T2V_CI_sampling_params, + run_perf_check=False, + run_consistency_check=False, + run_t2v_input_reference_check=False, + ), DiffusionTestCase( "cosmos3_nano_t2v", DiffusionServerArgs( diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 0f6da6e4f..843ad8f80 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -577,6 +577,14 @@ T2V_sampling_params = DiffusionSamplingParams( prompt=T2V_PROMPT, ) +SANA_VIDEO_T2V_CI_sampling_params = DiffusionSamplingParams( + prompt="A curious raccoon walks through a sunlit forest. motion score: 30.", + output_size="832x480", + num_frames=17, + fps=16, + extras={"num_inference_steps": 8, "guidance_scale": 6.0, "seed": 42}, +) + JOY_ECHO_T2V_CI_sampling_params = DiffusionSamplingParams( prompt=T2V_PROMPT, output_size="640x384", diff --git a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/hooks.py b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/hooks.py index a6fc0f2b1..a64c779bc 100644 --- a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/hooks.py +++ b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/hooks.py @@ -62,6 +62,8 @@ def _resolve_transformer_hook_compat(case: Any) -> TransformerHookCompat: normalize_reference_timestep=True, omit_reference_guidance=True, ) + if "sana-video" in model_path or "sana_video" in model_path: + return TransformerHookCompat(omit_reference_guidance=True) if "sana" in model_path: return TransformerHookCompat( omit_reference_guidance=True, diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 5e0c702a9..da5d86cfa 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -203,6 +203,9 @@ DEFAULT_SANA_WM_MODEL_NAME_FOR_TEST = "Efficient-Large-Model/SANA-WM_bidirection DEFAULT_SANA_WM_STREAMING_MODEL_NAME_FOR_TEST = ( "Efficient-Large-Model/SANA-WM_streaming" ) +DEFAULT_SANA_VIDEO_MODEL_NAME_FOR_TEST = ( + "Efficient-Large-Model/SANA-Video_2B_480p_diffusers" +) def print_value_formatted(description: str, value: int | float | str): diff --git a/python/sglang/multimodal_gen/test/unit/test_sana_video.py b/python/sglang/multimodal_gen/test/unit/test_sana_video.py new file mode 100644 index 000000000..73a3219ff --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_sana_video.py @@ -0,0 +1,87 @@ +from types import SimpleNamespace + +import torch + +from sglang.multimodal_gen.configs.pipeline_configs.sana_video import ( + SanaVideoPipelineConfig, +) +from sglang.multimodal_gen.configs.sample.sana_video import SanaVideoSamplingParams +from sglang.multimodal_gen.registry import get_model_info +from sglang.multimodal_gen.runtime.models.dits.sana_video import ( + SanaVideoRotaryPosEmbed, +) +from sglang.multimodal_gen.runtime.pipelines.sana_video import ( + select_sana_video_prompt_window, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import ( + LatentPreparationStage, +) + + +def test_sana_video_registry_resolution(monkeypatch): + monkeypatch.setattr( + "sglang.multimodal_gen.registry.maybe_download_model_index", + lambda _: {"_class_name": "SanaVideoPipeline"}, + ) + get_model_info.cache_clear() + model_info = get_model_info("Efficient-Large-Model/SANA-Video_2B_480p_diffusers") + + assert model_info is not None + assert model_info.pipeline_config_cls is SanaVideoPipelineConfig + assert model_info.sampling_param_cls is SanaVideoSamplingParams + get_model_info.cache_clear() + + +def test_sana_video_pipeline_latent_shape_and_frame_alignment(): + config = SanaVideoPipelineConfig() + sampling = SanaVideoSamplingParams() + + assert config.adjust_num_frames(81) == 81 + assert config.adjust_num_frames(80) == 77 + batch = SimpleNamespace(height=480, width=832, num_frames=81) + server_args = SimpleNamespace(pipeline_config=config) + latent_frames = LatentPreparationStage( + scheduler=None, transformer=None + ).adjust_video_length(batch, server_args) + assert latent_frames == 21 + # LatentPreparationStage applies temporal compression before calling the config. + assert config.prepare_latent_shape( + batch, batch_size=2, num_frames=latent_frames + ) == ( + 2, + 16, + 21, + 60, + 104, + ) + assert config.get_latent_dtype(torch.bfloat16) is torch.float32 + assert not config.enable_autocast + assert not config.vae_config.load_encoder + assert config.vae_config.load_decoder + assert (sampling.width, sampling.height, sampling.num_frames) == (832, 480, 81) + assert sampling.fps == 16 + assert sampling.num_inference_steps == 50 + assert sampling.guidance_scale == 6.0 + + +def test_select_sana_video_prompt_window_keeps_first_and_tail_tokens(): + tensor = torch.arange(10).view(1, 10, 1) + + selected = select_sana_video_prompt_window(tensor, max_sequence_length=4) + + assert selected.flatten().tolist() == [0, 7, 8, 9] + + +def test_sana_video_rotary_embeddings_follow_video_token_order(): + rotary = SanaVideoRotaryPosEmbed( + attention_head_dim=12, + patch_size=(1, 2, 2), + max_seq_len=16, + ) + + cos, sin = rotary(torch.zeros(1, 4, 3, 4, 4)) + + assert cos.shape == (1, 12, 1, 12) + assert sin.shape == (1, 12, 1, 12) + assert torch.isfinite(cos).all() + assert torch.isfinite(sin).all()