扩散模型(Diffusion Model)作为近年来生成式 AI 领域的核心技术,已广泛应用于图像生成、语音合成、文本生成等场景。不同于 GAN 的对抗训练模式,扩散模型通过模拟 “逐步加噪 - 逐步去噪” 的物理过程实现数据生成,具有训练稳定、生成质量高的显著优势。本文将从原理入手,拆解扩散模型的训练与推理全过程,并提供可落地的伪代码和 Python 实现示例。

一、扩散模型训练过程

1.1 训练核心逻辑

训练的本质是让模型(噪声预测器\epsilon_{\theta})尽可能准确地预测前向扩散过程中添加的噪声。具体步骤如下:
1. 采样一批真实数据x_{0}
2. 随机采样步长t\in [1,T]
3. 采样高斯噪声\epsilon \sim \mathbb{N}(0,I)
4. 根据前向扩散公式计算x_{t}
5. 将x_{t}t输入模型\epsilon _{\theta },得到预测噪声\epsilon _{\theta } (x_{t}, t)
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 推理核心逻辑

推理(生成)是反向扩散的过程:从纯噪声x_{T}出发,利用训练好的模型逐步去噪,最终得到接近真实数据的x_{0}。具体步骤如下:
1. 初始化x_{T}为随机高斯噪声;
2. 反向迭代:
    将x_{t}t输入模型,得到预测噪声;
    根据反向扩散公式计算x_{t-1}
3. 迭代结束后,x_{0}即为生成的结果。

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

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐