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
xiparameter is fixed. - Inference path only — no training-time camera-branch dropout.
- Online computation only — the precomputed
cam_pos_embedsshortcut 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=Trueonly — 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.
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 |
| required |
k_normed | Tensor |
| required |
v_raw | Tensor |
| required |
proj_q | Tensor |
| required |
proj_kv | Tensor |
| required |
rope_cos | Tensor |
| required |
rope_sin | Tensor |
| required |
k_scale | float |
| required |
Returns:
| Type | Description |
|---|---|
Tensor |
|
Tensor |
|
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.