扩散模型(DDPM)的训练与推理
·
扩散模型(Diffusion Model)作为近年来生成式 AI 领域的核心技术,已广泛应用于图像生成、语音合成、文本生成等场景。不同于 GAN 的对抗训练模式,扩散模型通过模拟 “逐步加噪 - 逐步去噪” 的物理过程实现数据生成,具有训练稳定、生成质量高的显著优势。本文将从原理入手,拆解扩散模型的训练与推理全过程,并提供可落地的伪代码和 Python 实现示例。
一、扩散模型训练过程
1.1 训练核心逻辑
训练的本质是让模型(噪声预测器)尽可能准确地预测前向扩散过程中添加的噪声。具体步骤如下:
1. 采样一批真实数据;
2. 随机采样步长;
3. 采样高斯噪声;
4. 根据前向扩散公式计算;
5. 将和
输入模型
,得到预测噪声
;
6. 计算预测噪声与真实噪声的 MSE 损失,反向传播更新模型参数。
1.2 训练代码示例
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
import numpy as np
T = 1000
beta = torch.linspace(1e-4, 0.02, T)
alpha = 1 - beta
alpha_bar = torch.cumprod(alpha, dim=0) # 累计乘积
def train_diffusion(model, dataloader, epochs, device):
model.to(device)
criterion = nn.MSELoss()
optimizer = optim.Adam(model.parameters(), lr=1e-4)
alpha_bar = alpha_bar.to(device)
for epoch in range(epochs):
model.train()
total_loss = 0.0
for batch in dataloader:
x0 = batch.to(device) # 假设输入为3通道图像,shape=(B,3,H,W)
B = x0.shape[0]
# 随机采样步长t
t = torch.randint(1, T+1, (B,), device=device)
# 采样高斯噪声
eps = torch.randn_like(x0)
# 计算xt
alpha_bar_t = alpha_bar[t-1].reshape(B, 1, 1, 1) # t从1开始,索引从0开始
xt = torch.sqrt(alpha_bar_t) * x0 + torch.sqrt(1 - alpha_bar_t) * eps
# 预测噪声
eps_theta = model(xt, t-1) # 嵌入层索引从0开始
# 计算损失并更新
loss = criterion(eps_theta, eps)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}, Loss: {total_loss/len(dataloader):.4f}")
二、扩散模型推理过程
2.1 推理核心逻辑
推理(生成)是反向扩散的过程:从纯噪声出发,利用训练好的模型逐步去噪,最终得到接近真实数据的
。具体步骤如下:
1. 初始化为随机高斯噪声;
2. 反向迭代:
将和
输入模型,得到预测噪声;
根据反向扩散公式计算;
3. 迭代结束后,即为生成的结果。
2.2 推理代码示例
def sample_diffusion(model, shape, device):
"""
扩散模型推理函数
:param model: 训练好的噪声预测模型
:param shape: 生成数据的形状,如(B, 3, 32, 32)
:param device: 运行设备
:return: 生成的数据x0
"""
model.eval()
alpha = alpha.to(device)
alpha_bar = alpha_bar.to(device)
beta = beta.to(device)
# 初始化xT为纯噪声
xt = torch.randn(shape, device=device)
with torch.no_grad(): # 推理阶段禁用梯度
for t in range(T, 0, -1):
# 1. 准备时间步(适配嵌入层)
t_tensor = torch.tensor([t-1], device=device).repeat(shape[0])
# 2. 预测噪声
eps_theta = model(xt, t_tensor)
# 3. 计算反向扩散的均值和方差
alpha_t = alpha[t-1]
alpha_bar_t = alpha_bar[t-1]
beta_t = beta[t-1]
# 均值μ_t
mean = (1 / torch.sqrt(alpha_t)) * (
xt - (1 - alpha_t) / torch.sqrt(1 - alpha_bar_t) * eps_theta
)
# 方差σ_t(简化为sqrt(β_t))
std = torch.sqrt(beta_t)
# 4. 采样噪声(t=1时为0)
if t > 1:
z = torch.randn_like(xt)
else:
z = torch.zeros_like(xt)
# 5. 更新xt为x_{t-1}
xt = mean + std * z
# 归一化到[0,1](可选,根据数据分布调整)
x0 = torch.clamp(xt, -1, 1)
x0 = (x0 + 1) / 2
return x0
更多推荐
所有评论(0)