扩散模型(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