Implementing SANA-WM for Controllable World Modeling with 6-DoF Camera Control
SANA-WM enables 720p, 1-minute video generation with precise 6-DoF camera control by injecting pose sequences into a 2.6B-parameter Linear-attention DiT through a specialized depth-wise MLP block.
SANA-WM (World Model) extends the NVlabs/Sana ecosystem to support controllable world modeling via explicit camera trajectory conditioning. Built on the efficient SANA architecture—featuring Linear-attention DiT blocks and DC-AE latent compression—this component adds lightweight 6-DoF pose injection to generate smooth, camera-controlled videos while maintaining memory efficiency for single-GPU deployment.
Understanding the SANA-WM Architecture
SANA-WM follows a three-stage pipeline: latent generation via diffusion, camera-pose injection through cross-attention mechanisms, and latent decoding via DC-AE VAE. The model processes text prompts alongside sequences of 6-DoF camera poses (3 position values + 3 orientation angles) to produce temporally coherent video outputs.
The DWMlp Block for Spatio-Temporal Mixing
At the heart of the world model is the DWMlp block, defined in diffusion/model/nets/basic_modules.py (lines 467-523). Unlike standard MLP layers, DWMlp extends the base Mlp class with a depth-wise 2-D convolution (self.conv) that captures long-range temporal dependencies across video frames:
# Conceptual structure from basic_modules.py
class DWMlp(nn.Module):
def __init__(self, in_features, hidden_features=None, ...):
super().__init__()
self.fc1 = nn.Linear(in_features, hidden_features)
self.conv = nn.Conv2d(...) # Depth-wise temporal aggregation
self.fc2 = nn.Linear(hidden_features, in_features)
def forward(self, x, H, W):
# Mixes spatial tokens with temporal context
...
This block enables the model to fuse camera pose information with visual tokens efficiently, allowing the 2.6B parameter model to maintain temporal consistency without quadratic attention costs.
Camera-Pose Conditioning Mechanism
Camera control is implemented through a lightweight conditioning branch. The system accepts pose sequences as N × 6 tensors representing [x, y, z, yaw, pitch, roll] for each frame. In diffusion/model/nets/sana.py, the main DiT backbone inserts these embeddings before the DWMlp layers, allowing pose information to influence the entire generation process through the depth-wise convolutions' temporal receptive field.
The GLUMBConv and GLUMBConvTemp modules (also in basic_modules.py) provide gated linear units with spatial-temporal aggregation, further stabilizing the integration of camera trajectories into the visual generation stream.
How 6-DoF Camera Control Works
The 6-DoF control system operates by projecting pose vectors into the latent space of the diffusion model. When you provide a sequence of camera poses, the pipeline executes the following data flow:
- Pose Embedding: Each 6-DoF vector is projected by a tiny MLP into the model's hidden dimension
- Token Injection: Pose tokens are broadcast and concatenated to the visual token stream before entering the transformer blocks
- Temporal Mixing: The
DWMlpprocesses these combined tokens using depth-wise convolutions that span the temporal dimension, ensuring smooth trajectory adherence throughout the generated video
Because the conditioning uses depth-wise operations rather than full cross-attention, the overhead remains minimal—preserving the 32× compression factor of DC-AE that allows 720p video generation on a single RTX 3090 GPU.
Implementation: Running Inference with SANA-WM
The SanaWMPipeline class in app/sana_pipeline.py provides a high-level Diffusers-compatible API for world model inference. Below is the complete workflow for generating camera-controlled videos.
Loading the Pipeline and Model Weights
First, instantiate the pipeline with the pre-trained SANA-WM checkpoint:
import torch
from app.sana_pipeline import SanaWMPipeline
pipe = SanaWMPipeline.from_pretrained(
"NVlabs/SANA-WM-2.6B-720p",
torch_dtype=torch.bfloat16,
)
pipe.enable_model_cpu_offload() # Offload VAE to CPU when not in use
Defining Camera Trajectories
Define your camera path as a tensor of shape (frames, 6), where columns represent [x, y, z, yaw, pitch, roll]:
# 60-frame trajectory for 1-minute video at 60fps
camera_poses = torch.tensor([
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0], # Frame 1: origin
[0.1, 0.0, 0.0, 5.0, -2.0, 0.0], # Frame 2: slight movement
# ... additional frames
], dtype=torch.float32)
Generating the Video
Pass the pose tensor to the pipeline alongside your text prompt:
generator = torch.Generator("cuda").manual_seed(42)
video = pipe(
prompt="A futuristic cityscape at dusk, viewed from a flying drone",
camera_poses=camera_poses, # 6-DoF conditioning input
height=720,
width=1280,
frames=60,
guidance_scale=6.5,
num_inference_steps=50,
generator=generator,
output_type="np",
)[0]
# Export to video file
from diffusers.utils import export_to_video
export_to_video(video, "sana_wm_output.mp4", fps=16)
The camera_poses argument triggers the pose encoder inside DWMlp automatically—no manual token handling is required.
Training Custom World Models
For fine-tuning or training from scratch, the repository provides distributed training scripts in train_video_scripts/train_video_ivjoint.py. The training loop integrates pose conditioning through the model's pose encoder, typically defined in diffusion/model/nets/sana_multi_scale_adaln.py.
Distributed Training Setup
Key utilities for multi-GPU training reside in diffusion/utils/dist_utils.py:
# Simplified training loop structure from train_video_ivjoint.py
from diffusion.utils.dist_utils import get_world_size
for step, batch in enumerate(train_loader):
# batch["pixel_values"] → latent video representations
# batch["camera_poses"] → (B, T, 6) pose tensors
latents = encoder(batch["pixel_values"]) # DC-AE encoding
pose_emb = model.pose_encoder(batch["camera_poses"]) # Pose projection
# Diffusion forward with pose conditioning
loss = diffusion_model(latents, pose_emb, **diffusion_kwargs)
# Distributed backpropagation
loss.backward()
optimizer.step()
if step % 100 == 0:
flush() # From dist_utils.py, ensures proper synchronization
The script handles global_world_size management and gradient accumulation, supporting both FP4 and BF16 precision for memory-efficient training of the 2.6B parameter architecture.
Key Source Files and Implementation Details
When implementing or modifying SANA-WM, reference these specific files:
app/sana_pipeline.py: ContainsSanaWMPipeline, the high-level wrapper exposing thegeneratemethod with camera pose argumentsdiffusion/model/nets/basic_modules.py: HousesDWMlp,GLUMBConv, andGLUMBConvTemp—the core blocks for spatio-temporal mixingdiffusion/model/nets/sana.py: Implements the main DiT backbone with linear attention and block-causal maskingdiffusion/model/nets/sana_multi_scale_adaln.py: Contains the pose encoder MLP used during trainingdiffusion/utils/dist_utils.py: Providesget_world_size()andflush()utilities for distributed training coordination
Summary
- SANA-WM adds 6-DoF camera control to the SANA family via a lightweight pose-injection mechanism
- The
DWMlpblock inbasic_modules.pyenables efficient spatio-temporal mixing through depth-wise convolutions - Camera poses are provided as
(frames, 6)tensors and processed automatically by theSanaWMPipeline - The 2.6B parameter model achieves 720p, 1-minute video generation on a single RTX 3090 thanks to Linear-attention DiT and 32× DC-AE compression
- Training scripts in
train_video_scripts/demonstrate proper distributed training with pose data using utilities fromdist_utils.py
Frequently Asked Questions
What hardware is required to run SANA-WM inference?
SANA-WM can generate 720p videos on a single NVIDIA RTX 3090 GPU (24GB VRAM) using enable_model_cpu_offload() to manage memory. For training the full 2.6B parameter model, multi-GPU setups are recommended, though the efficient Linear-attention and DC-AE compression reduce requirements compared to standard DiT architectures.
How does the camera pose conditioning work technically?
The system accepts sequences of 6-DoF vectors—representing 3D position and Euler angles—and projects them through a small MLP (DWMlp). These embeddings are injected into the token stream before the depth-wise convolution layers in basic_modules.py, allowing temporal mixing to propagate pose information across all spatial tokens while maintaining the model's linear computational complexity.
Can I fine-tune SANA-WM on custom video datasets with camera trajectories?
Yes. Use the train_video_ivjoint.py script in train_video_scripts/, ensuring your dataloader provides batches with both pixel_values (video frames) and camera_poses (6-dimensional tensors). The script handles distributed training automatically via get_world_size() and supports gradient accumulation for effective batch size management.
What is the difference between standard SANA and SANA-WM?
Standard SANA focuses on efficient image generation using Linear-attention DiT and DC-AE compression. SANA-WM extends this architecture for video by adding the DWMlp temporal mixing blocks and camera-pose conditioning branches, enabling minute-long video generation with explicit camera control while preserving the original model's memory efficiency.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →