Skip to content

DDPM 公式推导

DDPM(Denoising Diffusion Probabilistic Model)把生成过程拆成两个马尔可夫链:前向过程逐步加噪,反向过程逐步去噪。

前向扩散过程

给定真实数据 x0,前向过程定义为:

q(xt|xt1)=N(xt;1βtxt1,βtI)

记:

αt=1βt,α¯t=s=1tαs

递推可得直接从 x0 采样 xt 的公式:

q(xt|x0)=N(xt;α¯tx0,(1α¯t)I)

因此:

xt=α¯tx0+1α¯tϵ,ϵN(0,I)

反向去噪过程

模型学习反向分布:

pθ(xt1|xt)=N(xt1;μθ(xt,t),Σθ(xt,t))

DDPM 常让神经网络预测噪声 ϵθ(xt,t),再由噪声还原均值:

μθ(xt,t)=1αt(xtβt1α¯tϵθ(xt,t))

训练损失可简化为噪声预测误差:

Lsimple=Et,x0,ϵ[||ϵϵθ(xt,t)||2]

推导直觉

由前向采样公式:

xt=α¯tx0+1α¯tϵ

可以解出:

x0=xt1α¯tϵα¯t

真实噪声 ϵ 不可直接用于推理,所以训练网络预测 ϵθ。只要噪声预测准确,就能逐步把 xt 还原到更干净的 xt1

应用代码

python
import torch

T = 1000
betas = torch.linspace(1e-4, 0.02, T)
alphas = 1.0 - betas
alpha_bar = torch.cumprod(alphas, dim=0)

def q_sample(x0, t, noise=None):
    if noise is None:
        noise = torch.randn_like(x0)
    a_bar = alpha_bar[t].view(-1, 1, 1, 1)
    return torch.sqrt(a_bar) * x0 + torch.sqrt(1 - a_bar) * noise

def ddpm_loss(model, x0):
    b = x0.size(0)
    t = torch.randint(0, T, (b,), device=x0.device)
    noise = torch.randn_like(x0)
    xt = q_sample(x0, t, noise)
    pred_noise = model(xt, t)
    return ((noise - pred_noise) ** 2).mean()

小结

DDPM 的关键在于:前向过程有解析加噪公式,反向过程用神经网络预测噪声。训练时模型只需学习“给定带噪图像和时间步,噪声是多少”。