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

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 (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 (lines 107-108):

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 (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 (lines 88-96), the plugin is created with a FullStateDictConfig that disables CPU offloading and keeps the full state on each rank:

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 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 (lines 12-33) directly wraps sub-modules with torch.distributed.fsdp.FullyShardedDataParallel:

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:

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

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

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 for wrapping specific submodules like text encoders independently of the main model wrapper.

  • Production validation includes the 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 and the 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, 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) 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 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →