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 ¶
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"
)
chunk_plucker_post_attn_blocks class-attribute instance-attribute ¶
chunk_plucker_post_attn_blocks: int = 20
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.
SanaWmPipeline ¶
Bases: Module, CFGParallelMixin, SupportImageInput, SupportsComponentDiscovery, ProgressBarMixin, DiffusionPipelineProfilerMixin
Stage-1 SANA-WM image-to-video pipeline.
quant_config instance-attribute ¶
quant_config = (
getattr(od_config, "quantization_config", None)
if od_config is not None
else None
)
transformer instance-attribute ¶
transformer = SanaWmTransformer3DModel(
config=self.sana_wm_config,
quant_config=self.quant_config,
prefix=f"{prefix}.transformer"
if prefix
else "transformer",
)
SanaWmTransformer3DModel ¶
Bases: Module
SANA-WM Stage-1 DiT: bidirectional Gated DeltaNet blocks with a softmax attention block every softmax_every_n.
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)
]
)
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",
)
plucker_embedder instance-attribute ¶
plucker_embedder = SanaWmPatchEmbedMS3D(
self.patch_size,
self.config.chunk_plucker_channels,
self.config.hidden_size,
)
pos_embed instance-attribute ¶
raymap_embedder instance-attribute ¶
raymap_embedder = SanaWmPatchEmbedMS3D(
self.patch_size, 3, self.config.hidden_size
)
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 ¶
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.
normalize_sana_wm_payload ¶
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"].