Stable Diffusion made high-quality image generation run on a single consumer GPU. That was possible because of one architectural decision: denoise in a compressed latent space rather than in pixels. Every later member of the family keeps that decision, from SD 1.5 through SDXL to the transformer-based SD3 and SD3.5. What changes between versions is the denoiser, the text encoders and the training objective, and those changes drive where GPU memory and time go.
This article explains latent diffusion from the GPU's point of view. It covers the three networks and what each costs, the training objectives and a complete training step, a memory budget for fine-tuning, the tricks that make training stable, and what dominates inference latency and how to cut it. Code uses PyTorch and the Hugging Face diffusers library, with configuration values read from the model rather than hard-coded, because they differ between versions.
Three networks, one expensive loop
A Stable Diffusion model is three networks. The variational autoencoder (VAE) compresses an image eight times in each spatial dimension into a small latent tensor, and decodes latents back to pixels. One or more text encoders turn the prompt into a sequence of embeddings. The denoiser takes a noisy latent, a timestep and the text embeddings, and predicts how to move the latent toward a clean image. Only the denoiser runs at every sampling step, so it dominates cost.
| Model | Denoiser | Text encoders | Latent at native size | Objective |
|---|---|---|---|---|
| SD 1.5 | UNet, about 0.86B parameters | CLIP ViT-L/14 | 64 x 64 x 4 for 512 px | Noise prediction |
| SDXL base | UNet, about 2.6B | CLIP ViT-L and OpenCLIP ViT-bigG | 128 x 128 x 4 for 1024 px | Noise prediction |
| SD3.5 Large | MMDiT transformer, about 8.1B | Two CLIP models plus T5-XXL | 128 x 128 x 16 for 1024 px | Rectified flow |
The UNet interleaves convolutional residual blocks with attention at several resolutions. Self-attention mixes spatial positions, and cross-attention lets latent positions attend to text tokens. MMDiT, the SD3 design, patchifies the latent into tokens and runs a transformer in which image and text tokens share joint attention but keep separate weights per modality. For a transformer the token count sets the cost. A 128 x 128 latent with 2 x 2 patches gives 4,096 image tokens, and attention cost grows with the square of that number.
Why diffusion runs in latent space
Why latents? A 512 x 512 RGB image has 786,432 values. Its SD 1.5 latent has 64 x 64 x 4 = 16,384, which is 48 times fewer. Convolution cost scales roughly with spatial area, and attention cost with its square, so working at 1/64 of the pixel area is what made training and sampling affordable. The VAE is trained once, separately, with reconstruction, perceptual and adversarial losses. It is then frozen. Its quality caps the model's fine detail, which is why SD3 moved to 16 latent channels: more channels preserve small text and texture better.
Latents are not unit-variance as they come out of the encoder, so each model defines a scaling factor, and SD3 adds a shift factor, applied before diffusion. These constants differ between model families. Using the wrong one is a classic silent bug, so read them from the VAE's configuration instead of copying a number from a blog post.
Training objectives: noise, velocity and flow
Diffusion training picks a clean latent x0, a random timestep t and Gaussian noise ε, and builds a noisy sample. For the DDPM-style schedules used by SD 1.x and SDXL, x_t = √ᾱ_t · x0 + √(1 − ᾱ_t) · ε, where ᾱ_t falls from near 1 to near 0 across 1,000 training timesteps. The network then predicts a target, and three targets are common.
- Noise prediction (epsilon). Predict ε. Used by SD 1.x and SDXL base. Simple, but the prediction is poorly conditioned at the noisiest steps.
- Velocity prediction (v). Predict v = √ᾱ_t · ε − √(1 − ᾱ_t) · x0. Used by the 768-pixel SD 2.x models. Better behaved at both ends of the schedule.
- Rectified flow. Interpolate in a straight line, x_t = (1 − t) · x0 + t · ε, and predict the velocity ε − x0. Used by SD3 and SD3.5. Straighter paths allow fewer sampling steps, and SD3 samples t from a distribution that favours the middle of the range rather than uniformly.
The scheduler object stores which target a checkpoint expects in its configuration. Training a checkpoint on the wrong target does not crash. It produces a model that gets steadily worse.
A training step, line by line
Here is one fine-tuning step for an SDXL-class UNet in diffusers style. The VAE and text encoders are frozen, so they run under no_grad in bf16. Min-SNR weighting (Hang et al., 2023) caps the weight of low-noise timesteps, which otherwise dominate the gradient and slow convergence. A γ of 5 is the paper's suggested value.
import torch, torch.nn.functional as F
sf = vae.config.scaling_factor # read, never hard-code
ac = noise_scheduler.alphas_cumprod.to(device)
def train_step(pixels, prompt_embeds, added_cond, gamma=5.0):
with torch.no_grad():
latents = vae.encode(pixels.to(vae.dtype)).latent_dist.sample() * sf
noise = torch.randn_like(latents)
t = torch.randint(0, noise_scheduler.config.num_train_timesteps,
(latents.shape[0],), device=device)
noisy = noise_scheduler.add_noise(latents, noise, t)
with torch.autocast("cuda", dtype=torch.bfloat16):
pred = unet(noisy, t, encoder_hidden_states=prompt_embeds,
added_cond_kwargs=added_cond).sample # SDXL size/crop conditioning
if noise_scheduler.config.prediction_type == "epsilon":
target = noise
else: # "v_prediction"
target = noise_scheduler.get_velocity(latents, noise, t)
snr = ac[t] / (1 - ac[t])
w = torch.clamp(snr, max=gamma) / snr
if noise_scheduler.config.prediction_type == "v_prediction":
w = torch.clamp(snr, max=gamma) / (snr + 1)
loss = F.mse_loss(pred.float(), target.float(), reduction="none")
loss = (loss.mean(dim=(1, 2, 3)) * w).mean()
loss.backward()
torch.nn.utils.clip_grad_norm_(unet.parameters(), 1.0)
optimizer.step(); lr_scheduler.step(); optimizer.zero_grad(set_to_none=True)
return loss.item()Two more data tricks matter. Drop the caption, replacing it with an empty prompt, for roughly 10% of samples, so the model also learns the unconditional prediction that classifier-free guidance needs at inference. And bucket images by aspect ratio so that each batch has a single shape without cropping away subjects. SDXL adds micro-conditioning: the original image size and the crop coordinates are fed in as extra conditioning, so the model learns not to imitate low-resolution or badly cropped training images.
GPU memory budget for fine-tuning
Plan GPU memory before launching. A full fine-tune with AdamW in mixed precision keeps fp32 master weights, gradients and two optimizer moments, about 16 bytes per trained parameter. For SDXL's 2.6B-parameter UNet that is roughly 42 GB before a single activation. That fits an 80 GB card with gradient checkpointing and modest batch sizes. On smaller cards it needs FSDP or ZeRO sharding across GPUs, or an 8-bit optimizer. LoRA changes the picture. The base weights stay frozen in bf16, about 5.2 GB for the SDXL UNet, and only adapter matrices of a few million to tens of millions of parameters carry optimizer state. That is why LoRA fine-tuning of SDXL fits on a 24 GB card.
Activations scale with batch size, resolution and attention tokens. Gradient checkpointing recomputes block activations in the backward pass, trading roughly a third more compute for a large memory cut. Fused attention kernels such as PyTorch's scaled_dot_product_attention or FlashAttention avoid materialising the attention matrix. The frozen encoders are pure overhead during training. Precompute the latents and text embeddings once, store them, and train from the cache. That removes the VAE and text encoders from the step entirely, at the cost of fixing your augmentations at caching time.
Inference: where the time goes
Sampling cost is roughly steps × guidance passes × one denoiser forward, plus one VAE decode. Classifier-free guidance runs the denoiser on both the conditional and the unconditional input, usually as a batch of two, and combines them: ε = ε_uncond + w · (ε_cond − ε_uncond). Guidance therefore doubles the denoiser compute. The w value trades prompt adherence against saturation and artefacts.
Step count is the biggest lever. Multistep solvers such as DPM-Solver++ produce good SD 1.5 and SDXL images in about 20 to 30 steps instead of the 1,000 used in training. Rectified-flow models use Euler-type flow-matching schedulers. Distilled variants go further. Latent consistency models, SDXL Turbo and SD3.5 Large Turbo target one to a handful of steps, at some cost in diversity and fine control. Check each model card for its recommended steps and guidance, because distilled models often expect guidance to be off.
import torch
from diffusers import StableDiffusionXLPipeline, DPMSolverMultistepScheduler
pipe = StableDiffusionXLPipeline.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
torch_dtype=torch.float16, variant="fp16").to("cuda")
pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config)
pipe.enable_vae_tiling() # avoids the decode memory spike at large sizes
image = pipe("an isometric diagram of a GPU server rack, studio lighting",
num_inference_steps=25, guidance_scale=6.0).images[0]Other latency levers: torch.compile on the denoiser (pay the compile time once per shape), channels-last memory format for convolution-heavy UNets, fixed resolutions so compiled graphs are reused, and batching requests of the same shape. The VAE decode is a single call but has a large activation peak at high resolution. Tiled decode bounds it. The original SDXL VAE is known to overflow in fp16 and produce NaNs, so run it in fp32 or bf16, or use a VAE fine-tuned to be fp16-safe.
For capacity planning, estimate GPU seconds per image as steps × guidance factor × denoiser time per step, plus decode time, measured on your own hardware at your own resolution. Images per GPU-hour then follow directly. Requests do not batch as freely as LLM tokens, because every request in a batch must share resolution and usually step count, so group the queue by shape. If prompts repeat, cache text-encoder outputs, which removes the T5 cost for SD3-class models on those requests.
Worked example: a brand-style LoRA for SDXL
Worked example: a team wants product shots in their brand style from SDXL. They have 1,500 licensed, captioned images. Full fine-tuning would need sharded optimizer state across several GPUs for little benefit, so they train a LoRA of rank 32 on the UNet attention projections. They precompute latents and both text-encoder outputs at three aspect-ratio buckets near one megapixel, which removes the encoders from the training step. They use bf16 autocast, gradient checkpointing, a batch of 4 with gradient accumulation of 4, a learning rate of 1e-4, 10% caption dropout and min-SNR weighting.
They evaluate every 500 steps on a fixed prompt set with fixed seeds, and include prompts that do not mention the brand to catch style bleeding into everything. At serving time they load the LoRA into the base pipeline, use 25 DPM-Solver++ steps at guidance 6, and enable tiled decode. They measure p50 and p95 latency per resolution bucket before choosing a GPU type, rather than guessing from parameter counts.
Failure modes
- Black or NaN images. fp16 overflow, usually in the VAE. Decode in bf16 or fp32.
- Washed-out or noisy training results. Wrong latent scaling or shift factor, or a target that does not match the scheduler's prediction type.
- Overfitting and style bleed. Too many steps on a small set. Watch held-out prompts, lower the rank or the learning rate, or stop earlier.
- Guidance does nothing. No caption dropout during training, so the unconditional branch was never learned.
- Prompt silently truncated. CLIP encoders take 77 tokens. Longer prompts are cut unless the pipeline chunks them, and T5 accepts more.
- Out of memory only at inference. The CFG batch of two plus the VAE decode peak at high resolution. Tile the decode or lower the batch.
Trade-offs
Resolution is the trade-off people underestimate. Doubling the image side length quadruples the latent area, so convolution cost rises about four times. For a transformer denoiser the token count also quadruples, and full attention cost rises about sixteen times. Training at a lower resolution first and finishing at the target resolution is common for this reason, and serving at native resolution, then upscaling, is often cheaper than sampling large. UNets are cheaper per step and have a mature ecosystem of adapters and control networks. MMDiT models follow prompts and render text better but cost more per step and need large text encoders; T5-XXL alone is about 5B parameters. LoRA is cheap and swappable but captures less than full fine-tuning. Fewer steps cut latency linearly and eventually cost quality. Distillation keeps quality at very low step counts but narrows diversity and complicates further fine-tuning.
What to do next
- Read the VAE scaling factor, the scheduler prediction type and the recommended steps from your checkpoint's config.
- Precompute latents and text embeddings for your dataset and train the denoiser from the cache.
- Start with LoRA, bf16 autocast, gradient checkpointing, 10% caption dropout and min-SNR weighting.
- Evaluate on fixed seeds and prompts, including prompts that should not change.
- For serving, benchmark steps, guidance, tiled decode and torch.compile per resolution bucket.
- Keep learning: video diffusion, mixed-precision training, FlashAttention, FSDP sharding and torch.compile.