# LTX-2 Attention Backend Options: FlashAttention, PyTorch SDPA, and NADiffusion Decoder Explained

> Explore LTX-2 attention backend options including FlashAttention, PyTorch SDPA, and NADiffusion. Discover efficient choices for your models with hardware detection.

- Repository: [Lightricks/LTX-2](https://github.com/Lightricks/LTX-2)
- Tags: deep-dive
- Published: 2026-08-15

---

**LTX-2 supports FlashAttention 3/4, PyTorch SDPA, and MPS-SDPA with automatic hardware detection, plus a Neighborhood-Attention-based NADiffusion decoder for video generation.**

The LTX-2 video generation model from Lightricks provides a flexible, hardware-aware attention subsystem. This article breaks down every backend option, the automatic selection logic, and how the NADiffusion decoder integrates these primitives for high-performance video decoding.

---

## Understanding LTX-2 Attention Backends

The attention layer in [`ltx_core/model/transformer/attention.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/transformer/attention.py) implements multiple compute paths optimized for different hardware. Each backend trades off between raw speed, memory efficiency, and hardware availability.

### FlashAttention 3: Proven CUDA Performance

**FlashAttention 3** uses the original `flash_attn` CUDA kernel for unmasked attention. It avoids materializing the full attention matrix, reducing memory bandwidth and enabling longer sequences on NVIDIA GPUs.

Selection criteria:
- Requires CUDA-capable GPU
- Preferred on Ampere and Hopper architectures (sm_80, sm_90)
- Falls back automatically if kernel compilation fails

### FlashAttention 4: The "Cute" Implementation

**FlashAttention 4** (`flash_attn.cute`) provides an updated kernel API with equivalent speed to FA3. The LTX-2 codebase treats these as distinct options in the `AttentionFunction` enum, allowing explicit version pinning:

```python
from ltx_core.model.transformer.attention import AttentionFunction

# Force specific FlashAttention version

fa3 = AttentionFunction.FLASH_ATTENTION_3.to_callable()
fa4 = AttentionFunction.FLASH_ATTENTION_4.to_callable()

```

GPU architecture detection in `_select_primary_attention()` (lines 84-96) influences default preference—Hopper (sm_90) and newer may prefer FA4 when available.

### PyTorch SDPA: Cross-Platform Compatibility

**PyTorch SDPA** (`torch.nn.functional.scaled_dot_product_attention`) serves as the universal fallback. LTX-2 configures SDPA with a strict priority ordering:

```python
_SDPA_FULL_PRIORITY = (
    SDPBackend.CUDNN_ATTENTION,      # NVIDIA cuDNN fused kernel

    SDPBackend.FLASH_ATTENTION,      # torch.internal flash

    SDPBackend.EFFICIENT_ATTENTION,  # memory-efficient CUDA kernel

    SDPBackend.MATH,                 # pure PyTorch (always works)

)

```

`SDPBackend.CUDNN_ATTENTION` receives top priority because it often outperforms generic FlashAttention on NVIDIA hardware. The `sdpa_kernel(..., set_priority=True)` context manager applies this ordering at runtime.

### MPS-SDPA: Apple Silicon Optimization

**MPS-SDPA** (`mps-sdpa` package) targets Apple Silicon (M1/M2/M3/M4). It avoids materializing the full score matrix for long video sequences, yielding approximately **30× speedups** on M4 Pro compared to naive attention.

Detection logic:
- `_on_macos()` checks `sys.platform == "darwin"`
- Attempts `import mps_sdpa` to confirm package availability
- Automatically selected when both conditions pass

---

## Automatic Backend Selection

The `automatic_attention()` and `automatic_masked_attention()` functions implement zero-config backend selection. Both use `functools.cache` to guarantee singleton behavior per process.

### Selection Hierarchy

1. **Apple Silicon path**: If `_on_macos()` and `mps-sdpa` importable → `MPS_SDPA`
2. **CUDA path**: Query `torch.cuda.get_device_capability()`
   - Blackwell (sm_100+): Prioritize `FLASH_ATTENTION_4`
   - Hopper/Ampere (sm_80-sm_99): Prioritize `FLASH_ATTENTION_3`
   - Fallback: `SDPA_CUDNN` → `SDPA_FLASH` → `SDPA_EFFICIENT` → `SDPA_MATH`
3. **CPU-only**: Pure `SDPA_MATH`

```python
from ltx_core.model.transformer.attention import automatic_attention

# Hardware-determined backend

attention = automatic_attention()  # Cached singleton

output = attention(q, k, v, heads=8)  # Unmasked

```

### Masked Attention Variant

Causal or variable-length models require `automatic_masked_attention()`:

```python
from ltx_core.model.transformer.attention import automatic_masked_attention

masked_attn = automatic_masked_attention()
output = masked_attn(q, k, v, heads=8, mask=attention_mask)

```

The masked path follows identical hardware detection but selects backends supporting additive masking (SDPA variants, not raw FlashAttention).

---

## NADiffusion Decoder: Attention in Practice

The **NADiffusion decoder** ([`ltx_core/model/video_vae/diffusion_video_decoder.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/video_vae/diffusion_video_decoder.py)) demonstrates production integration of LTX-2 attention primitives. This video-VAE decoder converts latent tensors to pixel-space video using Neighborhood-Attention blocks.

### Architecture Overview

| Stage | Function | Channels | Kernel |
|:---|:---|:---|:---|
| 1-4 | Deterministic upsampling | (1024, 512, 256, 256) | (3,7,7), (3,5,5), (3,3,3), (3,3,3) |
| 5 | Diffusion-aware generation | 128 | (3,3,3) |

Stage 5 contains `DiffusionNABlock` modules that consume the attention backends. Three block variants exist:

- **`CombinedDiffusionNABlock`** (default): Full-volume context + residual MLP
- **`ChunkedDiffusionNABlock`**: Memory-constrained chunked processing
- **`DSLDiffusionNABlock`**: Fused CuTe DSL kernels for maximum inference speed

### Construction and Usage

```python
import torch
from ltx_core.model.video_vae.diffusion_video_decoder import DiffusionVideoDecoder

decoder = DiffusionVideoDecoder(
    in_channels=128,
    out_channels=3,
    patch_size=4,
    head_dim=64,
    default_num_inference_steps=2,
    model_output_type="v",  # v-parameterization

)

latent = torch.randn(1, 128, 16, 64, 64, device="cuda")  # [B,C,T,H,W]

video = decoder(latent)  # → [1, 3, 16, 256, 256]

```

### Tiling and Dynamic Shapes

The decoder handles large videos via `diffusion_tiling` utilities with automatic halo padding. Torch-compilation compatibility is ensured through `self.mark_dynamic_shapes()`, which preserves shape flexibility across batch sizes and resolutions.

---

## Explicit Backend Selection Examples

### Research Reproducibility

Force deterministic SDPA math kernel (slow but identical across platforms):

```python
from ltx_core.model.transformer.attention import AttentionFunction

math_attn = AttentionFunction.SDPA_MATH.to_callable()

```

### Maximum Performance on Known Hardware

Explicit FlashAttention 4 for Blackwell validation:

```python
if torch.cuda.get_device_capability() >= (10, 0):
    attn = AttentionFunction.FLASH_ATTENTION_4.to_callable()
else:
    raise RuntimeError("FA4 requires Blackwell+")

```

### Runtime Backend Inspection

```python
from ltx_core.model.transformer.attention import automatic_attention

attn_fn = automatic_attention()
print(attn_fn.__name__)  # Reveals selected backend: "flash_attention_4", "mps_sdpa", etc.

```

---

## Key Source Files Reference

| Path | Purpose |
|:---|:---|
| [`ltx_core/model/transformer/attention.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/transformer/attention.py) | `AttentionFunction` enum, backend selection, `automatic_attention()`, FlashAttention 3/4 kernels, MPS-SDPA wrapper |
| [`ltx_core/model/video_vae/diffusion_video_decoder.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/video_vae/diffusion_video_decoder.py) | `DiffusionVideoDecoder` class, 5-stage NA architecture, tiling integration |
| [`ltx_core/model/video_vae/transformer/blocks.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/video_vae/transformer/blocks.py) | Base `NABlock` and `DiffusionNABlock` definitions |
| [`ltx_core/model/video_vae/transformer/combined/block.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/video_vae/transformer/combined/block.py) | `CombinedDiffusionNABlock` (full-context default) |
| [`ltx_core/model/video_vae/transformer/dsl_kernels/block.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_core/model/video_vae/transformer/dsl_kernels/block.py) | `DSLDiffusionNABlock` (fused CuTe implementation) |

---

## Summary

- **Automatic selection** via `automatic_attention()` examines GPU architecture, CUDA availability, and Apple Silicon presence to choose between FlashAttention 3/4, PyTorch SDPA variants, or MPS-SDPA
- **Explicit control** through `AttentionFunction.to_callable()` enables reproducible research and hardware-specific optimization
- **PyTorch SDPA priority** orders CUDNN_ATTENTION > FLASH_ATTENTION > EFFICIENT_ATTENTION > MATH for broad compatibility
- **MPS-SDPA** delivers ~30× Apple Silicon speedups by avoiding full attention matrix materialization
- **NADiffusion decoder** integrates these primitives into a production video-VAE with 5-stage upsampling, configurable attention blocks, and torch-compile support

---

## Frequently Asked Questions

### How does LTX-2 choose between FlashAttention 3 and 4?

LTX-2 inspects `torch.cuda.get_device_capability()` in `_select_primary_attention()`. Blackwell GPUs (sm_100+) prefer FlashAttention 4, while Ampere and Hopper default to FlashAttention 3. Both require successful package import; if unavailable, the selector falls through to PyTorch SDPA.

### Can I use LTX-2 attention backends on CPU or Apple Silicon?

Yes. CPU-only systems automatically receive `SDPA_MATH`. Apple Silicon with the `mps-sdpa` package installed triggers `MPS_SDPA` selection, which provides substantial speedups over generic SDPA for long video sequences.

### What is the NADiffusion decoder used for?

The NADiffusion decoder converts compressed latent tensors into full-resolution video frames. It combines deterministic upsampling stages (1-4) with diffusion-aware attention blocks in stage 5, supporting variable resolutions through tiling and dynamic shape handling for production video generation pipelines.

### How do I force a specific attention backend for reproducibility?

Import `AttentionFunction` and call `.to_callable()` on any enum member: `AttentionFunction.SDPA_MATH.to_callable()` for pure PyTorch, `AttentionFunction.FLASH_ATTENTION_3.to_callable()` for specific CUDA kernels, etc. This bypasses automatic detection entirely.