# Implementing Block Causal Linear Attention for Long Video Generation with Sana

> Implement Block Causal Linear Attention for long video generation with Sana. This technique reduces complexity to linear, enabling generation of thousands of frames.

- Repository: [NVIDIA Research Projects/Sana](https://github.com/NVlabs/Sana)
- Tags: how-to-guide
- Published: 2026-05-19

---

**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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/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:

```python

# 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`](https://github.com/NVlabs/Sana/blob/main/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:

1. Input tensor `x` of shape `(B, N, C)` projects to Q/K/V via learned weights
2. Reshape to `(B, h, h_d, N)` where `h = C // dim` represents head count
3. Apply `self.kernel_func` (ReLU) to Q and K, ensuring positive embeddings
4. Compute normalized output via prefix sums:

```python

# 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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/nets/sana_multi_scale_video.py). At lines 99-104, the constructor selects `ChunkCausalAttention` when `attn_type="chunkcausal"` is specified:

```python

# 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:

```python
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:

```python
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:

```python
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 `ChunkCausalAttention` class in [`diffusion/model/nets/sana_blocks.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/nets/sana_blocks.py) (line 391) implements this mechanism, inheriting from `LiteLAReLURope` for 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"` in `SanaVideoMSBlock` at [`diffusion/model/nets/sana_multi_scale_video.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/app/sana_pipeline.py), which instantiates the appropriate `SanaVideoMSBlock` classes based on saved hyperparameters in the checkpoint file.