Skip to content

vllm_omni.diffusion.models.sana_wm.pipeline_sana_wm

Sana-WM pipeline integration.

This module wires the registry-visible surface, release-layout validation, and the Stage-1 sampling loop. The DiT, Gated-DeltaNet and camera stack are vLLM-Omni-native layers (sana_wm_transformer.py / ucpe.py); nothing here calls into NVlabs code at runtime.

SANA_WM_CONFIG_FILE module-attribute

SANA_WM_CONFIG_FILE = 'transformer/config.json'

SANA_WM_DEFAULT_GUIDANCE_SCALE module-attribute

SANA_WM_DEFAULT_GUIDANCE_SCALE = 5.0

SANA_WM_DEFAULT_NUM_INFERENCE_STEPS module-attribute

SANA_WM_DEFAULT_NUM_INFERENCE_STEPS = 60

SANA_WM_MODEL_ID module-attribute

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

SANA_WM_NATIVE_MAX_TOKENS module-attribute

SANA_WM_NATIVE_NUM_FRAMES module-attribute

SANA_WM_NATIVE_NUM_FRAMES = 321

SANA_WM_NUM_TRAIN_TIMESTEPS module-attribute

SANA_WM_NUM_TRAIN_TIMESTEPS = 1000

SANA_WM_OUTPUT_HEIGHT module-attribute

SANA_WM_OUTPUT_HEIGHT = 704

SANA_WM_OUTPUT_WIDTH module-attribute

SANA_WM_OUTPUT_WIDTH = 1280

SANA_WM_STAGE1_DIT_BASENAME module-attribute

SANA_WM_STAGE1_DIT_BASENAME = Path(
    SANA_WM_STAGE1_DIT_FILE
).name

SANA_WM_STAGE1_DIT_FILE module-attribute

SANA_WM_STAGE1_DIT_FILE = (
    "transformer/diffusion_pytorch_model.safetensors"
)

SANA_WM_STAGE1_DIT_SUBFOLDER module-attribute

SANA_WM_STAGE1_DIT_SUBFOLDER = 'transformer'

SANA_WM_STAGE1_PATTERNS module-attribute

SANA_WM_STAGE1_TEXT_ENCODER_ENV module-attribute

SANA_WM_STAGE1_TEXT_ENCODER_ENV = (
    "VLLM_OMNI_SANA_WM_STAGE1_TEXT_ENCODER"
)

SANA_WM_STAGE1_TEXT_ENCODER_FALLBACK_ID module-attribute

SANA_WM_STAGE1_TEXT_ENCODER_FALLBACK_ID = (
    "Efficient-Large-Model/gemma-2-2b-it"
)

SANA_WM_STAGE1_TEXT_ENCODER_ID module-attribute

SANA_WM_STAGE1_TEXT_ENCODER_ID = 'google/gemma-2-2b-it'

SANA_WM_VAE_CONFIG_FILE module-attribute

SANA_WM_VAE_CONFIG_FILE = 'vae/config.json'

SANA_WM_VAE_WEIGHT_FILE module-attribute

SANA_WM_VAE_WEIGHT_FILE = (
    "vae/diffusion_pytorch_model.safetensors"
)

logger module-attribute

logger = init_logger(__name__)

SanaWmLocalPaths dataclass

Resolved local file paths for a SANA-WM snapshot.

config instance-attribute

config: Path

root instance-attribute

root: Path

stage1_dit instance-attribute

stage1_dit: Path

vae_config instance-attribute

vae_config: Path

vae_weights instance-attribute

vae_weights: Path

SanaWmNativeParams dataclass

Resolved generation settings for one Stage-1 request.

Built by _native_params from the request payload and sampling params, and consumed by _run_native_backend — this is the production path.

cfg_scale class-attribute instance-attribute

height instance-attribute

height: int

num_frames instance-attribute

num_frames: int

num_inference_steps instance-attribute

num_inference_steps: int

seed instance-attribute

seed: int

width instance-attribute

width: int

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

build_sana_wm_download_patterns

build_sana_wm_download_patterns() -> tuple[str, ...]

Return the minimal HF allow-patterns needed for SANA-WM.

build_sana_wm_output_envelope

build_sana_wm_output_envelope(
    *,
    output: Any,
    output_type: str,
    metadata: dict[str, Any],
) -> dict[str, Any]

Wrap a Stage-1 result in the canonical output envelope.

normalize_diffusion_postprocess_output splits {"payload": ..., "metadata": ...} into the API-facing payload and the metadata groups, so the model-specific diagnostics ride along under a sana_wm group instead of the removed DiffusionOutput.custom_output field.

get_sana_wm_pre_process_func

get_sana_wm_pre_process_func(
    od_config: OmniDiffusionConfig,
)

resolve_or_download_sana_wm_checkpoint

resolve_or_download_sana_wm_checkpoint(
    model: str = SANA_WM_MODEL_ID,
    *,
    revision: str | None = None,
    cache_dir: str | None = None,
) -> SanaWmLocalPaths

Resolve a local SANA-WM tree or download the required HF files.

resolve_sana_wm_local_paths

resolve_sana_wm_local_paths(
    snapshot_dir: str | Path,
) -> SanaWmLocalPaths

validate_sana_wm_local_paths

validate_sana_wm_local_paths(
    paths: SanaWmLocalPaths,
) -> None