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:
-
Monitor memory allocation – Insert
torch.cuda.max_memory_allocated()probes before and after major operations (VAE encoding, transformer forward) to locate peak consumption. -
Enable mixed-precision – Add
--mixed_precision fp16(orbf16on RTX 30-series or newer) and--allow_tf32to halve activation memory. -
Activate gradient checkpointing – Include
--gradient_checkpointingto reduce transformer memory usage at the cost of increased compute time. -
Offload static components – Use
--offloadto relocate the VAE and text encoder to CPU during training steps. -
Install xformers – Verify
xformersis installed to enable memory-efficient attention kernels insana_blocks.py. -
Clear caches explicitly – Call
free_memory()andtorch.cuda.empty_cache()after preprocessing stages, as implemented in lines 88-92. -
Reduce micro-batch size – Lower
--train_batch_sizeto 1 and compensate with higher--gradient_accumulation_steps. -
Cache latents – For image datasets, enable
--cache_latentsto 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.pylines 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
--offloadflag when these components are not actively processing. - Utilize xformers attention kernels in
diffusion/model/nets/sana_blocks.pyfor memory-efficient attention computation. - Pre-cache VAE latents with
--cache_latentsto eliminate redundant encoding passes over static datasets. - Adjust batch size and accumulation steps to control simultaneous activation memory.
- Clear GPU caches using
free_memory()andtorch.cuda.empty_cache()after heavy preprocessing operations. - Process video data in chunks as demonstrated in
train_video_scripts/train_video_ivjoint_chunk.pyto 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →