From a53d3636ce050c638a7539645da513d23676ef88 Mon Sep 17 00:00:00 2001
From: Xiaoyu Zhang <1182563586@qq.com>
Date: Wed, 12 Aug 2026 10:07:24 +0800
Subject: [PATCH] [diffusion][model] Add native SANA-Video T2V support (#32921)
---
.../sglang-diffusion/compatibility_matrix.mdx | 21 +
.../configs/models/dits/__init__.py | 2 +
.../configs/models/dits/sana_video.py | 75 +++
.../configs/pipeline_configs/__init__.py | 4 +
.../configs/pipeline_configs/sana_video.py | 118 ++++
.../configs/sample/sana_video.py | 27 +
python/sglang/multimodal_gen/registry.py | 20 +
.../runtime/models/dits/sana_video.py | 513 ++++++++++++++++++
.../runtime/pipelines/sana_video.py | 152 ++++++
.../multimodal_gen/test/server/gpu_cases.py | 13 +
.../test/server/testcase_configs.py | 8 +
.../component_accuracy/hooks.py | 2 +
.../sglang/multimodal_gen/test/test_utils.py | 3 +
.../test/unit/test_sana_video.py | 87 +++
14 files changed, 1045 insertions(+)
create mode 100644 python/sglang/multimodal_gen/configs/models/dits/sana_video.py
create mode 100644 python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py
create mode 100644 python/sglang/multimodal_gen/configs/sample/sana_video.py
create mode 100644 python/sglang/multimodal_gen/runtime/models/dits/sana_video.py
create mode 100644 python/sglang/multimodal_gen/runtime/pipelines/sana_video.py
create mode 100644 python/sglang/multimodal_gen/test/unit/test_sana_video.py
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()