Skip to content

vllm_omni.diffusion.models.sana_wm.ucpe

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

Ported from NVlabs/Sana sana_camctrl_blocks.py for the inference-only SANA-WM bidirectional 1600M release. Scope is intentionally narrow:

  • Pinhole camera only (xi=0); the UCM xi parameter is fixed.
  • Inference path only — no training-time camera-branch dropout.
  • Online computation only — the precomputed cam_pos_embeds shortcut used by NVlabs' fused kernels is omitted because vLLM-Omni builds the matrices once per transformer forward and shares them across all blocks.
  • apply_vo=True only — the SANA-WM checkpoint always uses the inverse-output transform.
Public surface

prepare_cam_geometry(...) -> :class:SanaWmCamGeometry, built once per transformer forward and threaded through every block. It carries both the (apply_q, apply_kv, apply_output) closures used by the softmax camera branch and the projection/RoPE tensors used by cam_prep_func(...), the NVlabs fused camera-prep contract on the GDN camera branch.

Each closure transforms a tensor of shape (B, num_heads, N, D) where N = T * H * W (matching the GDN token layout) and D == head_dim, returning a tensor of the same shape.

SanaWmCamGeometry dataclass

UCPE geometry shared by every block of one transformer forward.

The ray grid, its SE(3) inverse and the camera RoPE tables depend only on the camera payload, the latent grid, the patch size and the head dim — all fixed for the whole forward — yet each block's camera branch used to rebuild them. At 20 blocks x 60 steps x 2 CFG branches that is ~2400 rebuilds of an 18480-token ray grid and its 4x4 inverses per request.

apply_q/apply_kv/apply_output serve the softmax camera branch; proj_q/proj_kv/rope_cos/rope_sin serve the fused cam_prep_func contract on the GDN branch. Both are derived from the same projection matrices, so they are built together.

apply_kv instance-attribute

apply_kv: Callable[[Tensor], Tensor]

apply_output instance-attribute

apply_output: Callable[[Tensor], Tensor]

apply_q instance-attribute

apply_q: Callable[[Tensor], Tensor]

proj_kv instance-attribute

proj_kv: Tensor

proj_q instance-attribute

proj_q: Tensor

rope_cos instance-attribute

rope_cos: Tensor

rope_sin instance-attribute

rope_sin: Tensor

cam_prep_func

cam_prep_func(
    q_normed: Tensor,
    k_normed: Tensor,
    v_raw: Tensor,
    *,
    proj_q: Tensor,
    proj_kv: Tensor,
    rope_cos: Tensor,
    rope_sin: Tensor,
    k_scale: float,
) -> tuple[Tensor, Tensor, Tensor, Tensor]

Native equivalent of NVlabs cam_prep_func.

The caller is responsible for the q/k RMSNorm, which must run through the model's norm modules so a tensor-parallel shard still normalises over the global channel width. This function only applies ReLU, the K scale, the UCPE ray projection and RoPE.

Parameters:

Name Type Description Default
q_normed Tensor

(B, N, H, D) camera Q, already RMSNorm'd.

required
k_normed Tensor

(B, N, H, D) camera K, already RMSNorm'd. K must also include the temporal short convolution.

required
v_raw Tensor

(B, N, H, D) raw camera V.

required
proj_q Tensor

(B, N, 4, 4) Q ray projection matrices.

required
proj_kv Tensor

(B, N, 4, 4) K/V ray projection matrices.

required
rope_cos Tensor

(N, D//2) interleaved-pair RoPE cosine table.

required
rope_sin Tensor

(N, D//2) interleaved-pair RoPE sine table.

required
k_scale float

D^-0.5 * spatial_tokens^-0.5.

required

Returns:

Type Description
Tensor

q_trans, k_trans, v_trans in (B, H, D, N) layout, and

Tensor

inflation_sq in (B, H, N).

prepare_cam_geometry

prepare_cam_geometry(
    *,
    camera_conditions: Tensor,
    spatial_shape: tuple[int, int, int],
    patch_size: tuple[int, int, int],
    head_dim: int,
    rotary_emb: Tensor | None,
) -> SanaWmCamGeometry

Build the per-forward UCPE geometry: ray projections, RoPE tables, apply fns.

Call once per transformer forward and share the result across every block — see :class:SanaWmCamGeometry for why.