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): Declares mixed_precision (e.g., bf16). Parsed by SanaConfig at diffusion/utils/config.py.
  • Accelerator: Created with the precision flag, enabling AMP and selecting the appropriate autocast dtype.
  • get_weight_dtype: Returns the concrete torch.dtype used to cast models and tensors.
  • Training loops: Scripts like train_scripts/sol_rl/train_sana.py (line 441) check the dtype to conditionally enable GradScaler (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: bf16 or fp16 in YAML files (e.g., configs/sol_rl/Sana1.0_1600M_linear.yaml).
  • Activation: The Accelerator in train_scripts/train.py automatically enables AMP based on the config string.
  • Type casting: All components use get_weight_dtype() from diffusion/model/utils.py to convert strings to torch.bfloat16 or torch.float16.
  • Hardware: BF16 requires Ampere/A100+ GPUs; FP16 works on most Tensor Core-equipped cards.
  • Override: Use --model.mixed_precision=fp16 to 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:

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 →