# How to Configure Mixed Precision Training (BF16/FP16) in Sana

> Learn to configure mixed precision training (bf16/fp16) in NVlabs/Sana by updating your YAML config. Accelerate your training with automatic mixed precision (AMP).

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

---

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

```yaml

# 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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train.py) (line 662), this enables torch-autocast automatically:

```python
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`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/utils.py) (line 598) converts configuration strings into concrete PyTorch dtypes:

```python
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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/train_scripts/sol_rl/train_sana.py) (line 441) check the dtype to conditionally enable `GradScaler` (required for FP16 but not BF16):

```python
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`](https://github.com/NVlabs/Sana/blob/main/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:

```yaml

# configs/sol_rl/Sana1.0_1600M_linear.yaml

model:
  mixed_precision: bf16

```

Launch training with the standard entry script:

```bash
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
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`](https://github.com/NVlabs/Sana/blob/main/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:

```python
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`](https://github.com/NVlabs/Sana/blob/main/configs/sol_rl/Sana1.0_1600M_linear.yaml)).
- **Activation**: The `Accelerator` in [`train_scripts/train.py`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/utils.py) handles conversion for both modes, and training scripts like [`train_scripts/sol_rl/train_sana.py`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/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`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train.py)) versus the inference pipeline ([`scripts/inference_sana_sprint.py`](https://github.com/NVlabs/Sana/blob/main/scripts/inference_sana_sprint.py)). Ensure the inference configuration in [`scripts/inference_sana_sprint.py`](https://github.com/NVlabs/Sana/blob/main/scripts/inference_sana_sprint.py) (line 271) matches the checkpoint's saved dtype for optimal performance.