Implementing Block Causal Linear Attention for Long Video Generation with Sana
Block Causal Linear Attention reduces video generation complexity from quadratic to linear by processing frames in causal chunks using kernel-based prefix sums, enabling Sana to generate videos with thousands of frames.
The NVlabs/Sana repository implements this mechanism to overcome the memory and computational bottlenecks of standard transformers when handling extended video sequences. By leveraging a ReLU-based kernel transformation and prefix-sum algebra in diffusion/model/nets/sana_blocks.py, Sana achieves O(N) complexity per layer while maintaining strict temporal causality across frame chunks.
What is Block Causal Linear Attention?
Standard self-attention scales quadratically with sequence length, consuming prohibitive memory for long videos. Block Causal Linear Attention solves this by combining kernel-based linearization with chunk-wise causal masking.
The mechanism rewrites attention as φ(q)·φ(k)ᵀ where φ represents an element-wise ReLU kernel. This transformation eliminates explicit softmax computation, reducing complexity from O(N²) to O(N) where N represents the token count (video frames).
Chunk-Wise Causality
Video frames are processed in fixed-size chunks (typically 8 or 16 frames). Within each chunk, attention follows causal constraints: each token attends only to previous tokens in the sequence. The linear kernel formulation enables efficient prefix-sum computation that respects these causal boundaries without explicit attention mask matrices.
Architecture and Implementation Details
The implementation centers on the ChunkCausalAttention class located at line 391 of diffusion/model/nets/sana_blocks.py.
Core Components in sana_blocks.py
ChunkCausalAttention inherits from LiteLAReLURope, providing base linear attention machinery with ReLU activation and Rotary Position Embeddings (RoPE). The class structure handles the dimension splitting required for multi-head linear attention:
# Conceptual structure from diffusion/model/nets/sana_blocks.py
class ChunkCausalAttention(LiteLAReLURope):
def __init__(self, in_dim, out_dim, heads, dim, eps=1e-8, ...):
super().__init__()
self.dim = dim # Per-head dimension (e.g., 32)
self.heads = heads # Number of linear heads
self.kernel_func = nn.ReLU() # Linearizing kernel φ
The Linear Attention Kernel
The kernel transformation defined at lines 350-380 of diffusion/model/nets/sana_blocks.py applies the ReLU function to Query and Key tensors before computing attention weights. The algebraic reformulation allows the attention output to be computed via matrix multiplications that scale linearly with sequence length:
- Input tensor
xof shape(B, N, C)projects to Q/K/V via learned weights - Reshape to
(B, h, h_d, N)whereh = C // dimrepresents head count - Apply
self.kernel_func(ReLU) to Q and K, ensuring positive embeddings - Compute normalized output via prefix sums:
# Simplified from the implementation logic
z = 1 / (k.sum(dim=-1, keepdim=True).transpose(-2, -1) @ q + self.eps)
vk = v @ k_rotated.transpose(-2, -1)
out = vk @ q_rotated
output = out * z # Linear attention normalization
Rotary Position Embeddings
The implementation applies apply_rotary_emb to Q and K tensors immediately after the kernel transformation (lines 350-380 of sana_blocks.py). This preserves temporal positional information exactly as implemented in Flash Attention, ensuring frame ordering is maintained despite the linear approximation.
Integration into Sana's Video Backbone
The attention block is instantiated through the SanaVideoMSBlock class in diffusion/model/nets/sana_multi_scale_video.py. At lines 99-104, the constructor selects ChunkCausalAttention when attn_type="chunkcausal" is specified:
# From diffusion/model/nets/sana_multi_scale_video.py
if attn_type == "chunkcausal":
self.attn = ChunkCausalAttention(
in_dim=hidden_size,
out_dim=hidden_size,
heads=hidden_size // linear_head_dim,
dim=linear_head_dim,
qk_norm=qk_norm,
# Additional configuration...
)
The linear_head_dim parameter determines the per-head dimension, typically set to 32 or 64, while hidden_size // linear_head_dim calculates the required head count automatically.
Practical Implementation Examples
Configuring Block Causal Attention via SanaVideoMSBlock
To utilize this mechanism in custom video generation pipelines:
import torch
from diffusion.model.nets.sana_multi_scale_video import SanaVideoMSBlock
# Initialize video block with block-causal linear attention
video_block = SanaVideoMSBlock(
hidden_size=512,
num_heads=8,
attn_type="chunkcausal", # Selects ChunkCausalAttention
linear_head_dim=32, # Dimension per linear head
qk_norm=True, # Enable Query/Key normalization
)
# Process video tokens: (batch, tokens, channels)
x = torch.randn(2, 64, 512) # 2 videos, 64 tokens (frames), 512 channels
output = video_block(x) # Forward pass uses linear attention
print(output.shape) # torch.Size([2, 64, 512])
Direct Instantiation of ChunkCausalAttention
For fine-grained control over the attention mechanism:
from diffusion.model.nets.sana_blocks import ChunkCausalAttention
import torch
batch, seq_len, dim = 2, 128, 512
heads = dim // 32 # 16 heads with 32-dim each
attn = ChunkCausalAttention(
in_dim=dim,
out_dim=dim,
heads=heads,
dim=32,
eps=1e-8,
use_bias=False,
qk_norm=True,
)
x = torch.randn(batch, seq_len, dim)
rotary_emb = torch.randn(1, 1, 32, seq_len) # RoPE tensor
output = attn(x, rotary_emb=rotary_emb)
assert output.shape == (batch, seq_len, dim)
Switching Between Attention Types
Sana supports multiple attention backends configurable via the attn_type parameter:
def build_video_block(attention_type: str):
"""Factory for different attention mechanisms."""
return SanaVideoMSBlock(
hidden_size=768,
num_heads=12,
attn_type=attention_type, # "flash", "chunkcausal", "linear", etc.
linear_head_dim=32,
)
# Standard Flash Attention (quadratic complexity)
flash_block = build_video_block("flash")
# Block Causal Linear Attention (linear complexity)
causal_linear_block = build_video_block("chunkcausal")
Memory Efficiency and Numerical Stability
The implementation includes safeguards for training stability. The fp32_attention option automatically casts intermediate tensors to FP32 precision to prevent NaN values during gradient computation. Optional qkv_store_buffer parameters allow storage of intermediate Q/K/V activations for visualization or analysis without disrupting the streaming pipeline, enabling efficient debugging of long-video generation sequences.
Summary
- Block Causal Linear Attention replaces quadratic self-attention with a kernel-based O(N) operation using ReLU transformations and prefix-sum algebra.
- The
ChunkCausalAttentionclass indiffusion/model/nets/sana_blocks.py(line 391) implements this mechanism, inheriting fromLiteLAReLURopefor rotary embedding support. - Causality is enforced through prefix-sum ordering rather than explicit masking, processing video frames in configurable fixed-size chunks.
- Integration occurs via
attn_type="chunkcausal"inSanaVideoMSBlockatdiffusion/model/nets/sana_multi_scale_video.py(lines 99-104). - The architecture supports FP32 fallback and memory-efficient buffering (
qkv_store_buffer) for stable training on long-video generation tasks.
Frequently Asked Questions
How does Block Causal Linear Attention differ from standard Flash Attention?
Flash Attention reduces memory usage through tiling and recomputation but maintains O(N²) computational complexity. Block Causal Linear Attention achieves O(N) complexity by approximating softmax attention with a ReLU kernel decomposition, trading some expressivity for linear scaling with sequence length. Both mechanisms use rotary embeddings, but only the linear variant efficiently processes videos containing thousands of frames on standard GPU hardware.
What chunk size should I use for video generation?
The optimal chunk size depends on your GPU memory constraints and temporal coherence requirements. The implementation typically uses chunks of 8 or 16 frames as balanced defaults. Smaller chunks reduce memory usage and increase parallelism, while larger chunks may improve temporal consistency across frame boundaries. Configure this through the SanaVideoMSBlock initialization parameters in diffusion/model/nets/sana_multi_scale_video.py.
Can Block Causal Linear Attention be used for image generation?
While technically possible, this attention mechanism is optimized for sequential video data where temporal causality matters. For single-image generation, standard Flash Attention or non-causal linear attention typically performs better because images lack the strict temporal ordering constraints that chunk-wise causality enforces.
Where is the attention mechanism configured in the Sana training pipeline?
The attention type is selected in diffusion/model/nets/sana_multi_scale_video.py (lines 99-104) through the attn_type constructor argument. During inference, this configuration is loaded from model checkpoints via app/sana_pipeline.py, which instantiates the appropriate SanaVideoMSBlock classes based on saved hyperparameters in the checkpoint file.
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 →