Consistency Models 深度解析:一步生成扩散模型

5767 字
29 分钟
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_models

2. 核心思想: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_2
def 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_loss

3.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_0

4.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 x

5.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 loss

7. 完整 PyTorch 实现#

7.1 Consistency Model 完整训练#

import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from 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_0

8. 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 image

9. 与其他方法的对比#

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 多步增强
LCMConsistency 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) 统一扩散和流匹配

文章分享

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

Consistency Models 深度解析:一步生成扩散模型
https://aiattnstudio.link/posts/consistency-models/
作者
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标签