# How to Train Sana with FSDP: A Complete Guide to Fully Sharded Data Parallel

> Learn how to train Sana with FSDP to fit 2B parameter video models on two GPUs. This guide covers sharding parameters, gradients, and optimizer states.

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

---

**Sana supports Fully Sharded Data Parallel (FSDP) through PyTorch's native distributed training, enabling the 2B-parameter video models to fit on just two GPUs by sharding parameters, gradients, and optimizer states across devices.**

The NVlabs/Sana repository implements FSDP integration through the Hugging Face Accelerate framework, allowing you to train massive text-to-image and text-to-video models without requiring enterprise-grade hardware. By activating FSDP via configuration flags and environment variables, you can distribute memory overhead across multiple GPUs while maintaining computational efficiency.

## What Is FSDP and Why Use It for Sana?

FSDP is a PyTorch-native parallelism strategy that shards model parameters, gradients, and optimizer states across all participating GPUs. For Sana's 2B-parameter **SanaMSVideo_2000M_P2_D20** model, this reduces per-GPU memory consumption by roughly **1 / world_size**, making it feasible to train with batch sizes of 1 on a 2-GPU node. Unlike Data Parallel (DP), which replicates the entire model on each device, FSDP keeps only a slice of parameters locally, dramatically reducing memory pressure during training.

## Enabling FSDP in Sana Training

### Configuration Files

The primary mechanism for activating FSDP is the **`use_fsdp`** flag in your training configuration. In [`diffusion/utils/config.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/config.py) (lines 216-218), the `TrainingConfig` class defines this boolean flag, which defaults to `false`.

To enable FSDP, use a YAML configuration that explicitly sets the flag, such as [`configs/sana_video_config/Sana_2000M_480px_AdamW_fsdp.yaml`](https://github.com/NVlabs/Sana/blob/main/configs/sana_video_config/Sana_2000M_480px_AdamW_fsdp.yaml) (lines 107-108):

```yaml
train:
  use_fsdp: true
  mixed_precision: "bf16"

```

### Environment Setup via set_fsdp_env()

When `use_fsdp` is set to `true`, the training scripts invoke `set_fsdp_env()` to populate environment variables required by Accelerate's FSDP integration. Located in [`train_scripts/train.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train.py) (lines 58-86), this function configures:

- **`ACCELERATE_USE_FSDP`**: Set to `true` to enable the FSDP backend
- **`FSDP_AUTO_WRAP_POLICY`**: Typically set to `TRANSFORMER_BASED_WRAP` for transformer architectures
- **`FSDP_SHARDING_STRATEGY`**: Set to `FULL_SHARD` for maximum memory efficiency
- **`FSDP_REDUCE_SCATTER_PRECISION`**: Set to `fp32` for numerical stability during gradient reduction

### Accelerator Plugin Configuration

The training scripts instantiate an `Accelerator` with a `FullyShardedDataParallelPlugin` when FSDP mode is active. In [`train_scripts/train_scm_ladd.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train_scm_ladd.py) (lines 88-96), the plugin is created with a `FullStateDictConfig` that disables CPU offloading and keeps the full state on each rank:

```python
from accelerate import FullyShardedDataParallelPlugin
from torch.distributed.fsdp.fully_sharded_data_parallel import FullStateDictConfig

fsdp_plugin = FullyShardedDataParallelPlugin(
    state_dict_config=FullStateDictConfig(offload_to_cpu=False, rank0_only=False)
)

accelerator = Accelerator(
    mixed_precision="bf16",
    fsdp_plugin=fsdp_plugin,
)

```

## End-to-End FSDP Training Workflow

The complete FSDP initialization follows this sequence:

1. **Parse Configuration**: `pyrallis` loads your YAML into a `SanaConfig` object and reads the `train.use_fsdp` boolean.

2. **Initialize Environment**: If FSDP is enabled, `set_fsdp_env()` writes the required environment variables before distributed initialization.

3. **Create Process Group**: `Accelerator` initializes the distributed process group via `InitProcessGroupKwargs`.

4. **Instantiate Plugin**: The `FullyShardedDataParallelPlugin` is configured with `FullStateDictConfig` for checkpoint compatibility.

5. **Build and Wrap Model**: `build_model()` returns the raw model, which `accelerator.prepare()` automatically wraps with FSDP according to the plugin settings.

6. **Execute Training Loop**: The training proceeds normally, with `save_checkpoint_fsdp` and `load_checkpoint_fsdp` in [`diffusion/utils/checkpoint.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/checkpoint.py) handling model state persistence.

## Manual Model Sharding for Text Encoders

For specific components like the T5-style text encoder in the WanVAE architecture, Sana provides a low-level alternative to the generic Accelerate plugin. The `shard_model()` helper in [`diffusion/model/wan/fsdp_utils.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/wan/fsdp_utils.py) (lines 12-33) directly wraps sub-modules with `torch.distributed.fsdp.FullyShardedDataParallel`:

```python
from diffusion.model.wan.fsdp_utils import shard_model

# Directly wrap a specific submodule with FSDP

shard_model(text_encoder, auto_wrap_policy=transformer_based_wrap)

```

This approach gives you granular control over which model components participate in sharding, useful when you want to keep certain layers (like frozen text encoders) unsharded for performance.

## Practical Code Examples

### Minimal FSDP Setup Script

To manually configure FSDP in a custom training script:

```python
from accelerate import Accelerator, FullyShardedDataParallelPlugin
from torch.distributed.fsdp.fully_sharded_data_parallel import FullStateDictConfig
import os

# Configure environment

os.environ["ACCELERATE_USE_FSDP"] = "true"
os.environ["FSDP_AUTO_WRAP_POLICY"] = "TRANSFORMER_BASED_WRAP"
os.environ["FSDP_SHARDING_STRATEGY"] = "FULL_SHARD"
os.environ["FSDP_REDUCE_SCATTER_PRECISION"] = "fp32"

# Initialize plugin

fsdp_plugin = FullyShardedDataParallelPlugin(
    state_dict_config=FullStateDictConfig(offload_to_cpu=False, rank0_only=False)
)

# Create accelerator

accelerator = Accelerator(
    mixed_precision="bf16",
    fsdp_plugin=fsdp_plugin,
)

# Prepare model (automatically wraps with FSDP)

model = build_model(cfg.model)
model = accelerator.prepare(model)

```

### Command-Line Training Launch

To train the 480p video model with FSDP on 2 GPUs:

```bash
bash train_video_scripts/train_video_ivjoint.sh \
    configs/sana_video_config/Sana_2000M_480px_AdamW_fsdp.yaml \
    --train.num_epochs=10 \
    --train.train_batch_size=1 \
    --train.use_fsdp=true \
    --np=2

```

The YAML configuration already contains `use_fsdp: true`, but you can override it via command-line arguments. The script automatically calls `set_fsdp_env()` before initializing the `Accelerator`.

### Loading FSDP Checkpoints for Inference

To resume training or perform inference from an FSDP checkpoint:

```python
from diffusion.utils.checkpoint import load_checkpoint_fsdp

model, optimizer, scheduler = load_checkpoint_fsdp(
    checkpoint_dir="output/sana_video/checkpoint",
    model=model,
    optimizer=optimizer,
    load_ema=False
)

```

This utility reconstructs the full parameter set from the sharded `pytorch_model_fsdp.bin` file, making checkpoints portable between FSDP and non-FSDP environments.

## Key Files and Implementation Details

- **[`train_scripts/train.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train.py)**: Main entry point containing `set_fsdp_env()` (lines 58-86) and the primary training loop with optional FSDP plugin creation.

- **[`train_scripts/train_scm_ladd.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train_scm_ladd.py)**: Demonstrates explicit `FullyShardedDataParallelPlugin` instantiation with custom `FullStateDictConfig` (lines 88-96).

- **[`diffusion/utils/config.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/config.py)**: Defines the `TrainingConfig.use_fsdp` boolean flag (lines 216-218) consumed by all training scripts.

- **[`diffusion/model/wan/fsdp_utils.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/wan/fsdp_utils.py)**: Contains `shard_model()` helper (lines 12-33) for direct FSDP wrapping using PyTorch's native API.

- **[`diffusion/utils/checkpoint.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/checkpoint.py)**: Implements FSDP-aware checkpoint management with `save_checkpoint_fsdp()` and `load_checkpoint_fsdp()`.

- **[`configs/sana_video_config/Sana_2000M_480px_AdamW_fsdp.yaml`](https://github.com/NVlabs/Sana/blob/main/configs/sana_video_config/Sana_2000M_480px_AdamW_fsdp.yaml)**: Reference configuration for 480p video training with FSDP enabled.

## Summary

- **FSDP activation** requires setting `use_fsdp: true` in your YAML config and ensuring `set_fsdp_env()` populates the required environment variables before `Accelerator` initialization.

- **Memory efficiency** comes from `FULL_SHARD` strategy, which distributes parameters, gradients, and optimizer states across GPUs, reducing per-device memory by the world size factor.

- **Checkpoint compatibility** is maintained through `FullStateDictConfig`, producing monolithic `pytorch_model_fsdp.bin` files that can be loaded without FSDP for inference.

- **Granular control** is available via `shard_model()` in [`fsdp_utils.py`](https://github.com/NVlabs/Sana/blob/main/fsdp_utils.py) for wrapping specific submodules like text encoders independently of the main model wrapper.

- **Production validation** includes the [`tests/bash/training/test_training_fsdp.sh`](https://github.com/NVlabs/Sana/blob/main/tests/bash/training/test_training_fsdp.sh) script, which verifies 2B-parameter video model training on 2-GPU nodes.

## Frequently Asked Questions

### What exactly does FSDP shard during Sana training?

According to the NVlabs/Sana source code, FSDP shards **parameters**, **gradients**, and **optimizer states** across all participating GPUs using the `FULL_SHARD` strategy. This means each GPU stores only 1/N of these tensors (where N is the world size), significantly reducing memory overhead compared to Data Parallel training.

### Can I train the 2B parameter Sana video model on a single GPU?

No, the **SanaMSVideo_2000M_P2_D20** model requires FSDP with at least 2 GPUs to fit in memory with a batch size of 1, as demonstrated in [`tests/bash/training/test_training_fsdp.sh`](https://github.com/NVlabs/Sana/blob/main/tests/bash/training/test_training_fsdp.sh) and the [`Sana_2000M_480px_AdamW_fsdp.yaml`](https://github.com/NVlabs/Sana/blob/main/Sana_2000M_480px_AdamW_fsdp.yaml) configuration. Single-GPU training would require gradient checkpointing and other memory optimization techniques not covered by the standard FSDP implementation.

### How do I resume training from an FSDP checkpoint?

Use the `load_checkpoint_fsdp()` function from [`diffusion/utils/checkpoint.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/checkpoint.py), passing your checkpoint directory and model instance. This function automatically detects the `pytorch_model_fsdp.bin` file and reconstructs the full state dict across all ranks, allowing seamless resumption of distributed training.

### What's the difference between using the Accelerator plugin versus manual shard_model()?

The **`FullyShardedDataParallelPlugin`** (used in [`train_scm_ladd.py`](https://github.com/NVlabs/Sana/blob/main/train_scm_ladd.py)) provides automatic wrapping of the entire model through Accelerate's `prepare()` method, suitable for the DiT (Diffusion Transformer) backbone. In contrast, **`shard_model()`** in [`diffusion/model/wan/fsdp_utils.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/wan/fsdp_utils.py) offers low-level control for wrapping specific sub-modules like the T5 text encoder, allowing you to exclude frozen parameters from sharding to optimize communication overhead.