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

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, these flags are defined at lines 62-70:


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


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


# 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 (lines 85-87), the attention mechanism utilizes:

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, cache clearing occurs at lines 88-92 and 112-114:


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


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

--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, the chunk loop at lines 1169-1172 demonstrates:


# 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.

  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:

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 to track memory usage across training stages:

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:

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 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 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 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 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.

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 →