# Debugging Out of Memory (OOM) Issues During High-Resolution Image Training in Sana

> Learn to debug Out of Memory OOM issues during high-resolution image training in Sana. Reduce GPU memory by 50% with mixed precision, gradient checkpointing, component offloading, and xformers.

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

---

**Enable mixed-precision training (fp16/bf16), gradient checkpointing, and component offloading while utilizing xformers memory-efficient attention to reduce GPU memory consumption by approximately 50% when training Sana on high-resolution images above 1024px.**

Sana’s training pipelines are built on the 🤗 Diffusers library and utilize large-scale vision-language models including the VAE, text encoder, and transformer backbone. When training with high-resolution inputs exceeding 1024 pixels, the activation memory footprint quickly exceeds available VRAM, causing CUDA out-of-memory errors. The NVlabs/Sana repository provides eight orthogonal optimization mechanisms to systematically debug and mitigate these OOM crashes without sacrificing model fidelity.

## Memory Optimization Mechanisms

Sana implements distinct strategies across the training stack to minimize GPU memory allocation. These techniques can be combined to compound memory savings during high-resolution training sessions.

### Mixed-Precision and TensorFloat-32

Training with **mixed-precision** (fp16 or bf16) reduces the bit-width of activations and weights, cutting memory usage by roughly half while maintaining training stability. On Ampere GPUs, enabling **TensorFloat-32** (TF-32) provides additional speed benefits without extra memory overhead.

In [`train_scripts/train_dreambooth_lora_sana.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train_dreambooth_lora_sana.py), these flags are defined at lines 62-70:

```python

# Line 62-66: TF-32 enablement for Ampere GPUs

if args.allow_tf32:
    torch.backends.cuda.matmul.allow_tf32 = True
    torch.backends.cudnn.allow_tf32 = True

# Line 65-70: Mixed precision configuration

accelerator = Accelerator(
    mixed_precision=args.mixed_precision,  # "fp16", "bf16", or "no"

    gradient_accumulation_steps=args.gradient_accumulation_steps,
    ...
)

```

### Gradient Checkpointing

**Gradient checkpointing** trades compute for memory by discarding intermediate activations during the forward pass and recomputing them during backward propagation. This significantly lowers peak memory usage at the cost of approximately 20-30% additional training time.

Enable this feature via the `--gradient_checkpointing` flag (line 96), which internally calls:

```python

# Line 84 in train_scripts/train_dreambooth_lora_sana.py

transformer.enable_gradient_checkpointing()

```

### Component Offloading

When preprocessing text embeddings or caching latents, the VAE and text encoder consume substantial VRAM while remaining idle during transformer updates. **Component offloading** moves these modules to CPU when not actively processing data.

The `--offload` flag (line 80) triggers the relocation logic at line 71:

```python

# Line 71: Offloading to CPU to free GPU memory

if args.offload:
    text_encoding_pipeline = text_encoding_pipeline.to("cpu")
    torch.cuda.empty_cache()

```

### XFormers Memory-Efficient Attention

The default PyTorch attention implementation materializes full N×N score matrices, causing quadratic memory growth with sequence length. Sana integrates **xformers memory-efficient attention** to replace dense kernels with block-wise implementations.

In [`diffusion/model/nets/sana_blocks.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/nets/sana_blocks.py) (lines 85-87), the attention mechanism utilizes:

```python
import xformers.ops as xops

# Line 85-87: Memory-efficient attention kernel

attn_output = xops.memory_efficient_attention(
    q, k, v, attn_bias=None, op=None
)

```

Ensure xformers is installed (`pip install xformers`) to activate this optimization automatically.

### Explicit Cache Clearing

Sana provides explicit memory deallocation utilities to prevent PyTorch's caching allocator from retaining unused tensors. The `free_memory()` helper is invoked after heavy operations like prompt encoding and VAE latent generation.

In [`train_scripts/train_dreambooth_lora_sana.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train_dreambooth_lora_sana.py), cache clearing occurs at lines 88-92 and 112-114:

```python

# Lines 88-92: Post text-encoding cleanup

free_memory()
torch.cuda.empty_cache()

# Lines 112-114: Post-VAE caching cleanup

if args.cache_latents:
    free_memory()
    torch.cuda.empty_cache()

```

The `free_memory()` wrapper is defined in [`diffusion/utils/optimizer.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/optimizer.py) and extends `torch.cuda.empty_cache()` with Python garbage collection.

### Latent Caching

For static image datasets, repeatedly encoding images through the VAE during each epoch wastes memory and compute. **Latent caching** pre-computes and stores VAE outputs on GPU, eliminating redundant forward passes.

Enable with `--cache_latents` (line 104), which executes the caching loop at lines 104-110:

```python

# Lines 104-110: Pre-compute and cache VAE latents

if args.cache_latents:
    with torch.no_grad():
        latents = vae.encode(batch["pixel_values"]).latent_dist.sample()
        latents = latents * vae.config.scaling_factor
    # Store latents for training loop reuse

```

### Batch Size and Gradient Accumulation

When memory constraints persist, reducing the **train batch size** while increasing **gradient accumulation steps** maintains the effective batch size while lowering simultaneous activation memory. This is controlled via arguments at lines 54-57:

```bash
--train_batch_size 1 \
--gradient_accumulation_steps 8

```

This configuration achieves the same optimization steps as batch size 8 with fractional memory requirements.

### Chunked Video Processing

For video-style training data, Sana processes long temporal sequences by splitting tensors into smaller chunks that are processed sequentially. This approach explicitly deletes each chunk after computation to bound memory usage.

In [`train_video_scripts/train_video_ivjoint_chunk.py`](https://github.com/NVlabs/Sana/blob/main/train_video_scripts/train_video_ivjoint_chunk.py), the chunk loop at lines 1169-1172 demonstrates:

```python

# Lines 1169-1172: Chunked processing with explicit deletion

for chunk in video_chunks:
    output = process_chunk(chunk)
    # Save output to disk or buffer

    del chunk  # Explicit deletion to free memory

    torch.cuda.empty_cache()

```

## OOM Debugging Workflow

When encountering out-of-memory errors, follow this systematic diagnostic approach to identify and resolve memory bottlenecks:

1. **Monitor memory allocation** – Insert `torch.cuda.max_memory_allocated()` probes before and after major operations (VAE encoding, transformer forward) to locate peak consumption.

2. **Enable mixed-precision** – Add `--mixed_precision fp16` (or `bf16` on RTX 30-series or newer) and `--allow_tf32` to halve activation memory.

3. **Activate gradient checkpointing** – Include `--gradient_checkpointing` to reduce transformer memory usage at the cost of increased compute time.

4. **Offload static components** – Use `--offload` to relocate the VAE and text encoder to CPU during training steps.

5. **Install xformers** – Verify `xformers` is installed to enable memory-efficient attention kernels in [`sana_blocks.py`](https://github.com/NVlabs/Sana/blob/main/sana_blocks.py).

6. **Clear caches explicitly** – Call `free_memory()` and `torch.cuda.empty_cache()` after preprocessing stages, as implemented in lines 88-92.

7. **Reduce micro-batch size** – Lower `--train_batch_size` to 1 and compensate with higher `--gradient_accumulation_steps`.

8. **Cache latents** – For image datasets, enable `--cache_latents` to avoid repeated VAE encoding passes.

## Implementation Examples

### Complete Training Launch Configuration

The following command combines all primary optimization flags for high-resolution training on limited VRAM:

```bash
accelerate launch \
  --mixed_precision fp16 \
  --gradient_checkpointing \
  --offload \
  train_scripts/train_dreambooth_lora_sana.py \
  --pretrained_model_name_or_path ./pretrained/sana \
  --instance_data_dir ./my_images \
  --instance_prompt "photo of my pet" \
  --train_batch_size 2 \
  --gradient_accumulation_steps 4 \
  --cache_latents \
  --allow_tf32 \
  --lr_scheduler cosine \
  --learning_rate 5e-5

```

### Memory Profiling Snippet

Add this utility function to [`train_scripts/train_dreambooth_lora_sana.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train_dreambooth_lora_sana.py) to track memory usage across training stages:

```python
import torch

def log_mem(stage: str):
    """Log current and peak GPU memory allocation."""
    allocated = torch.cuda.memory_allocated() / 1e9
    max_alloc = torch.cuda.max_memory_allocated() / 1e9
    print(f"[{stage}] cur: {allocated:.2f} GB, max: {max_alloc:.2f} GB")

# Usage throughout training loop

log_mem("after VAE encode")

# ... transformer forward pass ...

log_mem("after transformer forward")
torch.cuda.empty_cache()
log_mem("after cache clear")

```

### Direct XFormers Attention Usage

For custom modifications to attention mechanisms in [`diffusion/model/nets/sana_blocks.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/nets/sana_blocks.py):

```python
from xformers.ops import memory_efficient_attention as xformers_attn

def efficient_attention(q, k, v, dropout=0.0):
    """
    q, k, v: tensors of shape (batch, num_heads, seq_len, head_dim)
    """
    return xformers_attn(
        q, k, v, 
        dropout_p=dropout, 
        attn_bias=None,
        scale=None  # Uses default 1/sqrt(head_dim)

    )

```

## Summary

Debugging OOM issues in Sana high-resolution training requires a multi-layered approach targeting specific memory bottlenecks in the Diffusers pipeline:

- **Enable mixed-precision (fp16/bf16)** and **TF-32** to reduce tensor memory by 50% via flags in [`train_scripts/train_dreambooth_lora_sana.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/train_dreambooth_lora_sana.py) lines 62-70.
- **Activate gradient checkpointing** by calling `transformer.enable_gradient_checkpointing()` to trade compute for memory.
- **Offload VAE and text encoders** to CPU using the `--offload` flag when these components are not actively processing.
- **Utilize xformers attention** kernels in [`diffusion/model/nets/sana_blocks.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/model/nets/sana_blocks.py) for memory-efficient attention computation.
- **Pre-cache VAE latents** with `--cache_latents` to eliminate redundant encoding passes over static datasets.
- **Adjust batch size and accumulation steps** to control simultaneous activation memory.
- **Clear GPU caches** using `free_memory()` and `torch.cuda.empty_cache()` after heavy preprocessing operations.
- **Process video data in chunks** as demonstrated in [`train_video_scripts/train_video_ivjoint_chunk.py`](https://github.com/NVlabs/Sana/blob/main/train_video_scripts/train_video_ivjoint_chunk.py) to bound temporal memory growth.

## Frequently Asked Questions

### Why does Sana OOM specifically at high resolutions above 1024px?

High-resolution images increase activation memory quadratically in attention layers and linearly in convolutional VAE operations. Sana's transformer architecture processes flattened image patches, causing the sequence length (N) to grow with the square of image dimensions, which exponentially increases the attention score matrix memory (N²) if using standard dense attention implementations.

### How much VRAM can I save with gradient checkpointing versus latent caching?

Gradient checkpointing typically reduces peak memory by 30-40% for transformer backbones while increasing training time by 20-30%. Latent caching eliminates the VAE's forward pass memory entirely during training epochs, saving approximately 2-4 GB depending on batch size and VAE precision, with no runtime overhead after initial caching.

### Should I use fp16 or bf16 mixed precision for Sana training?

Use **bf16** (bfloat16) on NVIDIA Ampere GPUs (RTX 30-series, A100, H100) as it maintains the same dynamic range as fp32 while offering fp16's memory benefits, preventing gradient underflow issues. Use **fp16** on older Turing GPUs (RTX 20-series) with gradient scaling to maintain stability, or if bf16 is unavailable in your CUDA environment.

### Where is the `free_memory()` function defined in the Sana codebase?

The `free_memory()` helper is defined in [`diffusion/utils/optimizer.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/utils/optimizer.py) and imported by the training scripts. It wraps `torch.cuda.empty_cache()` with additional Python garbage collection calls to ensure complete memory deallocation after heavy operations like VAE encoding or prompt preprocessing.