Skip to content

vllm_omni.diffusion.models.sana_wm

SANA-WM diffusion model integration.

Modules:

Name Description
camera_control

Camera-control helpers for SANA-WM.

config

Config for the SANA-WM Stage-1 transformer.

pipeline_sana_wm

Sana-WM pipeline integration.

request

Request normalization for Sana-WM image-to-video.

sana_wm_transformer

SANA-WM Stage-1 transformer.

ucpe

SANA-WM UCPE (Unified Camera Pose Embedding) per-block attention transforms.

SANA_WM_MODEL_ID module-attribute

SANA_WM_MODEL_ID = (
    "BBBBruce/SANA-WM_bidirectional-stage1-diffusers"
)

SANA_WM_OUTPUT_HEIGHT module-attribute

SANA_WM_OUTPUT_HEIGHT = 704

SANA_WM_OUTPUT_WIDTH module-attribute

SANA_WM_OUTPUT_WIDTH = 1280

SanaWmConfig dataclass

SANA-WM Stage-1 architecture and runtime config.

Defaults mirror the first public HF release. from_json (reading transformer/config.json) is the normal construction path.

architecture_name class-attribute instance-attribute

architecture_name: str | None = (
    "SanaMSVideoCamCtrl_1600M_P1_D20"
)

attn_type class-attribute instance-attribute

attn_type: str = 'BidirectionalGDNTriton'

cam_attn_compress class-attribute instance-attribute

cam_attn_compress: int = 1

chi_prompt class-attribute instance-attribute

chi_prompt: list[str] = field(default_factory=list)

chunk_plucker_channels class-attribute instance-attribute

chunk_plucker_channels: int = 48

chunk_plucker_post_attn_blocks class-attribute instance-attribute

chunk_plucker_post_attn_blocks: int = 20

conv_kernel_size class-attribute instance-attribute

conv_kernel_size: int = 4

cross_norm class-attribute instance-attribute

cross_norm: bool = True

ffn_type class-attribute instance-attribute

ffn_type: str = 'GLUMBConvTemp'

fp32_attention class-attribute instance-attribute

fp32_attention: bool = True

hidden_size class-attribute instance-attribute

hidden_size: int = 2240

image_size class-attribute instance-attribute

image_size: int = 720

inference_flow_shift class-attribute instance-attribute

inference_flow_shift: float = 9.8

k_conv_only class-attribute instance-attribute

k_conv_only: bool = True

linear_head_dim class-attribute instance-attribute

linear_head_dim: int = 112

mixed_precision class-attribute instance-attribute

mixed_precision: str = 'bf16'

mlp_ratio class-attribute instance-attribute

mlp_ratio: float = 3.0

model_max_length class-attribute instance-attribute

model_max_length: int = 300

num_blocks class-attribute instance-attribute

num_blocks: int = 20

patch_size class-attribute instance-attribute

patch_size: tuple[int, int, int] = (1, 1, 1)

pos_embed_type class-attribute instance-attribute

pos_embed_type: str = 'wan_rope'

qk_norm class-attribute instance-attribute

qk_norm: bool = True

scheduler_type class-attribute instance-attribute

scheduler_type: str = 'flow_dpm-solver'

softmax_every_n class-attribute instance-attribute

softmax_every_n: int = 4

t_kernel_size class-attribute instance-attribute

t_kernel_size: int = 3

use_chunk_plucker_post_attn class-attribute instance-attribute

use_chunk_plucker_post_attn: bool = True

from_dict classmethod

from_dict(data: Mapping[str, Any]) -> SanaWmConfig

Build from a flat mapping of field name -> value.

Unknown keys (e.g. the diffusers _class_name marker or the latent_channels / prompt_channels constructor kwargs) are ignored. patch_size is coerced from a JSON list back to a tuple.

from_json classmethod

from_json(path: str | Path) -> SanaWmConfig

SanaWmPipeline

Bases: Module, CFGParallelMixin, SupportImageInput, SupportsComponentDiscovery, ProgressBarMixin, DiffusionPipelineProfilerMixin

Stage-1 SANA-WM image-to-video pipeline.

color_format class-attribute

color_format: str = 'RGB'

device instance-attribute

device = get_local_device()

dummy_run_num_frames class-attribute

dummy_run_num_frames: int = 0

od_config instance-attribute

od_config = od_config

prefix instance-attribute

prefix = prefix

quant_config instance-attribute

quant_config = (
    getattr(od_config, "quantization_config", None)
    if od_config is not None
    else None
)

release_paths instance-attribute

release_paths: SanaWmLocalPaths | None = None

sana_wm_config instance-attribute

sana_wm_config = SanaWmConfig()

support_image_input class-attribute

support_image_input: bool = True

text_encoder instance-attribute

text_encoder: Module | None = None

tokenizer instance-attribute

tokenizer: Any | None = None

transformer instance-attribute

transformer = SanaWmTransformer3DModel(
    config=self.sana_wm_config,
    quant_config=self.quant_config,
    prefix=f"{prefix}.transformer"
    if prefix
    else "transformer",
)

vae instance-attribute

vae: Module | None = None

weights_sources instance-attribute

weights_sources = []

forward

forward(
    req: DiffusionRequestBatch, *args: Any, **kwargs: Any
) -> DiffusionOutput

load_weights

load_weights(
    weights: Iterable[tuple[str, Any]],
) -> set[str]

predict_noise

predict_noise(**kwargs: Any) -> Tensor

Single transformer forward.

Overrides the mixin default, which assumes a tuple-returning transformer and takes result[0]; SANA-WM's returns the noise prediction directly.

resolve_checkpoint

resolve_checkpoint() -> SanaWmLocalPaths

SanaWmTransformer3DModel

Bases: Module

SANA-WM Stage-1 DiT: bidirectional Gated DeltaNet blocks with a softmax attention block every softmax_every_n.

attention_y_norm instance-attribute

attention_y_norm = RMSNorm(self.config.hidden_size)

blocks instance-attribute

blocks = nn.ModuleList(
    [
        SanaWmBlock(
            self.config,
            block_idx=i,
            quant_config=self.quant_config,
            prefix=f"{self.prefix}.blocks.{i}"
            if self.prefix
            else f"blocks.{i}",
        )
        for i in range(self.config.num_blocks)
    ]
)

config instance-attribute

config = config or SanaWmConfig()

final_layer instance-attribute

final_layer = SanaWmFinalLayer(
    self.config.hidden_size,
    self.patch_size,
    self._latent_channels,
    quant_config=self.quant_config,
    prefix=f"{self.prefix}.final_layer"
    if self.prefix
    else "final_layer",
)

patch_size instance-attribute

patch_size = _to_3tuple(self.config.patch_size)

plucker_embedder instance-attribute

plucker_embedder = SanaWmPatchEmbedMS3D(
    self.patch_size,
    self.config.chunk_plucker_channels,
    self.config.hidden_size,
)

pos_embed instance-attribute

pos_embed = nn.Parameter(
    torch.zeros(1, 484, self.config.hidden_size)
)

prefix instance-attribute

prefix = prefix

quant_config instance-attribute

quant_config = quant_config

raymap_embedder instance-attribute

raymap_embedder = SanaWmPatchEmbedMS3D(
    self.patch_size, 3, self.config.hidden_size
)

rope instance-attribute

rope = SanaWmWanRotaryPosEmbed(self.config.linear_head_dim)

t_block instance-attribute

t_block = nn.Sequential(nn.SiLU(), _t_block_linear)

t_embedder instance-attribute

t_embedder = SanaWmTimestepEmbedder(
    SANA_WM_STAGE1_TIMESTEP_CHANNELS,
    self.config.hidden_size,
    quant_config=self.quant_config,
    prefix=f"{self.prefix}.t_embedder"
    if self.prefix
    else "t_embedder",
)

x_embedder instance-attribute

x_embedder = SanaWmPatchEmbedMS3D(
    self.patch_size,
    self._latent_channels,
    self.config.hidden_size,
)

y_embedder instance-attribute

y_embedder = SanaWmTextEmbedder(
    self._prompt_channels,
    self.config.hidden_size,
    self.config.model_max_length,
    quant_config=self.quant_config,
    prefix=f"{self.prefix}.y_embedder"
    if self.prefix
    else "y_embedder",
)

forward

forward(
    hidden_states: Tensor,
    timestep: Tensor | float | int,
    *,
    encoder_hidden_states: Tensor | None = None,
    encoder_attention_mask: Tensor | None = None,
    camera_hidden_states: Tensor | None = None,
    plucker: Tensor | None = None,
    raymap: Tensor | None = None,
    spatial_raymap: Tensor | None = None,
) -> Tensor

load_weights

load_weights(
    weights: Iterable[tuple[str, Any]],
) -> set[str]

Stream checkpoint tensors into the eagerly-built modules.

Follows the wan2_2 idiom: one named_parameters lookup table plus the per-tensor weight_loader, which narrows full checkpoint tensors to the TP-local shard at copy time for both vLLM parallel layers and the plain parameters marked by _shard_param_across_tp. The source tensor is dropped each iteration — no copy of the checkpoint is retained.

get_sana_wm_pre_process_func

get_sana_wm_pre_process_func(
    od_config: OmniDiffusionConfig,
)

normalize_sana_wm_payload

normalize_sana_wm_payload(
    prompt: Mapping[str, Any],
) -> dict[str, Any]

Return a prompt copy with canonical Sana-WM request metadata.

The sana_wm block is read from the top-level sana_wm key, falling back to additional_information["sana_wm"] so the function is idempotent when called again on its own output. The first-frame image is read from multi_modal_data["image"].