Skip to content

vllm_omni.diffusion.models.sana_wm.camera_control

Camera-control helpers for SANA-WM.

This module mirrors the public NVlabs/Sana SANA-WM camera preparation path: WASD/IJKL action rollout or explicit camera-to-world poses are converted into relative poses, per-latent-frame ray metadata, and per-VAE-chunk Plucker maps. Model-local projection layers consume these tensors in the Stage-1 DiT.

SANA_WM_ALLOWED_ACTION_KEYS module-attribute

SANA_WM_ALLOWED_ACTION_KEYS = frozenset('wasdijkl')

SANA_WM_DEFAULT_PITCH_LIMIT_DEG module-attribute

SANA_WM_DEFAULT_PITCH_LIMIT_DEG = 85.0

SANA_WM_DEFAULT_ROTATION_SPEED_DEG module-attribute

SANA_WM_DEFAULT_ROTATION_SPEED_DEG = 1.2

SANA_WM_DEFAULT_TRANSLATION_SPEED module-attribute

SANA_WM_DEFAULT_TRANSLATION_SPEED = 0.05

SANA_WM_DEFAULT_VAE_STRIDE module-attribute

SANA_WM_DEFAULT_VAE_STRIDE = (8, 32, 32)

SanaWmCameraCondition dataclass

action class-attribute instance-attribute

action: str | None = None

height class-attribute instance-attribute

height: int = 704

intrinsics class-attribute instance-attribute

intrinsics: Any | None = None

num_frames class-attribute instance-attribute

num_frames: int | None = None

pitch_limit_deg class-attribute instance-attribute

poses class-attribute instance-attribute

poses: Any | None = None

rotation_speed_deg class-attribute instance-attribute

rotation_speed_deg: float = (
    SANA_WM_DEFAULT_ROTATION_SPEED_DEG
)

translation_speed class-attribute instance-attribute

width class-attribute instance-attribute

width: int = 1280

action_rollout_num_frames

action_rollout_num_frames(action: str) -> int

Frames action rolls out: the identity start pose plus one per step.

Totals the durations arithmetically, so an oversized request is rejected before anything allocates per-frame state.

action_string_to_c2w

action_string_to_c2w(
    action: str,
    *,
    translation_speed: float = SANA_WM_DEFAULT_TRANSLATION_SPEED,
    rotation_speed_deg: float = SANA_WM_DEFAULT_ROTATION_SPEED_DEG,
    pitch_limit_deg: float = SANA_WM_DEFAULT_PITCH_LIMIT_DEG,
) -> ndarray

Roll out an OpenCV-convention (N + 1, 4, 4) c2w trajectory.

build_plucker_condition

build_plucker_condition(
    condition: SanaWmCameraCondition,
    *,
    vae_stride: tuple[int, int, int]
    | list[int] = SANA_WM_DEFAULT_VAE_STRIDE,
) -> dict[str, Tensor]

Build native Stage-1 camera tensors from normalized request metadata.

compute_raymap

compute_raymap(
    intrinsics: Tensor,
    poses: Tensor,
    height: int,
    width: int,
    *,
    use_plucker: bool = True,
) -> Tensor

Compute SANA-WM geometry ray maps.

Parameters:

Name Type Description Default
intrinsics Tensor

(T, 4) tensor in [fx, fy, cx, cy] order.

required
poses Tensor

(T, 4, 4) camera-to-world matrices.

required
height int

raymap height in latent pixels.

required
width int

raymap width in latent pixels.

required
use_plucker bool

return [direction, moment] if true, otherwise [origin, direction].

True

get_pose_inverse

get_pose_inverse(transform: Tensor) -> Tensor

Invert homogeneous rigid transforms with shape (..., 4, 4).

intrinsics_to_vec4_array

intrinsics_to_vec4_array(
    intrinsics: Any,
    *,
    num_frames: int,
    height: int,
    width: int,
) -> ndarray

Normalize the intrinsics payload to an (F, 4) [fx, fy, cx, cy] array.

Accepts either None (derive from the output resolution) or the {fx, fy, cx, cy} mapping — the only form the request contract exposes.