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 可以高效压缩,不丢失感知质量 ││ 高层语义才需要扩散过程 ││ → 潜空间扩散专注于语义生成 ││ │└─────────────────────────────────────────────────────────────┘@dataclassclass 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 的感知压缩层,有三个关键原因:
@dataclassclass 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 configs3. 条件机制:交叉注意力
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 embeddings3.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 image5. 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 x6. 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 results7. 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 versions7.2 SD 1.x 详细配置
@dataclassclass 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 improvements8. 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 torchimport torch.nn as nnimport torch.nn.functional as Ffrom dataclasses import dataclassimport 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_bar10. 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 contributions10.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× 压缩,计算量降低两个数量级 |
| 条件注入 | 交叉注意力机制,任意条件(文本/分割/边缘)可插拔 |
| CFG | Classifier-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/
