Latent Diffusion Models
Performing diffusion in latent space can make diffusion much more efficient in terms of generation and training.
Table of Contents
1. Motivation
Directly training and predicting in pixel space is computationally expensive. Therefore, a natural flow is to try training and predicting in a lower-dimensional space.
2. Latent Diffusion Models
Given the motivation and observations, LDMs adopts a two-stage solution. In the first stage, an encoder-decoder pair is trained to connect latent space and pixel space. In the second stage, a diffusion model is trained to operate in the learned latent space.
2.1. Stage One: Encoder and Decoder
An image \(x \in \mathbb{R}^{H\times W\times 3}\) is encoded by an encoder \(\mathcal{E}\) into a lower-dimensional latent representation \(z = \mathcal{E}(x) \in \mathbb{R}^{h \times w \times c}\) by a factor \(f = H/h = W/w = 2^{m}\) where \(m \in \mathbb{N}\).
The decoder \(\mathcal{D}\) reconstructs the image from the latent, i.e., \(\tilde{x}=\mathcal{D}(z)=\mathcal{D}(\mathcal{E}(x))\)
The encoder and decoder are trained only in stage one, and are frozen in latter stages for reuse.
The training objective involves a perceptual loss and a patch-based adversarial objective.
2.2. Stage Two: Diffusion in Latent Space
At this stage, it’s just applying diffusion technique in the learned latent space, as normal pixel space do. More specifically, it’s denoising Gaussian noise \(z_{T}\) to image latent \(z_{0}\). The authors uses U-Net here to predict the denoised noise from the input latent, as the original DDPM do.
2.3. Flexible Multimodal Conditioning via Cross-Attention
Maybe we can discuss a little bit about multimodal conditioning in latent space. The conditioning could be text (e.g., tasks like generating images from text prompts), images (e.g., image editing), etc.
LDMs support this by inserting cross-attention into the UNet backbone.
- First, a domain-specific encoder \(\tau_{\theta}\) is introduced to encode conditioning input \(y\) and obtain an intermediate representation \(\tau_{\theta}(y) \in \mathbb{R}^{M \times d_{\tau}}\)
Then, the cross attention layers treat projected latents as query tokens, while projected labels as keys and values, i.e., where \(\phi_{i}\) denotes a (flattened) latent.
\[ Q = W_{Q}^{(i)}\phi_{i}(z_{t}) \\ K = W_{K}^{(i)} \tau_{\theta}(y) \\ V = W_{V}^{(i)} \tau_{\theta}(y) \]