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 totrueto enable the FSDP backendFSDP_AUTO_WRAP_POLICY: Typically set toTRANSFORMER_BASED_WRAPfor transformer architecturesFSDP_SHARDING_STRATEGY: Set toFULL_SHARDfor maximum memory efficiencyFSDP_REDUCE_SCATTER_PRECISION: Set tofp32for 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:
-
Parse Configuration:
pyrallisloads your YAML into aSanaConfigobject and reads thetrain.use_fsdpboolean. -
Initialize Environment: If FSDP is enabled,
set_fsdp_env()writes the required environment variables before distributed initialization. -
Create Process Group:
Acceleratorinitializes the distributed process group viaInitProcessGroupKwargs. -
Instantiate Plugin: The
FullyShardedDataParallelPluginis configured withFullStateDictConfigfor checkpoint compatibility. -
Build and Wrap Model:
build_model()returns the raw model, whichaccelerator.prepare()automatically wraps with FSDP according to the plugin settings. -
Execute Training Loop: The training proceeds normally, with
save_checkpoint_fsdpandload_checkpoint_fsdpindiffusion/utils/checkpoint.pyhandling 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
-
train_scripts/train.py: Main entry point containingset_fsdp_env()(lines 58-86) and the primary training loop with optional FSDP plugin creation. -
train_scripts/train_scm_ladd.py: Demonstrates explicitFullyShardedDataParallelPlugininstantiation with customFullStateDictConfig(lines 88-96). -
diffusion/utils/config.py: Defines theTrainingConfig.use_fsdpboolean flag (lines 216-218) consumed by all training scripts. -
diffusion/model/wan/fsdp_utils.py: Containsshard_model()helper (lines 12-33) for direct FSDP wrapping using PyTorch's native API. -
diffusion/utils/checkpoint.py: Implements FSDP-aware checkpoint management withsave_checkpoint_fsdp()andload_checkpoint_fsdp(). -
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: truein your YAML config and ensuringset_fsdp_env()populates the required environment variables beforeAcceleratorinitialization. -
Memory efficiency comes from
FULL_SHARDstrategy, 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 monolithicpytorch_model_fsdp.binfiles that can be loaded without FSDP for inference. -
Granular control is available via
shard_model()infsdp_utils.pyfor wrapping specific submodules like text encoders independently of the main model wrapper. -
Production validation includes the
tests/bash/training/test_training_fsdp.shscript, 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →