Implementing LoRA Fine-Tuning with DreamBooth for Sana
This guide explains how to create personalized LoRA adapters for the Sana diffusion model using DreamBooth techniques, enabling efficient fine-tuning on consumer GPUs while keeping the base model parameters frozen.
The NVlabs/Sana repository provides a complete implementation for customizing the Sana text-to-image diffusion model through Low-Rank Adaptation (LoRA). By combining DreamBooth data augmentation with parameter-efficient fine-tuning, you can teach the model new visual concepts using minimal compute resources and storage.
Dataset Preparation and Validation
The training script train_scripts/train_dreambooth_lora_sana.py accepts data from either a Hugging Face dataset (--dataset_name) or a local image folder (--instance_data_dir). Validation logic at lines 79-99 ensures exactly one data source is specified.
The DreamBoothDataset class handles two concurrent image streams: instance images for your target concept and class images for prior preservation. It applies standard augmentation including resizing, center-cropping, random horizontal flips, and normalization to the [-1, 1] range (lines 24-84).
Prior Preservation Configuration
When using --with_prior_preservation, the dataset automatically manages class images to prevent overfitting to your specific instances. Specify the class prompt via --class_prompt and the storage directory via --class_data_dir. The dataset repeats instance images according to the --repeats parameter to balance the sampling ratio.
Injecting LoRA Adapters into the Sana Transformer
Only the transformer backbone receives LoRA weights; the VAE and text encoder remain frozen for computational efficiency. The adapter configuration uses PEFT's LoraConfig with default targeting of attention layers:
target_modules = (
[layer.strip() for layer in args.lora_layers.split(",")]
if args.lora_layers is not None
else ["to_k", "to_q", "to_v"]
)
transformer_lora_config = LoraConfig(
r=args.rank,
lora_alpha=args.rank,
init_lora_weights="gaussian",
target_modules=target_modules,
)
transformer.add_adapter(transformer_lora_config)
This injection occurs at lines 86-99 in train_scripts/train_dreambooth_lora_sana.py. The default rank is 4, but you can adjust this via the --rank argument to trade off between adapter capacity and file size.
Training Loop and Optimization
The training loop runs for args.max_train_steps and supports text-embedding caching (compute_text_embeddings) and latent caching (--cache_latents) to accelerate training by avoiding redundant encodings (starting around line 311).
Custom Save and Load Hooks
The script registers accelerator hooks to handle LoRA state dicts separately from the base model weights. The save_model_hook extracts trainable parameters using get_peft_model_state_dict and persists them via SanaPipeline.save_lora_weights (lines 106-124). During resumption, the load_model_hook restores the adapter using SanaPipeline.lora_state_dict and set_peft_model_state_dict (lines 124-143), ensuring seamless checkpoint recovery.
Optimizer Selection
You can choose between AdamW (with optional 8-bit quantization via bitsandbytes) or Prodigy for adaptive learning rate optimization. Configure this through the script's argument parser to match your hardware constraints and convergence requirements.
Running DreamBooth Training
Execute the full training pipeline with this command structure:
python -m train_scripts.train_dreambooth_lora_sana \
--pretrained_model_name_or_path path/to/sana_base \
--instance_data_dir ./my_photos \
--instance_prompt "photo of a my-cat" \
--output_dir ./sana-cat-lora \
--rank 8 \
--train_batch_size 4 \
--num_train_epochs 3 \
--learning_rate 5e-5 \
--with_prior_preservation \
--class_prompt "photo of a cat" \
--class_data_dir ./class_images \
--push_to_hub \
--hub_model_id myusername/sana-cat-lora
Key parameters include --rank for adapter dimensionality and --instance_prompt which serves as the trigger word for subsequent generation. The --with_prior_preservation flag is essential for maintaining the model's generalization ability for the base class.
Inference with LoRA Weights
Load the fine-tuned adapter using SanaPipeline from app/sana_pipeline.py:
from app.sana_pipeline import SanaPipeline
import torch
# Load base Sana model
pipe = SanaPipeline().from_pretrained("output/Sana_D20/SANA.pth")
# Attach the LoRA weights produced by the training script
pipe.load_lora_weights("./sana-cat-lora")
# Generate an image with the custom token
result = pipe(
prompt="photo of a my-cat playing with a ball",
guidance_scale=7.5,
num_inference_steps=50,
generator=torch.Generator().manual_seed(123),
)
result[0].save("cat_generated.png")
The pipeline automatically switches between classifier-free guidance and Perturbed Attention Guidance (PAG) based on the pag_scale and attn_type parameters (lines 66-89 in app/sana_pipeline.py).
Exporting and Sharing Models
Automatic Model Card Generation
The save_model_card function (lines 66-94) generates a README.md inside the output directory containing validation images, the trigger word, and ready-to-use code snippets. When using --push_to_hub, this card accompanies the LoRA weights to the Hugging Face Hub, enabling immediate community sharing and version control.
Summary
- Dataset Flexibility: Use either Hugging Face datasets or local folders with automatic prior-preservation image handling to maintain model fidelity.
- Efficient Architecture: LoRA adapters attach only to transformer attention layers (
to_q,to_k,to_vby default), leaving VAE and text encoder weights frozen and reducing memory overhead. - Robust Checkpointing: Custom hooks in
train_dreambooth_lora_sana.pyhandle conversion between PEFT state dicts and the pipeline's expected format viaSanaPipelineutilities. - Seamless Inference: The
SanaPipelineclass loads LoRA weights viaload_lora_weights()without requiring base model reinitialization, supporting both classifier-free and PAG guidance modes. - Hub Integration: Built-in model card generation and
--push_to_hubsupport streamline distribution of custom adapters through the Hugging Face ecosystem.
Frequently Asked Questions
What is the default LoRA rank for Sana DreamBooth training?
The default rank is 4, but you can specify any dimensionality using the --rank argument. Higher ranks increase the adapter's capacity to learn complex features but require more GPU memory and produce larger checkpoint files.
Which layers of the Sana transformer receive LoRA adapters?
By default, adapters attach to the query, key, and value projection layers (to_q, to_k, to_v). You can customize the target modules via the --lora_layers parameter to experiment with different fine-tuning strategies or reduce the parameter count further.
How do I resume training from a checkpoint?
The script handles resumption automatically through custom accelerator hooks. The load_model_hook restores adapter weights using SanaPipeline.lora_state_dict and set_peft_model_state_dict, ensuring training continues seamlessly from the exact optimization state where it stopped.
Can I use 8-bit Adam for memory-efficient training?
Yes, the script supports 8-bit AdamW optimization via the bitsandbytes library. Specify this optimizer choice through the command-line arguments to significantly reduce VRAM requirements during training, enabling fine-tuning on consumer-grade GPUs with limited memory.
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 →