Consistency Models 深度解析:一步生成扩散模型
1. 从多步采样到一步生成
1.1 扩散模型的核心瓶颈
DDPM 需要 100-1000 步才能生成一张高质量图像。即使是加速后的 DDIM 也需要 20-50 步。每次生成都意味着:
DDPM 多步采样的代价:
模型前向: 1000 步 × N 层 × 复杂度 = 高昂计算成本 延迟: 每步 ~50ms → 1000步 ≈ 50秒 → 无法实时 显存: 所有中间激活值需要保存 能耗: 移动端 / 边缘部署几乎不可能
实际需求: ✓ "我要生成一张猫的照片" → 等待 50 秒? ✗ "我要实时生成视频帧" → 根本不可能 ✗ "在手机上运行 SD" → 显存和延迟双重爆炸1.2 加速采样的三条技术路线
┌──────────────────────────────────────────────────────────────┐│ 扩散模型加速采样的三条路线 │├──────────────────────────────────────────────────────────────┤│ ││ 路线 1: 更少的采样步数 (DDIM / PNDM) ││ ──────────────────────────────────────────────────────── ││ 思想: 跳过部分中间步,通过 ODE 插值重建 ││ 效果: 50-100 步仍可接受,20 步开始明显质量下降 ││ 缺点: 仍有固定步数,延迟有下限 ││ ││ 路线 2: 预测更少的变量 (Latent Consistency Models / LCM) ││ ──────────────────────────────────────────────────────── ││ 思想: 预测 x_0 或一致性函数,减少步数 ││ 效果: 2-8 步生成,质量接近全步 ││ 缺点: 需要额外引导,对某些模型结构有依赖 ││ ││ 路线 3: 一步直接生成 (Consistency Models) ││ ──────────────────────────────────────────────────────── ││ 思想: 训练一个模型,直接从噪声一步映射到数据 ││ 效果: 1 步生成,质量通过蒸馏提升 ││ 缺点: 需要从预训练扩散模型蒸馏或从零训练 ││ │└──────────────────────────────────────────────────────────────┘1.3 Consistency Models 的核心洞察
Song 等人 (OpenAI) 的核心洞察是:DDPM 的前向过程和反向过程本质上是一条 ODE 轨迹——任意时间步的 x 都沿着这条轨迹流向数据分布或噪声分布。既然如此,我们可以直接学习这条轨迹的”一致性函数”:轨迹上的任意点,都应该映射到轨迹的起点(即数据分布)。这就是 Consistency Models 的核心思想——训练一个函数 f,使得 f(x_t) = x_0 对所有 t ∈ [0, T] 成立。
1.4 论文信息
论文: "Consistency Models"作者: Yang Song, Prafulla Dhariwal, Mark Chen, Ilya Sutskever单位: OpenAI发表于: ICML 2023引用: > 2,000 次开源: https://github.com/openai/consistency_models2. 核心思想:ODE 轨迹与一致性映射
2.1 扩散模型的 ODE 视角
DDPM 的前向过程 q(x_t | x_0) 可以被等价的 ODE(随机微分方程)描述:
def diffusion_ode(): """ 扩散模型的随机微分方程 (SDE) 视角。
关键洞察: 前向过程和反向过程都在同一条轨迹上移动。 """
return { "前向 SDE (添加噪声)": { "方程": "dx = σ_max^t · dW_t", "解释": " Wiener 过程 W_t 驱动 x 趋向噪声", "边界": "x_0 → x_T (T → ∞ 时,x_T ≈ N(0, σ_max²))", }, "反向 ODE (生成采样)": { "方程": "dx = [x - ∇_x log p_θ(x, t)] dt", "解释": "得分函数 ∇ log p 指向数据密度梯度", "关键": "这是确定性路径,从噪声流向数据", }, "轨迹一致性": { "观察": "反向 ODE 定义了一条从 x_T 到 x_0 的路径", "性质": "轨迹上任意点 x_t 都是同一个扩散过程的"快照"", "需求": "我们希望一个函数 f 满足: f(x_t) = x_0 对所有 t", }, }
def consistency_mapping(): """ 一致性映射。
目标: 学习一个函数 f_θ 使得: f_θ(x_t) = x_0 对所有 t ∈ [0, T] 其中 x_0 和 x_t 是同一轨迹上的点。
这意味着: - f_θ 沿着 ODE 轨迹是"一致的" - 轨迹上的所有点映射到同一点 (轨迹起点) - 一旦学会,从任意中间步 x_t 一步到 x_0 """
return { "数学定义": "f_θ(x_t) = x_0", "物理含义": "轨迹上任意时刻的状态都"记住"了起点", "推理优势": "从任意 x_t,直接一步到达 x_0", "训练挑战": "x_0 是未知的 (我们只有噪声样本)", }2.2 一致性函数的数学推导
def consistency_function_derivation(): """ 一致性函数的数学推导。
核心问题: 如何训练 f_θ 使得 f_θ(x_t) ≈ x_0?
方法 1: 自监督学习 ───────────────────────────────────────────── 1. 采样 x_0 ~ p_data (真实数据) 2. 采样 t_1, t_2 ~ Uniform(ε, T), 其中 t_1 < t_2 3. 通过前向过程: x_{t_1} = α_{t_1} x_0 + σ_{t_1} ε_1 x_{t_2} = α_{t_2} x_0 + σ_{t_2} ε_2 4. 损失: |f_θ(x_{t_1}) - x_{t_2}|²
关键观察: - 如果 t_2 → 0,则 x_{t_2} ≈ x_0 - f_θ(x_{t_1}) 应该接近 x_{t_2} (都接近 x_0) - 不需要知道真实的 x_0!
方法 2: 蒸馏 (从预训练扩散模型) ───────────────────────────────────────────── 1. 使用预训练的 DDPM 作为教师 2. 采样 x_{t+ε},用教师模型计算 x_0^teacher 3. 损失: |f_θ(x_t) - x_0^teacher|²
方法 3: TTCT (Two-Track Consistency Training) ───────────────────────────────────────────── 结合方法 1 和 2,用更强的教师信号 """
return { "TTCT 损失": { "公式": "L = λ_1 · |f_θ(x_t) - f_φ(x_{t+ε})|² + λ_2 · |f_θ(x_t) - x_0^teacher|²", "第一项": "自监督一致性 (不依赖教师)", "第二项": "蒸馏一致性 (利用预训练知识)", "优点": "既能从零训练,又保留蒸馏的效率", }, }2.3 EDM 框架整合
Consistency Models 被设计为可以直接在 EDM(Elucidating the Design Space of Diffusion Models)框架上工作:
def edm_integration(): """ Consistency Models 与 EDM 的整合。
EDM 的核心设计: - 使用网络输出 c_out = σ_data · F_θ(x, t) - 其中 F_θ 是 UNet,σ_data 是数据标准差 - 采样使用 EDM 的 ODE 求解器
Consistency Model 的 EDM 版本: - f_θ(x_t, t) 直接输出 x_0 - 网络: F_θ(x, t) 输出 c_skip(t)·x + c_out(t)·F_θ(x, t) - 其中 c_skip 和 c_out 是时间依赖的跳跃连接权重
优势: - 与 SD 的 EDM 实现兼容 - 可以直接加载 SD 的预训练权重进行蒸馏 """
return { "网络输出": { "标准 EDM": "c_out = σ_data · F_θ(x, t)", "Consistency Model": "x_0 = c_skip(t)·x + c_out(t)·F_θ(x, t)", }, "时间依赖权重": { "c_skip(t)": "控制输入 x 的残差比例", "c_out(t)": "控制网络输出的比例", "设计": "与 EDM 相同,通过 MLP 从 t 计算", }, }3. 训练方法:TTCT 双轨训练
3.1 TTCT 详解
TTCT(Two-Track Consistency Training)是 Consistency Models 提出的双轨训练方法:
TTCT 训练流程:
轨道 1: 自监督一致性 ─────────────────────────────────────────────────────── x_0 ~ p_data ↓ t_1, t_2 ~ Uniform(ε, T), t_1 < t_2 ↓ x_{t_1} = α_{t_1} x_0 + σ_{t_1} ε_1 x_{t_2} = α_{t_2} x_0 + σ_{t_2} ε_2 ↓ 损失: L_1 = |f_θ(x_{t_1}) - x_{t_2}|²
轨道 2: 蒸馏一致性 (从预训练教师) ─────────────────────────────────────────────────────── x_ε ~ N(0, σ²) ↓ 用预训练教师模型计算 x_0^teacher ↓ 损失: L_2 = |f_θ(x_ε) - x_0^teacher|²
总损失: L = λ_1 · L_1 + λ_2 · L_2def ttct_loss(): """ TTCT 损失函数实现。 """ def compute_ttct_loss(student_model, teacher_model, p_data, optimizer, config): """ 参数: student_model: 要训练的 Consistency Model teacher_model: 预训练的 DDPM (用于蒸馏) p_data: 真实数据分布 config: 训练配置 """ # 采样一个 batch x_0 = sample_batch(p_data, batch_size=config.batch_size)
# ============ 轨道 1: 自监督一致性 ============ # 采样两个时间步 epsilon = config.epsilon # 最小时间步 T = config.T # 最大时间步
t_1 = torch.rand(batch_size) * (T - epsilon) + epsilon t_2 = torch.rand(batch_size) * (T - t_1) + t_1 # t_2 > t_1
# 加噪 noise_1 = torch.randn_like(x_0) noise_2 = torch.randn_like(x_0)
x_t1 = alpha(t_1).view(-1, 1, 1, 1) * x_0 + \ sigma(t_1).view(-1, 1, 1, 1) * noise_1
x_t2 = alpha(t_2).view(-1, 1, 1, 1) * x_0 + \ sigma(t_2).view(-1, 1, 1, 1) * noise_2
# 学生模型预测 pred_t1 = student_model(x_t1, t_1)
# 自监督损失: f_θ(x_{t_1}) 应该接近 x_{t_2} # (因为两者都接近 x_0) loss_1 = F.mse_loss(pred_t1, x_t2.detach())
# ============ 轨道 2: 蒸馏一致性 ============ # 采样噪声 x_epsilon = torch.randn_like(x_0) * sigma(epsilon)
# 教师模型预测 (全步 DDPM) with torch.no_grad(): x_0_teacher = teacher_model(x_epsilon, epsilon)
# 学生模型预测 pred_eps = student_model(x_epsilon, epsilon)
# 蒸馏损失 loss_2 = F.mse_loss(pred_eps, x_0_teacher.detach())
# ============ 总损失 ============ loss = config.lambda_1 * loss_1 + config.lambda_2 * loss_2
return loss, {"self_supervised": loss_1, "distillation": loss_2}
return compute_ttct_loss3.2 从零训练 vs 蒸馏
def training_modes(): """ Consistency Models 的两种训练模式。 """ return { "Training-from-Scratch": { "方法": "不使用预训练教师,只用自监督一致性", "损失": "L = |f_θ(x_{t_1}) - x_{t_2}|²", "优点": "不需要预训练扩散模型", "缺点": "需要更多数据和训练时间", "适用": "从头构建新模型", "收敛": "慢,但最终质量可达", }, "SCT (Sampling Consistency Training)": { "方法": "从零训练 + 使用单步学生作为教师", "损失": "L = |f_θ(x_t) - f_φ(x_{t+Δ})|²", "优点": "渐进提升,不需要预训练教师", "缺点": "需要迭代训练教师-学生", "适用": "持续改进模型", }, "CFG (Consistency Fine-tuning)": { "方法": "在预训练 CM 基础上微调,加入 CFG", "损失": "L = |f_θ(x_t, c) - x_0^teacher(x_t, c)|²", "优点": "继承预训练 CM + 支持条件生成", "缺点": "需要额外的条件编码器", "适用": "快速适配下游任务", }, "CD (Consistency Distillation)": { "方法": "从预训练扩散模型蒸馏", "损失": "L = |f_θ(x_t) - x_0^teacher(x_t)|²", "优点": "快速收敛,质量高", "缺点": "需要预训练扩散模型", "适用": "已有 SD 模型,想加速采样", }, }4. 网络架构
4.1 与 EDM UNet 的兼容性
Consistency Models 直接复用 EDM 的 UNet 架构,只需修改输出头:
class ConsistencyModelEDM(nn.Module): """ 基于 EDM UNet 的 Consistency Model。
核心修改: - 网络输出改为预测 x_0 (而非噪声或 v-prediction) - 添加时间依赖的 c_skip 和 c_out 跳跃连接 """
def __init__(self, unet_backbone): super().__init__() self.unet = unet_backbone
# 时间嵌入 → 跳跃连接权重 self.c_skip_net = nn.Sequential( nn.Linear(256, 1), nn.Sigmoid(), ) self.c_out_net = nn.Sequential( nn.Linear(256, 1), nn.Sigmoid(), )
def forward(self, x, t): """ 参数: x: (B, C, H, W) 输入 (可以是任意时间步的 x_t) t: (B,) 时间步 (0 到 T) 返回: x_0: (B, C, H, W) 预测的原始数据 """ # 时间嵌入 t_emb = get_timestep_embedding(t, 128) t_emb = self.unet.time_embed(t_emb) # (B, 256)
# EDM UNet 前向 h = self.unet.prep(x) h = self.unet.downsample(h) h = self.unet.attention(h, t_emb) h = self.unet.upsample(h)
# EDM 输出: F_θ(x, t) f_theta = self.unet.final(h) # (B, C, H, W)
# 时间依赖跳跃连接 c_skip = self.c_skip_net(t_emb) # (B, 1) ∈ (0, 1) c_out = self.c_out_net(t_emb) # (B, 1) ∈ (0, 1)
# Consistency Model 输出: x_0 = c_skip · x + c_out · F_θ x_0 = c_skip.view(-1, 1, 1, 1) * x + \ c_out.view(-1, 1, 1, 1) * f_theta
return x_04.2 边界条件处理
def boundary_conditions(): """ Consistency Models 的边界条件。
t = 0 时: - x_t = x_0 (无噪声) - f_θ(x_0, 0) 应该 = x_0 (恒等映射)
t = T 时: - x_T ≈ 纯噪声 - f_θ(x_T, T) 应该 ≈ x_0 (仍然预测数据) """
return { "t = 0 (边界)": { "条件": "f_θ(x, 0) = x", "实现": "强制 c_skip(0) = 1, c_out(0) = 0", "原因": "此时 x 就是真实数据,不需要预测", }, "t = T (边界)": { "条件": "f_θ(x, T) = E[x_0 | x_T = x]", "实现": "网络在高噪声时预测数据分布的均值", "原因": "从纯噪声中重建数据", }, "中间 t": { "条件": "f_θ(x_t, t) 应该介于 x_t 和 x_0 之间", "实现": "c_skip(t) 和 c_out(t) 通过 MLP 学习", "观察": "t 越大,c_skip 越小(x_t 越偏离 x_0)", }, }5. 推理:多步生成与 TTA
5.1 基本推理流程
@torch.no_grad()def cm_sample(model, x_T, num_steps=1, guide_scale=1.0): """ Consistency Model 推理。
参数: model: 训练好的 Consistency Model x_T: (B, C, H, W) 初始噪声 num_steps: 采样步数 (1 = 一步生成) guide_scale: CFG 引导强度 返回: x_0: (B, C, H, W) 生成的图像 """
if num_steps == 1: # ============ 一步生成 (核心) ============ t = torch.full((x_T.shape[0],), 1.0, device=x_T.device)
if guide_scale == 1.0: # 无条件生成 x_0 = model(x_T, t) else: # CFG 引导 # 需要条件和无条件输入 x_0_cond = model(x_T, t, condition) # 条件输入 x_0_uncond = model(x_T, t, uncondition) # 无条件输入 x_0 = x_0_uncond + guide_scale * (x_0_cond - x_0_uncond)
return x_0
else: # ============ 多步生成 (TTA 增强) ============ return cm_sample_tta(model, x_T, num_steps)
@torch.no_grad()def cm_sample_tta(model, x_T, num_steps=4): """ TTST (Test-Time Self-Taught) 多步增强。
核心思想: - 多步生成比单步更稳定 - 但步数仍然远少于 DDPM - 通过 self-distillation 持续改进 """
# 初始化 x = x_T.clone() T = 1.0
# 时间步调度 (从大到小) timesteps = torch.linspace(T, 0, num_steps + 1).tolist()
for i in range(num_steps): t = timesteps[i] t_next = timesteps[i + 1]
# 单步预测 x_0_pred = model(x, torch.full((x.shape[0],), t, device=x.device))
# 如果 t_next = 0,则这就是最终输出 if t_next == 0: return x_0_pred
# 否则,用预测的 x_0 重建下一步 # x_{t_next} = α_{t_next} · x_0 + σ_{t_next} · noise # 其中 noise = (x - α_t · x_0) / σ_t alpha_t = alpha(torch.tensor(t)) sigma_t = sigma(torch.tensor(t)) alpha_next = alpha(torch.tensor(t_next)) sigma_next = sigma(torch.tensor(t_next))
noise = (x - alpha_t * x_0_pred) / sigma_t.clamp(min=1e-6) x = alpha_next * x_0_pred + sigma_next * noise
return x5.2 TTST 推理详解
def ttst_explained(): """ TTST (Test-Time Self-Taught) 推理详解。
目标: 用多步推理提升单步 CM 的质量。
原理: 1. 从噪声 x_T 开始 2. 用 CM 预测 x_0 3. 用预测的 x_0 重建中间步 x_{t-Δ} 4. 再从 x_{t-Δ} 预测 x_0 5. 重复 N 次,质量逐步提升
为什么有效: - 每一步都在"纠正"上一步的误差 - 与 ODE 求解器的思想类似 - 但不需要梯度,计算量小 """
return { "单步 TTST": { "步骤": "x_T → f_θ(x_T) = x_0", "质量": "★★★☆☆", "速度": "★★★★★", }, "4步 TTST": { "步骤": "x_T → x_{0.75} → x_{0.5} → x_{0.25} → x_0", "质量": "★★★★☆", "速度": "★★★★☆", }, "8步 TTST": { "步骤": "更细粒度的多步重建", "质量": "★★★★★", "速度": "★★★☆☆", }, "对比 DDPM": { "DDPM 100步": "质量 ★★★★★, 速度 ★☆☆☆☆", "CM 1步": "质量 ★★★☆☆, 速度 ★★★★★", "CM 4步 TTST": "质量 ★★★★☆, 速度 ★★★★☆", }, }6. CFG 条件一致性
6.1 为 Consistency Model 添加条件控制
原始 Consistency Models 是无条件模型。通过 Consistency Fine-tuning(CFG),可以支持条件控制:
def cfg_consistency_model(): """ CFG Consistency Model。
方法: 微调预训练 CM,添加条件信息。 损失: L = |f_θ(x_t, c) - x_0^teacher(x_t, c)|² """
class CFGConsistencyModel(nn.Module): def __init__(self, base_cm, text_encoder): super().__init__() self.base_cm = base_cm self.text_encoder = text_encoder
# 额外的条件 embedding 投影 self.cfg_head = nn.Sequential( nn.Linear(768 + 768, 256), # cond + uncond → time_embed dim nn.SiLU(), nn.Linear(256, 256), )
def forward(self, x, t, cond_embed=None): """ 参数: x: (B, C, H, W) 输入 t: (B,) 时间步 cond_embed: (B, D) 条件 embedding (如文本) """ # 文本编码 if cond_embed is None: cond_embed = torch.zeros(x.shape[0], 768, device=x.device)
# 无条件 (全零 embedding) uncond_embed = torch.zeros_like(cond_embed)
# 条件预测 x_cond = self.base_cm(x, t, cond_embed)
# 无条件预测 x_uncond = self.base_cm(x, t, uncond_embed)
# CFG 合并 cfg_scale = 7.5 x_0 = x_uncond + cfg_scale * (x_cond - x_uncond)
return x_0
return CFGConsistencyModel
def cfg_training(): """ CFG Consistency Model 训练。 """
def train_cfg_cm(student_cm, teacher_cm, p_data, p_text, optimizer, config): """ 参数: student_cm: 要训练的 CFG CM teacher_cm: 预训练的 CFG 扩散模型 p_data: 图像分布 p_text: 文本分布 """ # 采样 x_0 = sample_batch(p_data) text = sample_batch(p_text)
# 文本编码 cond_embed = student_cm.text_encoder(text)
# 加噪 t = torch.rand(x_0.shape[0]) * (config.T - config.epsilon) + config.epsilon noise = torch.randn_like(x_0) x_t = alpha(t).view(-1, 1, 1, 1) * x_0 + sigma(t).view(-1, 1, 1, 1) * noise
# 学生预测 pred = student_cm(x_t, t, cond_embed)
# 教师预测 (使用相同的条件) with torch.no_grad(): target = teacher_cm(x_t, t, cond_embed)
# CFG 损失 loss = F.mse_loss(pred, target)
return loss7. 完整 PyTorch 实现
7.1 Consistency Model 完整训练
import torchimport torch.nn as nnimport torch.nn.functional as Fimport mathfrom typing import Optional
def alpha_sigma(t, sigma_data=0.5, p_mean=-1.2, p_std=1.2): """ EDM 风格的 alpha 和 sigma 调度。 """ # t 是 [0, 1] 归一化的时间步 # sigma(t) 使用指数调度 sigma = sigma_data ** (1 + 0.5 * (t - 0.5)) # 非线性调度
# alpha 从 t 计算 # 使用 sigmoid 中心化 c = -2.5 * (t - 0.5) alpha = 1 / (1 + math.e ** c)
return alpha, sigma
class ConsistencyModel(nn.Module): """ Consistency Model 完整实现。
核心设计: - 输入任意时间步的 x_t,输出 x_0 - 使用 EDM UNet 作为 backbone - 时间依赖的 c_skip 和 c_out 跳跃连接 """
def __init__(self, unet: nn.Module, sigma_data=0.5): super().__init__() self.unet = unet self.sigma_data = sigma_data
# 时间嵌入维度 (与 UNet 匹配) time_embed_dim = 256
# c_skip(t) 和 c_out(t) 的 MLP self.c_skip_net = nn.Sequential( nn.Linear(time_embed_dim, 256), nn.SiLU(), nn.Linear(256, 1), ) self.c_out_net = nn.Sequential( nn.Linear(time_embed_dim, 256), nn.SiLU(), nn.Linear(256, 1), )
# 时间步嵌入 self.time_embed = nn.Sequential( nn.Linear(128, time_embed_dim), nn.SiLU(), nn.Linear(time_embed_dim, time_embed_dim), )
def get_time_embedding(self, t): """获取时间步的正弦位置嵌入。""" half_dim = 64 emb = math.log(10000) / (half_dim - 1) emb = torch.exp(torch.arange(half_dim, device=t.device) * -emb) emb = t[:, None] * emb[None, :] emb = torch.cat([emb.sin(), emb.cos()], dim=-1) return self.time_embed(emb)
def forward(self, x, t): """ 前向传播。
参数: x: (B, C, H, W) 输入 (任意时间步的 x_t) t: (B,) 时间步 (0 到 1) 返回: x_0: (B, C, H, W) 预测的原始数据 """ # 时间嵌入 t_emb = self.get_time_embedding(t) # (B, 256)
# EDM UNet 前向 h = self.unet.prep(x) h = self.unet.downsample(h, t_emb) h = self.unet.attention(h, t_emb) h = self.unet.upsample(h, t_emb)
# 网络输出 F_θ(x, t) f_theta = self.unet.final(h)
# 时间依赖跳跃连接 c_skip = self.c_skip_net(t_emb) # (B, 1) c_out = self.c_out_net(t_emb) # (B, 1)
# 边界条件: t=0 时,c_skip=1, c_out=0 # 边界条件: t=1 时,c_skip=0, c_out=1 # 通过 Sigmoid 实现自动插值 c_skip = torch.sigmoid(c_skip) c_out = torch.sigmoid(c_out)
# Consistency Model 输出 x_0 = c_skip.view(-1, 1, 1, 1) * x + c_out.view(-1, 1, 1, 1) * f_theta
return x_0
class ConsistencyModelTrainer: """Consistency Model 训练器。"""
def __init__(self, model: ConsistencyModel, teacher_model: Optional[nn.Module] = None, device="cuda", lambda_1=1.0, lambda_2=1.0): self.model = model.to(device) self.teacher = teacher_model.to(device) if teacher_model else None self.device = device self.lambda_1 = lambda_1 self.lambda_2 = lambda_2 self.step = 0
def train_step(self, x_0: torch.Tensor, cfg_emb: Optional[torch.Tensor] = None): """ 单步训练。
参数: x_0: (B, C, H, W) 真实数据 cfg_emb: (B, D) 条件 embedding (可选) """ B = x_0.shape[0]
# ============ 轨道 1: 自监督一致性 ============ epsilon = 1e-3 T = 1.0
# 采样两个时间步 (t_1 < t_2) t_1 = torch.rand(B, device=self.device) * (T - epsilon) + epsilon t_2 = torch.rand(B, device=self.device) * (T - t_1) + t_1
# 加噪 noise_1 = torch.randn_like(x_0) noise_2 = torch.randn_like(x_0)
alpha_1, sigma_1 = alpha_sigma(t_1) alpha_2, sigma_2 = alpha_sigma(t_2)
x_t1 = alpha_1.view(-1, 1, 1, 1) * x_0 + sigma_1.view(-1, 1, 1, 1) * noise_1 x_t2 = alpha_2.view(-1, 1, 1, 1) * x_0 + sigma_2.view(-1, 1, 1, 1) * noise_2
# 学生预测 pred_1 = self.model(x_t1, t_1)
# 自监督损失: f_θ(x_{t_1}) ≈ x_{t_2} loss_1 = F.mse_loss(pred_1, x_t2.detach())
# ============ 轨道 2: 蒸馏一致性 ============ if self.teacher is not None: # 用教师模型计算目标 t_eps = torch.full((B,), epsilon, device=self.device) alpha_eps, sigma_eps = alpha_sigma(t_eps) x_eps = alpha_eps * x_0 + sigma_eps * torch.randn_like(x_0)
with torch.no_grad(): # 教师预测 x_0 target = self.teacher(x_eps, t_eps)
# 学生预测 pred_eps = self.model(x_eps, t_eps)
# 蒸馏损失 loss_2 = F.mse_loss(pred_eps, target.detach())
total_loss = self.lambda_1 * loss_1 + self.lambda_2 * loss_2 else: total_loss = loss_1 loss_2 = torch.tensor(0.0)
# 反向传播 total_loss.backward()
return {"total": total_loss.item(), "self_sup": loss_1.item(), "distill": loss_2.item()}
@torch.no_grad() def sample(self, batch_size=1, steps=1, guide_scale=1.0): """ 推理采样。 """ # 初始化噪声 x = torch.randn(batch_size, 3, 32, 32, device=self.device)
if steps == 1: # 一步生成 t = torch.ones(batch_size, device=self.device) x_0 = self.model(x, t) else: # TTST 多步生成 timesteps = torch.linspace(1.0, 0.0, steps + 1, device=self.device) for i in range(steps): t = timesteps[i] t_next = timesteps[i + 1]
pred = self.model(x, t.expand(batch_size))
if t_next == 0: return pred
# 重建下一步 alpha_t, sigma_t = alpha_sigma(t) alpha_next, sigma_next = alpha_sigma(t_next)
noise = (x - alpha_t * pred) / sigma_t.clamp(min=1e-6) x = alpha_next * pred + sigma_next * noise
x_0 = x
return x_08. LCM:Latent Consistency Models
8.1 LCM 的核心思想
LCM(Latent Consistency Models, Luo et al., 2023)是 Consistency Models 在 Stable Diffusion 上的应用:
LCM vs 标准 CM:
Consistency Model: - 在像素空间操作 - 需要处理高分辨率 (512×512×3) - 计算量大
LCM (Latent Consistency Model): - 在潜空间操作 (使用 SD 的 VAE) - 处理低维潜变量 (64×64×4) - 计算量大幅降低 - 可以直接加载 SD 权重进行蒸馏
核心公式: x_0 = c_skip(t) · x + c_out(t) · ε_θ(x, t)
LCM 的特殊处理: - 使用 SD 的 UNet 作为 backbone - 直接蒸馏 ε_θ 的预测 - 输出 x_0 而非噪声 ε8.2 LCM 的训练与推理
def lcm_training(): """ LCM 训练 (与 Consistency Model 相同,但操作在潜空间)。 """
def train_lcm(sd_unet, vae, text_encoder, dataloader, optimizer): """ LCM 训练。 """ for batch in dataloader: images, texts = batch
# VAE 编码到潜空间 with torch.no_grad(): z_0 = vae.encode(images)
# 采样时间步 t = torch.rand(B) * 0.998 + 0.002 # 避免 t=0, t=1
# 加噪 noise = torch.randn_like(z_0) z_t = alpha(t) * z_0 + sigma(t) * noise
# 文本编码 context = text_encoder(texts)
# LCM 预测 x_0 # 使用 SD UNet,但输出头替换为 CM 的 c_skip/c_out epsilon_pred = sd_unet(z_t, t, context) x_0_pred = (z_t - sigma(t) * epsilon_pred) / alpha(t)
# 教师目标 with torch.no_grad(): target = x_0_pred # 可以用更好的教师
loss = F.mse_loss(x_0_pred, target)
optimizer.zero_grad() loss.backward() optimizer.step()
return train_lcm
@torch.no_grad()def lcm_sample(lcm_model, text_encoder, prompt, num_steps=4, guidance_scale=7.0): """ LCM 推理 — 4 步生成 SD 图像。 """ # 文本编码 context = text_encoder([prompt]) uncond = text_encoder([""])
# 初始化噪声 z = torch.randn(1, 4, 64, 64, device=device)
# 调度 timesteps = [0.9, 0.7, 0.4, 0.0]
for i, t in enumerate(tqdm(timesteps)): t_tensor = torch.tensor([t], device=device)
# CFG z_cond = lcm_model(z, t_tensor, context) z_uncond = lcm_model(z, t_tensor, uncond) z_0 = z_uncond + guidance_scale * (z_cond - z_uncond)
if t == 0.0: break
# 重建下一步 alpha_t = alpha_of_t(t) alpha_next = alpha_of_t(timesteps[i + 1]) sigma_t = sigma_of_t(t) sigma_next = sigma_of_t(timesteps[i + 1])
noise = (z - alpha_t * z_0) / sigma_t.clamp(min=1e-6) z = alpha_next * z_0 + sigma_next * noise
# VAE 解码 image = vae.decode(z)
return image9. 与其他方法的对比
9.1 加速方法全面对比
def comparison_all(): """ 扩散模型加速方法的全面对比。 """ return { "DDIM": { "步数": "20-50 步", "质量": "★★★★☆ (20步) - ★★★★★ (50步)", "速度": "★★★☆☆", "方法": "确定性 ODE 采样", "优点": "无需训练,通用", "缺点": "仍有固定步数延迟", }, "LCM": { "步数": "2-8 步", "质量": "★★★★☆", "速度": "★★★★★", "方法": "潜空间一致性模型", "优点": "步数极少,兼容 SD", "缺点": "需要蒸馏训练", }, "Consistency Models": { "步数": "1 步 (可达 ★★★★☆)", "质量": "★★★★☆ (TTST) - ★★★☆☆ (单步)", "速度": "★★★★★", "方法": "ODE 轨迹一致性蒸馏", "优点": "一步生成,理论优美", "缺点": "单步质量有限,需 TTST 增强", }, "LCM + TTST": { "步数": "4-8 步", "质量": "★★★★★", "速度": "★★★★☆", "方法": "LCM + 多步自改进", "优点": "质量和速度的最佳平衡", "缺点": "比单步慢", }, "对抗蒸馏": { "步数": "1 步", "质量": "★★★★☆", "速度": "★★★★★", "方法": "GAN + 扩散模型蒸馏", "优点": "一步高质量", "缺点": "训练不稳定,mode collapse", }, }9.2 Consistency Models vs DDPM
def cm_vs_ddpm(): """ Consistency Models vs DDPM 对比。 """ return { "采样": { "DDPM": "100-1000 步随机采样", "CM": "1-N 步确定性/随机采样", }, "损失": { "DDPM": "噪声预测 |ε - ε_θ|²", "CM": "一致性 |f_θ(x_t) - f_θ(x_{t+Δ})|²", }, "训练": { "DDPM": "单阶段,从零训练", "CM": "双轨 (自监督 + 蒸馏),或从零训练", }, "条件": { "DDPM": "天然支持 (cross-attention)", "CM": "需要 CFG 微调", }, "理论基础": { "DDPM": "变分推断 + ELBO", "CM": "ODE 轨迹一致性 + 最优传输", }, "生成模式": { "DDPM": "自回归多步", "CM": "直接映射 (类似 GAN)", }, }10. 总结
10.1 核心要点
| 维度 | 关键要点 |
|---|---|
| 核心思想 | ODE 轨迹一致性:轨迹上任意点映射到轨迹起点 x_0 |
| 数学框架 | f_θ(x_t) = x_0 对所有 t ∈ [0, T] |
| 训练目标 | TTCT = 自监督一致性 + 蒸馏一致性 |
| 边界条件 | t=0 时 f_θ(x, 0) = x,t=T 时预测数据均值 |
| 推理方式 | 单步生成 或 TTST 多步增强 |
| LCM | Consistency Models + 潜空间 + SD 权重 |
| 条件控制 | CFG Consistency Fine-tuning |
| 与 DDPM 关系 | CM 压缩了 DDPM 的多步采样为单步映射 |
10.2 采样速度对比
生成一张 512×512 图像所需的前向次数:
DDPM (100步) ████████████████████████████████████ 100步 ~50秒DDPM (1000步) ██████████████████████████████████████████████ 1000步 ~500秒DDIM (50步) ██████████████████ 50步 ~25秒DDIM (20步) ████████ 20步 ~10秒LCM (4步) █ 4步 ~2秒CM (1步) █ 1步 ~0.5秒CM (TTST 4步) ████ 4步 ~2秒GAN (1步) █ 1步 ~0.5秒
质量对比:GAN (1步) ████████████████████████████████ ~85% 质量CM (1步) ██████████████████████████ ~80% 质量LCM (4步) █████████████████████████████████████ ~95% 质量CM (TTST 8步) ██████████████████████████████████████ ~97% 质量DDIM (50步) ████████████████████████████████████████ ~99% 质量DDPM (100步) ███████████████████████████████████████████████ 100% 质量10.3 一句话总结
Consistency Models 的核心创新是将扩散模型的”多步去噪”视为一条 ODE 轨迹,并训练一个”一致性函数”f_θ,使得轨迹上的任意点都能一步映射到数据分布的起点 x_0——这通过 TTCT 双轨训练(自监督一致性 + 蒸馏一致性)和 EDM 兼容的 UNet 架构实现。CM 可以从零训练也可以从预训练扩散模型蒸馏,支持 1 步到多步推理;在潜空间的变体 LCM 则与 Stable Diffusion 无缝结合,实现了 4 步生成高质量图像,代表了扩散模型从”慢速多步”向”快速一步”演进的重要技术路线。
10.4 推荐资源
论文: - Consistency Models (Song et al., 2023): "Consistency Models" ICML 2023, OpenAI - LCM (Luo et al., 2023): "Latent Consistency Models: Synthesizing Convergence-Offering Latent Space for Image Generation" - EDM (Karras et al., 2022): "Elucidating the Design Space of Diffusion Models"
代码: - openai/consistency_models: https://github.com/openai/consistency_models - openai/scoreODE: https://github.com/openai/scoreODE - Latent-Consistency-Models/sd-webui-lcm: LCM for WebUI
相关: - DDIM (Song et al., 2021): DDIM 采样加速 - PF-ODE (Lu et al., 2022): Rectified Flow,另一种 ODE 加速 - flow-matching: (Lipman et al., 2022) 统一扩散和流匹配文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!

