Latent Diffusion Model 深度剖析:从像素扩散到潜在空间生成

6669 字
33 分钟
Latent Diffusion Model 深度剖析:从像素扩散到潜在空间生成

1. LDM 的核心动机:为什么要在潜空间里扩散?#

1.1 像素扩散的计算瓶颈#

原始 DDPM 在像素空间直接生成图像遇到了严重的效率问题:

直接像素扩散的计算代价(以 512×512 RGB 图像为例):
图像分辨率: 512 × 512 × 3 = 786,432 像素
压缩后潜空间: 64 × 64 × 4 = 16,384 维度
压缩比: 48 倍!
UNet 前向计算量:
像素空间: O(512² × C) × 层数
潜空间: O(64² × C') × 层数
加速比: 64 倍!
推理时间:
DDPM (像素): ~20-60 分钟 (512×512)
LDM (潜空间): ~1-5 秒 (512×512)

1.2 感知压缩 vs 语义压缩#

LDM 论文(High-Resolution Image Synthesis with Latent Diffusion Models, Rombach et al., CVPR 2022)提出了一个关键洞察:扩散模型中不同层负责不同类型的压缩

┌─────────────────────────────────────────────────────────────┐
│ 扩散模型的多层压缩 │
│ │
│ 感知压缩 (Perceptual Compression) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 低层: 像素级细节 │ │
│ │ 如: 纹理、边缘、局部颜色 │ │
│ │ 压缩方式: VAE (无损重建局部细节) │ │
│ │ 压缩比: 4-48 倍 → 可以很激进 │ │
│ └─────────────────────────────────────────────────────┘ │
│ ↓ │
│ 语义压缩 (Semantic Compression) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 高层: 语义级概念 │ │
│ │ 如: 物体类别、空间布局、姿态 │ │
│ │ 压缩方式: DDPM (学习语义流形) │ │
│ │ 压缩比: 不能太激进 (会丢失语义) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 核心洞察: │
│ 低层信息(纹理、颜色)不需要用扩散建模 │
│ → VAE 可以高效压缩,不丢失感知质量 │
│ 高层语义才需要扩散过程 │
│ → 潜空间扩散专注于语义生成 │
│ │
└─────────────────────────────────────────────────────────────┘
@dataclass
class CompressionConfig:
"""
LDM 的压缩配置。
决定感知压缩的程度。
"""
# SD 1.x 配置
image_size: int = 512
latent_size: int = 64 # 512 / 8 = 64
latent_channels: int = 4 # 通道数 4 (而非 RGB 的 3)
downsample_factor: int = 8 # 8× 压缩
# 压缩比
compression_ratio: float = (512 * 512 * 3) / (64 * 64 * 4) # = 48
# SD 2.x 配置 (更保守的压缩)
# downsample_factor: 8, latent_channels: 4 (相同)
# SDXL 配置
# latent_size: 128 (1024 / 8)
# latent_channels: 4
# 更大的潜空间: 128×128×4 = 65536 (vs 64×64×4 = 16384)

1.3 LDM vs DDPM 的架构对比#

DDPM (像素空间):
[噪声 x_T] → [UNet] → [去噪 x_0] (在 512×512 像素空间)
问题: 每一步 UNet 都要处理 512×512 的高分辨率特征
LDM (潜空间):
[图像] → [VAE encoder] → [潜变量 z] (64×64×4)
[噪声 z_T] → [UNet (潜)] → [去噪 z_0] (在 64×64 潜空间)
[z_0] → [VAE decoder] → [重建图像]
优势: UNet 只处理 64×64 的低分辨率潜空间

1.4 论文信息#

论文: "High-Resolution Image Synthesis with Latent Diffusion Models"
作者: Robin Rombach, Andreas Blattmann, Dominik Lorenz, Patrick Esser,
Björn Ommer
单位: LMU Munich, IWR Heidelberg
发表于: CVPR 2022 (Oral)
引用: > 10,000 次 ★★★★★
开源: https://github.com/CompVis/stable-diffusion
配套: https://ommer-lab.com/research/latent-diffusion-models/

1.5 一句话概括 LDM#

LDM 的核心思想是把扩散过程从昂贵的像素空间搬到低维的潜在空间——用 VAE 编码器把图像压缩 48 倍得到潜变量,在潜空间里用 UNet 执行扩散去噪,最后用 VAE 解码器重建图像。这使得高分辨率图像生成的计算量降低了两个数量级,同时保持了几乎无损的视觉质量。

2. VAE:感知压缩的载体#

2.1 为什么 VAE 适合感知压缩?#

VAE(Variational AutoEncoder)用于 LDM 的感知压缩层,有三个关键原因:

@dataclass
class VAEProperties:
"""
VAE 用于感知压缩的优势。
"""
property_1: str = (
"无损局部重建: VAE 的解码器可以精确重建局部纹理和颜色,"
"而扩散过程不需要学习这些低层细节"
)
property_2: str = (
"隐式正则化: 潜空间是连续可微的,便于梯度反向传播,"
"且 KL 散度项鼓励潜空间平滑"
)
property_3: str = (
"独立于扩散过程: VAE 可以单独预训练,不需要参与扩散训练的梯度计算,"
"节省大量显存"
)

2.2 LDM 中使用的 VAE 架构#

LDM 使用的是一个改进的 VAE,架构参考了 VQ-VAE 和标准 VAE 的融合:

class LDMVAE(nn.Module):
"""
LDM (Stable Diffusion) 使用的 VAE。
基于 VQ-VAE + 标准 VAE 的融合设计。
"""
def __init__(self, latent_channels=4):
super().__init__()
self.latent_channels = latent_channels
# ============ 编码器 ============
# 输入: 512×512×3
self.encoder = nn.Sequential(
# 块 1: 512→256
nn.Conv2d(3, 128, kernel_size=3, padding=1),
ResBlock(128, 128),
ResBlock(128, 128),
nn.Conv2d(128, 128, kernel_size=3, stride=2, padding=1), # down
# 块 2: 256→128
nn.Conv2d(128, 256, kernel_size=3, padding=1),
ResBlock(256, 256),
ResBlock(256, 256),
nn.Conv2d(256, 256, kernel_size=3, stride=2, padding=1), # down
# 块 3: 128→64
nn.Conv2d(256, 512, kernel_size=3, padding=1),
ResBlock(512, 512),
ResBlock(512, 512),
nn.Conv2d(512, 512, kernel_size=3, stride=2, padding=1), # down
# 块 4: 64→64 (保持)
nn.Conv2d(512, 512, kernel_size=3, padding=1),
ResBlock(512, 512),
ResBlock(512, 512),
# 最终: 64×64, 512 ch → 潜空间
nn.GroupNorm(32, 512),
nn.SiLU(),
nn.Conv2d(512, latent_channels, kernel_size=3, padding=1), # 4 ch
)
# 量化层: 实际用的是 VQ 或 EMA 量化
# LDM 1.x 使用的是 KL 正则化 VAE(而非 VQ)
# SD 2.x 使用的是 EMA 量化
self.quant_conv = nn.Conv2d(latent_channels, latent_channels, kernel_size=1)
# ============ 解码器 ============
self.decoder = nn.Sequential(
# 潜空间 → 64×64, latent_channels
nn.Conv2d(latent_channels, 512, kernel_size=3, padding=1),
ResBlock(512, 512),
ResBlock(512, 512),
# 块 1: 64→128
nn.ConvTranspose2d(512, 512, kernel_size=4, stride=2, padding=1), # up
nn.Conv2d(512, 512, kernel_size=3, padding=1),
ResBlock(512, 512),
ResBlock(512, 512),
# 块 2: 128→256
nn.ConvTranspose2d(512, 256, kernel_size=4, stride=2, padding=1), # up
nn.Conv2d(256, 256, kernel_size=3, padding=1),
ResBlock(256, 256),
ResBlock(256, 256),
# 块 3: 256→512
nn.ConvTranspose2d(256, 128, kernel_size=4, stride=2, padding=1), # up
nn.Conv2d(128, 128, kernel_size=3, padding=1),
ResBlock(128, 128),
ResBlock(128, 128),
# 最终: 512×512×3
nn.GroupNorm(32, 128),
nn.SiLU(),
nn.Conv2d(128, 3, kernel_size=3, padding=1),
)
def encode(self, x):
"""
编码图像到潜空间。
参数:
x: (B, 3, 512, 512) 原始图像
返回:
z: (B, 4, 64, 64) 潜变量
"""
h = self.encoder(x)
z = self.quant_conv(h)
return z
def decode(self, z):
"""
从潜空间解码到图像。
参数:
z: (B, 4, 64, 64) 潜变量
返回:
x_recon: (B, 3, 512, 512) 重建图像
"""
return self.decoder(z)
def forward(self, x):
"""
完整 VAE 前向(训练时使用)。
"""
z = self.encode(x)
x_recon = self.decode(z)
return x_recon, z
class ResBlock(nn.Module):
"""
Residual Block,用于 VAE 和 UNet。
"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.norm1 = nn.GroupNorm(32, in_channels)
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)
self.norm2 = nn.GroupNorm(32, out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
# 残差连接
if in_channels != out_channels:
self.shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1)
else:
self.shortcut = nn.Identity()
def forward(self, x):
h = self.norm1(x)
h = F.silu(h)
h = self.conv1(h)
h = self.norm2(h)
h = F.silu(h)
h = self.conv2(h)
return h + self.shortcut(x)

2.3 潜空间维度对比#

def latent_dimensions():
"""
不同模型配置的潜空间维度。
"""
configs = {
"LDM-4 (Stable Diffusion 1.x)": {
"image_size": 512,
"latent_size": 64, # 512 / 8
"latent_channels": 4,
"total_dims": 64 * 64 * 4, # 16384
"compression": "48×",
},
"LDM-4 (Stable Diffusion 2.x)": {
"image_size": 768,
"latent_size": 96, # 768 / 8
"latent_channels": 4,
"total_dims": 96 * 96 * 4, # 36864
"compression": "48×",
},
"SDXL": {
"image_size": 1024,
"latent_size": 128, # 1024 / 8
"latent_channels": 4,
"total_dims": 128 * 128 * 4, # 65536
"compression": "48×",
"变化": "更大的基础分辨率 + refiner 双阶段",
},
}
return configs

3. 条件机制:交叉注意力#

3.1 LDM 如何注入条件信息?#

LDM 支持多种条件:文本、边界图、语义分割、类别标签等。所有条件通过交叉注意力机制注入:

┌─────────────────────────────────────────────────────────────┐
│ 交叉注意力条件注入 │
│ │
│ 文本编码: CLIP Text Encoder │
│ "a cat sitting on a red couch" │
│ ↓ │
│ text_embeddings: (B, seq_len, d_text) │
│ ↓ │
│ 潜变量: z ∈ ℝ^{B×4×64×64} │
│ ↓ │
│ UNet 中的交叉注意力层: │
│ Query: UNet 特征 (B, H×W, d_model) │
│ Key: text_embeddings (B, seq_len, d_text) │
│ Value: text_embeddings (B, seq_len, d_text) │
│ ↓ │
│ 输出: 条件感知的潜变量特征 │
│ │
└─────────────────────────────────────────────────────────────┘
class CrossAttention(nn.Module):
"""
交叉注意力层。
用于将条件信息(文本)注入 UNet。
"""
def __init__(self, query_dim, context_dim=None, num_heads=8):
super().__init__()
self.num_heads = num_heads
self.head_dim = query_dim // num_heads
self.scale = self.head_dim ** -0.5
# 线性投影
self.to_q = nn.Linear(query_dim, query_dim, bias=False)
self.to_k = nn.Linear(context_dim or query_dim, query_dim, bias=False)
self.to_v = nn.Linear(context_dim or query_dim, query_dim, bias=False)
self.to_out = nn.Linear(query_dim, query_dim)
def forward(self, x, context=None, mask=None):
"""
参数:
x: (B, N, query_dim) UNet 特征
context: (B, seq_len, context_dim) 条件 embedding
mask: (B, seq_len) 注意力掩码
"""
if context is None:
context = x
B, N, _ = x.shape
seq_len = context.shape[1]
# 投影
q = self.to_q(x) # (B, N, D)
k = self.to_k(context) # (B, seq_len, D)
v = self.to_v(context) # (B, seq_len, D)
# 分头
q = q.view(B, N, self.num_heads, self.head_dim).transpose(1, 2) # (B, H, N, d)
k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
# 注意力
attn = (q @ k.transpose(-2, -1)) * self.scale # (B, H, N, seq_len)
if mask is not None:
# BERT-style 掩码: True = 需要 mask
mask_expanded = mask.view(B, 1, 1, seq_len).expand(-1, self.num_heads, N, -1)
attn = attn.masked_fill(mask_expanded, float("-inf"))
attn = F.softmax(attn, dim=-1)
out = attn @ v # (B, H, N, d)
# 合并头
out = out.transpose(1, 2).contiguous().view(B, N, -1)
return self.to_out(out)

3.2 文本编码器:CLIP Text Encoder#

class CLIPTextEncoder(nn.Module):
"""
LDM 使用的文本编码器。
来自 OpenAI 的 CLIP ViT-L/14 文本编码器。
"""
def __init__(self, model_name="openai/clip-vit-l-14"):
super().__init__()
from transformers import CLIPTextModel, CLIPTokenizer
self.tokenizer = CLIPTokenizer.from_pretrained(model_name)
self.text_encoder = CLIPTextModel.from_pretrained(model_name)
# 冻结文本编码器
for param in self.text_encoder.parameters():
param.requires_grad_(False)
self.text_embed_dim = self.text_encoder.config.hidden_size # 768
@property
def device(self):
return next(self.parameters()).device
def encode(self, texts):
"""
编码文本到 embedding。
参数:
texts: str 或 list[str]
返回:
text_embeddings: (B, seq_len, 768)
"""
if isinstance(texts, str):
texts = [texts]
# Tokenize
tokens = self.tokenizer(
texts,
padding="max_length",
max_length=77,
truncation=True,
return_tensors="pt",
).to(self.device)
# 编码
with torch.no_grad():
outputs = self.text_encoder(**tokens)
embeddings = outputs.last_hidden_state # (B, seq_len, 768)
return embeddings

3.3 完整的条件注入 UNet#

class UNetLDM(nn.Module):
"""
LDM 的 UNet,支持交叉注意力条件注入。
"""
def __init__(self, latent_channels=4, model_channels=320, context_dim=768):
super().__init__()
self.latent_channels = latent_channels
self.model_channels = model_channels
# 时间步嵌入
self.time_embed = TimeEmbedding(model_channels)
# 初始卷积
self.input_conv = nn.Conv2d(latent_channels, model_channels, kernel_size=3, padding=1)
# 下采样路径
self.down_blocks = nn.ModuleList([
DownBlock(model_channels, model_channels),
DownBlock(model_channels, model_channels * 2),
DownBlock(model_channels * 2, model_channels * 4),
DownBlock(model_channels * 4, model_channels * 4),
])
# 中间块(带交叉注意力)
self.mid_block = MidBlock(model_channels * 4, context_dim)
# 上采样路径
self.up_blocks = nn.ModuleList([
UpBlock(model_channels * 8, model_channels * 2, context_dim),
UpBlock(model_channels * 4, model_channels, context_dim),
UpBlock(model_channels * 2, model_channels, context_dim),
UpBlock(model_channels * 2, model_channels, context_dim),
])
# 最终输出
self.output_conv = nn.Sequential(
nn.GroupNorm(32, model_channels),
nn.SiLU(),
nn.Conv2d(model_channels, latent_channels, kernel_size=3, padding=1),
)
def forward(self, x, t, context):
"""
参数:
x: (B, 4, 64, 64) 带噪声的潜变量
t: (B,) 时间步
context: (B, seq_len, context_dim) 条件 embedding
"""
# 时间步嵌入
t_emb = self.time_embed(t) # (B, model_channels)
# 初始卷积
h = self.input_conv(x)
# 下采样
hs = []
for block in self.down_blocks:
h = block(h, t_emb)
hs.append(h)
# 中间 + 交叉注意力
h = self.mid_block(h, t_emb, context)
# 上采样 + 跳跃连接
for block, skip in zip(self.up_blocks, reversed(hs)):
h = block(h, t_emb, context, skip)
# 输出
return self.output_conv(h)

4. 整体训练流程#

4.1 LDM 训练的两个阶段#

┌─────────────────────────────────────────────────────────────┐
│ 阶段 A: VAE 预训练 │
│ 数据: 大量图像 │
│ 目标: 训练 VAE 学习感知压缩 │
│ 损失: VAE 重建损失 + KL 散度 │
│ 特点: 与扩散过程独立,可冻结 ❄️ │
├─────────────────────────────────────────────────────────────┤
│ 阶段 B: 扩散训练 (主体) │
│ 数据: 图像 + 条件 (文本/分割/...) │
│ 目标: 在潜空间训练 UNet 学习去噪 │
│ 损失: DDPM 噪声预测损失 │
│ 条件注入: 交叉注意力 (文本 → UNet) │
└─────────────────────────────────────────────────────────────┘
训练技巧:
1. VAE 冻结: 训练扩散时只更新 UNet,不更新 VAE
2. 文本编码器冻结: 只更新 UNet 的交叉注意力层参数
3. 分类器-free 引导: 10% 概率 drop 条件,保持无分类器生成能力

4.2 LDM 训练循环#

def train_ldm(vae, unet, text_encoder, dataloader, optimizer, device):
"""
LDM 训练循环。
"""
vae.eval() # VAE 冻结
text_encoder.eval() # 文本编码器冻结
unet.train() # UNet 训练
for batch in tqdm(dataloader):
images, texts = batch["image"].to(device), batch["text"]
# ============ 1) 编码到潜空间 ============
with torch.no_grad():
# 图像 → VAE → 潜变量
z = vae.encode(images) # (B, 4, 64, 64)
# ============ 2) 采样时间步 ============
batch_size = z.shape[0]
t = torch.randint(0, diffusion_steps, (batch_size,), device=device)
# ============ 3) 加噪 ============
noise = torch.randn_like(z)
alpha_bar = diffusion_schedule[t] # 预设的噪声调度
z_t = alpha_bar.sqrt().view(-1, 1, 1, 1) * z + \
(1 - alpha_bar).sqrt().view(-1, 1, 1, 1) * noise
# ============ 4) 编码文本条件 ============
with torch.no_grad():
context = text_encoder.encode(texts) # (B, seq_len, 768)
# ============ 5) 随机丢弃条件 (Classifier-Free Guidance) ============
drop_mask = torch.rand(batch_size, device=device) < 0.1
context = context.masked_fill(drop_mask.unsqueeze(1), 0.0)
# ============ 6) UNet 去噪预测 ============
noise_pred = unet(z_t, t, context) # (B, 4, 64, 64)
# ============ 7) 计算损失 ============
loss = F.mse_loss(noise_pred, noise)
# ============ 8) 反向传播 ============
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(unet.parameters(), max_norm=1.0)
optimizer.step()
return loss.item()

4.3 推理:文本到图像#

@torch.no_grad()
def generate_with_ldm(vae, unet, text_encoder, prompt, num_inference_steps=50,
guidance_scale=7.5):
"""
LDM 推理: 文本 → 图像。
参数:
prompt: str 文本描述
num_inference_steps: 推理步数 (DDPM 约 1000, DDIM 可降到 20-50)
guidance_scale: CFG 引导强度
"""
vae.eval()
unet.eval()
text_encoder.eval()
device = next(unet.parameters()).device
# ============ 1) 编码文本 ============
context = text_encoder.encode([prompt]) # (1, seq_len, 768)
uncond_context = text_encoder.encode([""]) # (1, seq_len, 768)
# ============ 2) 初始化噪声 ============
B = 1
z_T = torch.randn(B, 4, 64, 64, device=device)
# ============ 3) DDIM 采样 ============
# 使用 DDIM 可大幅减少步数
timesteps = torch.linspace(999, 0, num_inference_steps, device=device).long()
for i, t in enumerate(tqdm(timesteps)):
t_tensor = t.expand(B)
# 条件预测
noise_pred_cond = unet(z_T, t_tensor, context)
# 无条件预测 (CFG)
noise_pred_uncond = unet(z_T, t_tensor, uncond_context)
# CFG 合并
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
# DDIM 去噪步
alpha_t = diffusion_schedule[t]
alpha_prev = diffusion_schedule[timesteps[i - 1]] if i > 0 else 1.0
# 估算 x_0
x0_pred = (z_T - (1 - alpha_t).sqrt() * noise_pred) / alpha_t.sqrt()
# 计算 x_{t-1}
z_T = alpha_prev.sqrt() * x0_pred + (1 - alpha_prev).sqrt() * noise_pred
# ============ 4) 解码到图像 ============
image = vae.decode(z_T)
return image

5. Cross-Attention Networks(CIN):LDM 的核心结构单元#

5.1 什么是 CIN?#

LDM 论文中把集成交叉注意力层的 UNet block 称为 Cross-Attention Networks (CIN)

class CrossAttentionBlock(nn.Module):
"""
Cross-Attention Network (CIN)。
LDM 论文中提出,是 ResNet Block + 交叉注意力的组合。
"""
def __init__(self, channels, context_dim, num_heads=8):
super().__init__()
self.res_block = ResBlock(channels, channels)
self.cross_attn = CrossAttention(channels, context_dim, num_heads)
def forward(self, x, context):
"""
参数:
x: (B, C, H, W) 特征图
context: (B, seq_len, D) 条件 embedding
"""
# 先残差
h = self.res_block(x)
B, C, H, W = h.shape
# Reshape: (B, C, H, W) → (B, H*W, C)
h = h.flatten(2).transpose(1, 2)
# 交叉注意力
h = self.cross_attn(h, context)
# Reshape 回: (B, H*W, C) → (B, C, H, W)
h = h.transpose(1, 2).view(B, C, H, W)
return h + x # 残差连接

5.2 LDM UNet 的完整结构#

class DownBlock(nn.Module):
"""下采样块。"""
def __init__(self, in_ch, out_ch):
super().__init__()
self.res1 = ResBlock(in_ch, out_ch)
self.res2 = ResBlock(out_ch, out_ch)
self.downsample = nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=2, padding=1)
def forward(self, x, t_emb):
x = self.res1(x)
x = self.res2(x)
skip = x
x = self.downsample(x)
return x, skip
class MidBlock(nn.Module):
"""带交叉注意力的中间块。"""
def __init__(self, channels, context_dim):
super().__init__()
self.res1 = ResBlock(channels, channels)
self.cross_attn = CrossAttention(channels, context_dim)
self.res2 = ResBlock(channels, channels)
def forward(self, x, t_emb, context):
x = self.res1(x)
B, C, H, W = x.shape
x = x.flatten(2).transpose(1, 2)
x = self.cross_attn(x, context)
x = x.transpose(1, 2).view(B, C, H, W)
x = self.res2(x)
return x
class UpBlock(nn.Module):
"""上采样块(含跳跃连接和交叉注意力)。"""
def __init__(self, in_ch, out_ch, context_dim):
super().__init__()
self.res1 = ResBlock(in_ch, out_ch)
self.res2 = ResBlock(out_ch, out_ch)
self.upsample = nn.ConvTranspose2d(out_ch, out_ch, kernel_size=4, stride=2, padding=1)
self.cross_attn = CrossAttention(out_ch, context_dim)
def forward(self, x, t_emb, context, skip):
# 上采样
x = self.upsample(x)
# 拼接跳跃连接
x = torch.cat([x, skip], dim=1)
# ResBlock
x = self.res1(x)
x = self.res2(x)
# 交叉注意力
B, C, H, W = x.shape
x_flat = x.flatten(2).transpose(1, 2)
x_flat = self.cross_attn(x_flat, context)
x = x_flat.transpose(1, 2).view(B, C, H, W)
return x

6. Classifier-Free Guidance(CFG)#

6.1 CFG 的原理#

CFG 是 LDM 生成高质量图像的关键技术:

核心思想:
- 训练时: 10% 概率用空文本 (unconditional) 替代条件文本
- 推理时: 同时预测条件输出和无条件输出,用差值增强条件
数学推导:
ε_θ(z_t, ∅) = 无条件噪声预测
ε_θ(z_t, c) = 条件噪声预测
引导后的预测:
ε̃_θ(z_t, c) = ε_θ(z_t, ∅) + w · (ε_θ(z_t, c) - ε_θ(z_t, ∅))
= (1 + w) · ε_θ(z_t, c) - w · ε_θ(z_t, ∅)
其中 w ≥ 1 是引导强度
w = 1: 无引导
w = 7.5: SD 默认值,强条件引导
w → ∞: 等价于用分类器引导(但无需训练分类器)
def cfg_guidance(noise_pred_uncond, noise_pred_cond, guidance_scale=7.5):
"""
Classifier-Free Guidance 合并。
参数:
noise_pred_uncond: (B, 4, H, W) 无条件预测
noise_pred_cond: (B, 4, H, W) 条件预测
guidance_scale: float 引导强度
返回:
引导后的预测
"""
# ε̃ = ε_uncond + w · (ε_cond - ε_uncond)
return noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
# 推理时使用:
for t in timesteps:
# 1) 无条件预测(空文本)
noise_uncond = unet(z_t, t, uncond_emb)
# 2) 条件预测(实际文本)
noise_cond = unet(z_t, t, text_emb)
# 3) CFG 合并
noise_pred = cfg_guidance(noise_uncond, noise_cond, guidance_scale=7.5)

6.2 CFG 的效果#

def cfg_effect():
"""
不同 guidance_scale 的效果。
"""
results = {
"w = 1.0 (无引导)": {
"图像质量": "模糊、缺少细节",
"文本相关性": "低",
"多样性": "高",
},
"w = 3.0": {
"图像质量": "中等",
"文本相关性": "中低",
"多样性": "较高",
},
"w = 7.5 (SD 默认)": {
"图像质量": "清晰、细节丰富",
"文本相关性": "高",
"多样性": "适中",
},
"w = 12.0": {
"图像质量": "过于锐利、可能过饱和",
"文本相关性": "极高(可能夸张)",
"多样性": "低",
},
"w = 20.0": {
"图像质量": "失真、artifacts",
"文本相关性": "过度强调",
"多样性": "极低",
},
}
return results

7. Stable Diffusion 版本演进#

7.1 SD 1.x → SD 2.x → SDXL#

def sd_version_comparison():
"""
Stable Diffusion 版本对比。
"""
versions = {
"SD 1.x": {
"发布时间": "2022.08",
"基础模型": "LDM-4 (latent 64×64×4)",
"文本编码器": "OpenCLIP ViT-L/14 (冻结)",
"UNet 参数量": "~860M",
"分辨率": "512×512",
"训练数据": "LAION-5B 筛选 ( Aesthetic score > 5.0)",
"改进": "简单高效,开源社区广泛使用",
},
"SD 2.x": {
"发布时间": "2022.11",
"基础模型": "latent 64×64×4 (768×768)",
"文本编码器": "OpenCLIP ViT-L/14 + 自训练版本",
"UNet 参数量": "~1B (更深)",
"分辨率": "768×768 (基础)",
"训练数据": "LAION-5B 筛选 + NSFW 过滤",
"关键改进": "更高的基础分辨率,更好的 prompt following",
},
"SDXL": {
"发布时间": "2023.07",
"基础模型": "latent 128×128×4 (1024×1024)",
"文本编码器": "OpenCLIP ViT-L/14 + OpenCLIP ViT-G/14 (双编码器)",
"UNet 参数量": "~3.5B (两阶段: base + refiner)",
"分辨率": "1024×1024 (基础) + 更高分辨率 (refiner)",
"训练数据": "高质量数据集,审美分数筛选",
"关键改进": "双文本编码器 + refiner + T5 XXL 文本编码器",
},
}
return versions

7.2 SD 1.x 详细配置#

@dataclass
class SD1Config:
"""Stable Diffusion 1.x 配置。"""
# VAE
latent_channels: int = 4
latent_size: int = 64 # 512 / 8
# UNet
in_channels: int = 4
model_channels: int = 320 # base channel dimension
out_channels: int = 4
num_res_blocks: int = 2
attention_resolutions: tuple = (4, 2, 1) # 在哪些分辨率加注意力
channel_mult: tuple = (1, 2, 4, 4)
# 交叉注意力
num_heads: int = 8
num_heads_ema: int = -1 # 使用 num_heads
# 文本编码器
context_dim: int = 768 # CLIP ViT-L/14
use_clip: bool = True
# 训练
image_size: int = 512
num_classes: int = -1 # 无类别条件
def channel_list(self):
"""每层的通道数。"""
return [self.model_channels * m for m in self.channel_mult]

7.3 SDXL 的关键改进#

def sdxl_improvements():
"""
SDXL 相比 SD 1.x 的关键改进。
"""
improvements = {
"双文本编码器": {
"编码器1": "OpenCLIP ViT-L/14 (768d)",
"编码器2": "OpenCLIP ViT-G/14 (1280d)",
"融合": "concat 后投影到 2048d",
"效果": "更强的文本理解,支持更长的文本描述",
},
"更大的潜空间": {
"SD 1.x": "64×64×4 = 16384 维度",
"SDXL": "128×128×4 = 65536 维度 (4倍)",
"效果": "更高的基础分辨率 (1024px vs 512px)",
},
"两阶段生成": {
"Base": "基础 UNet 生成 1024px",
"Refiner": "额外 UNet 做高分辨率 refinement",
"效果": "更精细的细节和纹理",
},
"改进的 UNet": {
"残差块": "更多的残差连接",
"注意力": "64×64 → 32×32 → 16×16 → 8×8 多分辨率",
"参数": "860M → 3.5B (含 refiner)",
},
}
return improvements

8. LDM 的条件控制扩展#

8.1 文生图(Text-to-Image)#

def text2image_ldm():
"""
LDM 文本到图像生成。
条件: CLIP 文本 embedding
"""
return {
"条件类型": "文本",
"条件编码器": "CLIP Text Encoder",
"条件维度": "768 (ViT-L/14)",
"注入方式": "交叉注意力 (Query: UNet, Key/Value: Text)",
"CFG": True,
"典型引导强度": "7.5",
}
def semantic_segmentation_control():
"""
语义分割控制。
条件: 语义分割图
"""
return {
"条件类型": "分割图",
"条件编码器": "语义分割模型 (如 Segmenter)",
"条件维度": "与潜变量 H×W 匹配",
"注入方式": "拼接 (concatenate) 到潜变量通道",
"CFG": True,
}

8.2 图像变换(Image-to-Image)#

def image2image_ldm():
"""
Image-to-Image (Img2Img)。
用 LDM 的加噪-去噪过程实现图像变换。
"""
return {
"原理": "给定初始图像 + 文本,采样出符合文本的变体",
"步骤": [
"1. 用 VAE 把输入图像编码到潜空间: z_0",
"2. 按 strength 参数加噪: z_t = α·z_0 + (1-α)·noise",
"3. 以 z_t 为起点,文本为条件,执行扩散去噪",
"4. 用 VAE 解码 z_0' → 图像",
],
"参数": {
"strength": "0.0(保持原图) - 1.0(完全重绘)",
"0.2": "轻微变化",
"0.5": "适度变化",
"0.8": "大幅变化",
},
}
def inpainting_ldm():
"""
图像修复 (Inpainting)。
根据 mask 重新生成指定区域。
"""
return {
"输入": "原图 + mask + 文本描述",
"方法": [
"1. 原图 VAE 编码 → z_0",
"2. mask 下采样 → 二值 mask_m (64×64)",
"3. 在 mask 区域注入噪声: z_t[mask] = noise",
"4. UNet 输入: (z_t, mask_m, text_emb)",
"5. 仅在 mask 区域计算损失",
"6. VAE 解码 → 修复后图像",
],
}

8.3 超分辨率(Super-Resolution)#

def latent_upscaler():
"""
潜在空间超分辨率。
LDM 的一个扩展应用。
"""
return {
"低分辨率输入": "64×64×4 潜变量",
"方法": "在潜空间做上采样 + 去噪",
"UNet 修改": "输入通道从 4 → 8 (原潜变量 + 插值潜变量)",
"效果": "4× 超分 (64→256)",
"后续": "配合 VAE 解码 → 2048×2048 输出",
}

9. 完整实现#

import torch
import torch.nn as nn
import torch.nn.functional as F
from dataclasses import dataclass
import math
# ============ 时间步嵌入 ============
class SinusoidalPositionEmbedding(nn.Module):
"""正弦位置嵌入。"""
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, t):
device = t.device
half_dim = self.dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
emb = t[:, None] * emb[None, :]
emb = torch.cat([emb.sin(), emb.cos()], dim=-1)
return emb
class TimeEmbedding(nn.Module):
"""时间步嵌入(MLP 版本)。"""
def __init__(self, dim):
super().__init__()
self.mlp = nn.Sequential(
SinusoidalPositionEmbedding(dim),
nn.Linear(dim, dim * 4),
nn.SiLU(),
nn.Linear(dim * 4, dim),
)
def forward(self, t):
return self.mlp(t)
# ============ ResBlock ============
class ResBlock(nn.Module):
def __init__(self, channels, out_channels=None):
super().__init__()
out_channels = out_channels or channels
self.norm1 = nn.GroupNorm(32, channels)
self.conv1 = nn.Conv2d(channels, out_channels, 3, padding=1)
self.norm2 = nn.GroupNorm(32, out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
self.shortcut = nn.Conv2d(channels, out_channels, 1) if channels != out_channels else nn.Identity()
def forward(self, x):
h = self.norm1(x)
h = F.silu(h)
h = self.conv1(h)
h = self.norm2(h)
h = F.silu(h)
h = self.conv2(h)
return h + self.shortcut(x)
# ============ 交叉注意力 ============
class CrossAttention(nn.Module):
def __init__(self, query_dim, context_dim=None, num_heads=8):
super().__init__()
self.num_heads = num_heads
self.head_dim = query_dim // num_heads
self.scale = self.head_dim ** -0.5
self.to_q = nn.Linear(query_dim, query_dim)
self.to_k = nn.Linear(context_dim or query_dim, query_dim)
self.to_v = nn.Linear(context_dim or query_dim, query_dim)
self.to_out = nn.Linear(query_dim, query_dim)
def forward(self, x, context):
B, N, D = x.shape
h = self.num_heads
d = self.head_dim
q = self.to_q(x).view(B, N, h, d).transpose(1, 2)
k = self.to_k(context).view(B, -1, h, d).transpose(1, 2)
v = self.to_v(context).view(B, -1, h, d).transpose(1, 2)
attn = (q @ k.transpose(-2, -1)) * self.scale
attn = F.softmax(attn, dim=-1)
out = (attn @ v).transpose(1, 2).reshape(B, N, D)
return self.to_out(out)
class TransformerBlock(nn.Module):
"""Self-Attention + Cross-Attention Block。"""
def __init__(self, channels, context_dim, num_heads=8):
super().__init__()
self.norm1 = nn.LayerNorm(channels)
self.attn1 = nn.MultiheadAttention(channels, num_heads, batch_first=True)
self.norm2 = nn.LayerNorm(channels)
self.cross_attn = CrossAttention(channels, context_dim, num_heads)
def forward(self, x, context):
x = x + self.attn1(self.norm1(x), self.norm1(x), self.norm1(x))[0]
x = x + self.cross_attn(self.norm2(x), context)
return x
# ============ UNet ============
class UNetLDM(nn.Module):
"""
LDM 的完整 UNet。
支持交叉注意力条件注入。
"""
def __init__(self, latent_channels=4, model_channels=320,
context_dim=768, num_heads=8):
super().__init__()
self.latent_channels = latent_channels
self.model_channels = model_channels
# 时间步嵌入
self.time_embed = TimeEmbedding(model_channels)
# 输入投影
self.input_conv = nn.Conv2d(latent_channels, model_channels, 3, padding=1)
# 通道列表
ch = [1, 2, 4, 4] # channel multipliers
self chs = [model_channels * c for c in ch] # [320, 640, 1280, 1280]
# 下采样
self.down_blocks = nn.ModuleList()
self.down_transformers = nn.ModuleList()
for i in range(len(ch)):
in_ch = self.chs[i - 1] if i > 0 else model_channels
out_ch = self.chs[i]
self.down_blocks.append(nn.ModuleList([
ResBlock(in_ch, out_ch),
ResBlock(out_ch, out_ch),
]))
if i >= 1: # 在后半段加 transformer
self.down_transformers.append(
TransformerBlock(out_ch, context_dim, num_heads)
)
else:
self.down_transformers.append(nn.Identity())
# 中间层
self.mid_block1 = ResBlock(self.chs[-1], self.chs[-1])
self.mid_transformer = TransformerBlock(self.chs[-1], context_dim, num_heads)
self.mid_block2 = ResBlock(self.chs[-1], self.chs[-1])
# 上采样
self.up_blocks = nn.ModuleList()
self.up_transformers = nn.ModuleList()
for i in range(len(ch)):
skip_ch = self.chs[-(i + 1)]
out_ch = self.chs[-(i + 2)] if i < len(ch) - 1 else model_channels
self.up_blocks.append(nn.ModuleList([
ResBlock(skip_ch * 2, out_ch),
ResBlock(out_ch, out_ch),
]))
self.up_transformers.append(
TransformerBlock(out_ch, context_dim, num_heads)
)
# 输出
self.output_conv = nn.Sequential(
ResBlock(model_channels, model_channels),
nn.Conv2d(model_channels, latent_channels, 3, padding=1),
)
def forward(self, x, t, context):
"""
参数:
x: (B, 4, 64, 64) 潜变量
t: (B,) 时间步
context: (B, seq_len, context_dim) 条件
"""
# 时间嵌入
t_emb = self.time_embed(t)
# 输入
h = self.input_conv(x)
# 下采样
hs = []
for i, (blocks, transformer) in enumerate(zip(
zip(self.down_blocks[::2], self.down_blocks[1::2]),
self.down_transformers
)):
for block in blocks:
h = block(h)
if i < len(self.down_transformers) - 1:
# Transformer: 需要 reshape
B, C, H, W = h.shape
h_flat = h.flatten(2).transpose(1, 2)
h_flat = transformer(h_flat, context)
h = h_flat.transpose(1, 2).view(B, C, H, W)
else:
h = transformer(h, context) if not isinstance(transformer, nn.Identity) else h
hs.append(h)
h = F.avg_pool2d(h, 2)
# 中间
h = self.mid_block1(h)
B, C, H, W = h.shape
h = self.mid_transformer(h.flatten(2).transpose(1, 2), context)
h = h.transpose(1, 2).view(B, C, H, W)
h = self.mid_block2(h)
# 上采样
for i, (blocks, transformer) in enumerate(zip(
zip(self.up_blocks[::2], self.up_blocks[1::2]),
self.up_transformers
)):
h = F.interpolate(h, scale_factor=2, mode="nearest")
h = torch.cat([h, hs.pop()], dim=1)
for block in blocks:
h = block(h)
B, C, H, W = h.shape
h = transformer(h.flatten(2).transpose(1, 2), context)
h = h.transpose(1, 2).view(B, C, H, W)
return self.output_conv(h)
class LDM(nn.Module):
"""
完整的 Latent Diffusion Model。
"""
def __init__(self, latent_channels=4, model_channels=320, context_dim=768):
super().__init__()
self.vae = LDMVAE(latent_channels)
self.unet = UNetLDM(latent_channels, model_channels, context_dim)
self.text_encoder = CLIPTextEncoder()
self.diffusion_steps = 1000
def training_parameters(self):
"""只训练 UNet,冻结 VAE 和文本编码器。"""
vae_params = list(self.vae.parameters())
text_params = list(self.text_encoder.parameters())
unet_params = list(self.unet.parameters())
for p in vae_params + text_params:
p.requires_grad = False
return self.unet.parameters()
@torch.no_grad()
def encode(self, x):
"""图像 → 潜变量。"""
return self.vae.encode(x)
@torch.no_grad()
def decode(self, z):
"""潜变量 → 图像。"""
return self.vae.decode(z)
def forward(self, x, texts, t=None):
"""
训练前向。
"""
# 编码
z = self.encode(x)
# 时间步
if t is None:
t = torch.randint(0, self.diffusion_steps, (x.shape[0],), device=x.device)
# 加噪
noise = torch.randn_like(z)
alpha_bar = self._get_alpha_bar(t)
z_noisy = alpha_bar.sqrt().view(-1, 1, 1, 1) * z + \
(1 - alpha_bar).sqrt().view(-1, 1, 1, 1) * noise
# 文本编码
context = self.text_encoder.encode(texts)
# 预测噪声
noise_pred = self.unet(z_noisy, t, context)
# 损失
loss = F.mse_loss(noise_pred, noise)
return {"loss": loss}
def _get_alpha_bar(self, t):
"""简化的噪声调度。"""
steps = self.diffusion_steps
s = 0.008
t_normalized = t.float() / steps
alpha_bar = torch.cos((t_normalized + s) / (1 + s) * torch.pi * 0.5) ** 2
return alpha_bar

10. LDM 的影响与遗产#

10.1 LDM 对生成式 AI 的贡献#

def ldm_impact():
"""
LDM 对生成式 AI 的核心贡献。
"""
contributions = {
"效率突破": {
"贡献": "把高分辨率图像生成从 20+ 分钟降到 1-5 秒",
"方法": "潜空间压缩 48 倍",
"影响": "使消费级 GPU 本地运行成为可能",
},
"开源生态": {
"贡献": "Stable Diffusion 完全开源",
"方法": "模型权重、代码、训练流程全部公开",
"影响": "催生了 AUTOMATIC1111 WebUI、ComfyUI 等生态",
},
"多条件控制": {
"贡献": "统一的交叉注意力条件注入框架",
"方法": "任意条件编码器 → 交叉注意力 → UNet",
"影响": "ControlNet、T2I-Adapter 等控制方法的基础",
},
"学术影响": {
"贡献": "CVPR 2022 Oral, 10000+ 引用",
"方法": "感知压缩 + 语义压缩的理论框架",
"影响": "DiT、SDXL 等后续工作的基础",
},
}
return contributions

10.2 LDM 催生的重要应用#

开源社区生态:
├── Stable Diffusion WebUI (AUTOMATIC1111) — 最大的开源 GUI
├── ComfyUI — 节点式工作流(更灵活)
├── ControlNet — 姿态/边缘/深度图控制
├── LoRA — 低秩适配器(个性化风格)
├── Dreambooth — 个性化主体训练
├── SDXL — 更强的 1024px 生成
└── Stable Diffusion 3 — MM-DiT 架构
商业应用:
├── Adobe Firefly — Adobe 的生成式 AI
├── Runway ML — 视频生成
├── Midjourney — AI 艺术创作
└── 无数图像生成 SaaS 服务

11. 总结#

11.1 核心要点#

维度关键要点
核心创新感知压缩(VAE)+ 语义压缩(DDPM)的两阶段压缩框架
效率提升潜空间 48× 压缩,计算量降低两个数量级
条件注入交叉注意力机制,任意条件(文本/分割/边缘)可插拔
CFGClassifier-Free Guidance,无需分类器即可实现条件引导
开源影响Stable Diffusion 开源催生了整个 AI 图像生成生态
版本演进SD 1.x → SD 2.x → SDXL (分辨率/文本编码器/双阶段)
后续发展DiT 替代 UNet → SD3 (MM-DiT)

11.2 LDM vs Pixel Diffusion#

DDPM (像素扩散) LDM (潜在扩散)
效率 低 (512²) 高 (64², 48×压缩)
质量 高 (无损) 高 (VAE 重建损失可接受)
条件控制 像素级条件 语义级条件注入
CFG 像素级 潜空间
应用 低分辨率 高分辨率 (512/768/1024)
代表 DALL-E 2 (无开源) Stable Diffusion (开源)

11.3 一句话总结#

LDM 的核心洞察是把扩散过程分为两个层次:低层的感知细节(纹理、颜色、边缘)交给 VAE 高效压缩和重建,高层的语义内容(物体、布局、风格)交给扩散模型在低维潜空间里学习。通过在潜空间中执行 DDPM 去噪,LDM 把高分辨率图像生成的计算量降低了 48 倍,同时通过交叉注意力机制支持灵活的多条件控制——这就是 Stable Diffusion 能够在消费级 GPU 上运行、让开源社区掀起 AI 图像革命的技术基础。

11.4 推荐资源#

论文:
- LDM (Rombach et al., 2022): "High-Resolution Image Synthesis with Latent Diffusion Models"
- SDXL (Podell et al., 2023): "SDXL: Improving Latent Diffusion Models for High Resolution Image Synthesis"
代码:
- CompVis/stable-diffusion (官方实现)
- AUTOMATIC1111/stable-diffusion-webui (WebUI)
- comfyanonymous/ComfyUI (节点式工作流)
相关:
- ControlNet (Zhang et al., 2023): 条件控制扩展
- DiT (Peebles & Xie, 2023): UNet → Transformer 的后续演进
- SD3/MM-DiT: Stable Diffusion 3 的新架构

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

Latent Diffusion Model 深度剖析:从像素扩散到潜在空间生成
https://aiattnstudio.link/posts/latent-diffusion-model/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签