How to Configure Mixed Precision Training (BF16/FP16) in Sana
Enable BF16 or FP16 training in the NVlabs/Sana repository by setting mixed_precision: bf16 (or fp16) in your YAML configuration file, which the Hugging Face Accelerator uses to enable automatic mixed precision (AMP), while the get_weight_dtype() utility casts model weights and encoders to the corresponding PyTorch dtype.
Sana supports mixed precision training to reduce GPU memory usage and accelerate computation on modern hardware. This guide explains how the training pipeline handles precision configuration, from YAML declarations to low-level dtype casting, based on the actual implementation in the NVlabs/Sana source code.
Where Precision is Configured in Sana
The repository implements mixed precision through three coordinated layers: configuration files, the Accelerate library integration, and utility functions that map string names to PyTorch dtypes.
YAML Configuration Files
Every model configuration in configs/ contains a mixed_precision entry that declares the training dtype. For example, the Sol-RL configuration explicitly sets BF16:
# configs/sol_rl/Sana1.0_1600M_linear.yaml
model:
mixed_precision: bf16
fp32_attention: true
This value is parsed by SanaConfig in diffusion/utils/config.py (line 69) and passed throughout the training pipeline.
Accelerate Integration
The training scripts instantiate an Accelerator object that reads the precision string directly from the config. In train_scripts/train.py (line 662), this enables torch-autocast automatically:
accelerator = Accelerator(mixed_precision=config.model.mixed_precision)
This single line activates automatic mixed precision for the entire training loop without manual autocast context managers.
The get_weight_dtype Utility
The get_weight_dtype function in diffusion/model/utils.py (line 598) converts configuration strings into concrete PyTorch dtypes:
from diffusion.model.utils import get_weight_dtype
weight_dtype = get_weight_dtype(config.model.mixed_precision)
# Returns torch.bfloat16 for "bf16", torch.float16 for "fp16", torch.float32 for "fp32"
All pipelines—training, inference, and ControlNet—use this helper to cast model weights, VAE, and text encoders to the appropriate precision.
How Components Work Together
The following flow illustrates how the precision setting propagates from configuration to execution:
- Config file (
*.yaml): Declaresmixed_precision(e.g.,bf16). Parsed bySanaConfigatdiffusion/utils/config.py. - Accelerator: Created with the precision flag, enabling AMP and selecting the appropriate autocast dtype.
- get_weight_dtype: Returns the concrete
torch.dtypeused to cast models and tensors. - Training loops: Scripts like
train_scripts/sol_rl/train_sana.py(line 441) check the dtype to conditionally enableGradScaler(required for FP16 but not BF16):
mixed_precision_dtype = {"fp16": torch.float16, "bf16": torch.bfloat16}.get(config.mixed_precision)
enable_amp = mixed_precision_dtype is not None
scaler = GradScaler(enabled=enable_amp and mixed_precision_dtype == torch.float16)
Hardware Requirements and Compatibility
BF16 requires GPUs with native bfloat16 support, specifically NVIDIA Ampere architecture (A100, RTX 30-series/40-series) or newer. When BF16 is unavailable, the accelerator gracefully falls back to FP32, ensuring training continues with higher memory usage.
FP16 works on any CUDA device with Tensor Cores (most GPUs from the Pascal generation onward). The code includes specific handling for Apple MPS devices in train_scripts/train_dreambooth_lora_sana.py (line 827), where BF16 requests on MPS fall back to CPU-compatible modes.
Configuration Methods
You can set mixed precision through three primary interfaces: static YAML configs, runtime command-line overrides, or programmatically during inference.
Method 1: Setting Precision in YAML Configs
For training from scratch, modify the model configuration file:
# configs/sol_rl/Sana1.0_1600M_linear.yaml
model:
mixed_precision: bf16
Launch training with the standard entry script:
bash train_scripts/train.sh \
configs/sol_rl/Sana1.0_1600M_linear.yaml \
--data.data_dir="[data/toy_data]" \
--train.num_epochs=10
The script reads the mixed_precision field, constructs the Accelerator with BF16 support, and the entire pipeline executes in bfloat16.
Method 2: Overriding via Command Line
For quick experiments without editing files, override the config value using dot-notation arguments processed by pyrallis:
bash train_scripts/train.sh \
configs/sol_rl/Sana1.0_1600M_linear.yaml \
--model.mixed_precision=fp16 \
--train.num_epochs=5
The extra flag overwrites the YAML value before SanaConfig initialization, allowing rapid A/B testing between precision modes. Note that some specialized scripts like train_dreambooth_lora_sana.py also expose a direct --mixed_precision CLI flag (line 561).
Method 3: Loading BF16 Checkpoints for Inference
Inference scripts respect the same precision configuration. When loading a BF16 checkpoint, ensure the config matches the model weights:
import torch
from app.sana_pipeline import SanaPipeline
from diffusion.utils.config import SanaConfig
# Load configuration with BF16 precision
config = SanaConfig.from_yaml("configs/sana1-5_config/1024ms/Sana_1600M_1024px_allqknorm_bf16_lr2e5.yaml")
pipeline = SanaPipeline(config) # Creates Accelerator(mixed_precision="bf16")
# Load checkpoint
pipeline.from_pretrained("hf://Efficient-Large-Model/SANA1.5_1.6B_1024px/checkpoints/SANA1.5_1.6B_1024px.pth")
# Generate
generator = torch.Generator().manual_seed(42)
image = pipeline(
prompt="a cyberpunk cat with neon lights",
height=1024,
width=1024,
num_inference_steps=20,
generator=generator,
)[0]
image.save("sana_bf16_demo.png")
The pipeline automatically casts the diffusion model, VAE, and text encoder to torch.bfloat16, reducing GPU memory consumption by approximately 50% compared to FP32 while preserving generation quality on compatible hardware.
Summary
- Configuration: Set
mixed_precision: bf16orfp16in YAML files (e.g.,configs/sol_rl/Sana1.0_1600M_linear.yaml). - Activation: The
Acceleratorintrain_scripts/train.pyautomatically enables AMP based on the config string. - Type casting: All components use
get_weight_dtype()fromdiffusion/model/utils.pyto convert strings totorch.bfloat16ortorch.float16. - Hardware: BF16 requires Ampere/A100+ GPUs; FP16 works on most Tensor Core-equipped cards.
- Override: Use
--model.mixed_precision=fp16to change precision without editing config files.
Frequently Asked Questions
Does Sana support both BF16 and FP16 for training?
Yes. The repository supports both BF16 (bfloat16) and FP16 (float16) through the mixed_precision configuration field. The get_weight_dtype utility in diffusion/model/utils.py handles conversion for both modes, and training scripts like train_scripts/sol_rl/train_sana.py conditionally configure GradScaler for FP16 while allowing BF16 to run without gradient scaling.
What happens if I request BF16 on unsupported hardware?
The accelerator automatically falls back to FP32 (full precision) if your GPU does not support native bfloat16 operations. Additionally, specific scripts like train_scripts/train_dreambooth_lora_sana.py include MPS (Apple Silicon) compatibility checks that route BF16 requests to CPU-compatible fallback paths to prevent runtime errors.
How do I verify mixed precision is active during training?
Monitor the Accelerator initialization logs at startup, which report the selected precision mode. In the training loop, check that mixed_precision_dtype in train_scripts/sol_rl/train_sana.py evaluates to torch.bfloat16 or torch.float16 rather than None, and verify that GPU memory usage decreases by approximately 40-50% compared to FP32 baselines.
Can I use different precision modes for training and inference?
Yes. Training and inference configurations are independent. You can train in FP16 and evaluate in BF16, or vice versa, by providing different YAML configurations to the training script (train_scripts/train.py) versus the inference pipeline (scripts/inference_sana_sprint.py). Ensure the inference configuration in scripts/inference_sana_sprint.py (line 271) matches the checkpoint's saved dtype for optimal performance.
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 →