尧图精选

扩散模型前向过程详解:从数学原理到代码实现与调参

🕒 发布时间:2026/9/28 8:48:04 📁 来源:尧图网络
扩散模型这两年火得一塌糊涂但很多人上手就是跑 Stable Diffusion WebUI、调 LoRA、换采样器问到底层怎么工作的答不上来。我见过太多人卡在会用但不懂的阶段一旦遇到出图效果不对、训练不收敛、采样步数不知道怎么设就完全抓瞎。前向扩散过程Forward Diffusion Process是整个扩散模型的地基你把这个搞明白了后面反向去噪、噪声调度、采样策略这些东西全都能串起来。这篇就专门拆前向扩散从数学原理到代码实现从噪声调度的设计动机到实际调参的经验一次讲透。不管你是刚接触扩散模型的新手还是已经能跑通 DDPM 但没细看公式的老手这篇都能帮你把这块知识补扎实。1. 前向扩散到底在做什么从一张图变成纯噪声1.1 一句话说清楚前向过程的本质前向扩散过程说白了就一件事把一张正常的图片一步一步地加噪声直到它变成完全随机的噪声图。这个过程是固定的不需要学习没有可训练参数就是一个预先定义好的数学操作。你可以把它想象成往一杯清水里滴墨水。每一步滴一点水就浑一点滴到足够多步之后整杯水变成均匀的灰色你再也看不出原来清水长什么样。前向扩散就是这个滴墨水的过程只不过滴的是高斯噪声而且每一步滴多少是有严格数学定义的。DDPMDenoising Diffusion Probabilistic Models里定义的前向过程是一个马尔可夫链总共 T 步通常 T1000。每一步只依赖上一步的结果公式长这样q(x_t | x_{t-1}) N(x_t; sqrt(1-β_t) * x_{t-1}, β_t * I)翻译成人话就是第 t 步的图像 x_t是从第 t-1 步的图像 x_{t-1} 出发先乘一个缩放系数 sqrt(1-β_t)再加上一个方差为 β_t 的高斯噪声得到的。β_t 是每一步的噪声强度是一个预先设定好的超参数。1.2 为什么是乘缩放系数再加噪声而不是直接加噪声这个问题很多人第一眼看到公式时会疑惑。为什么不直接 x_t x_{t-1} noise非要搞个 sqrt(1-β_t) 的缩放原因在于方差保持。如果每一步只是加噪声不缩放那么随着步数增加图像的像素值方差会越来越大数值会爆炸。加上 sqrt(1-β_t) 这个系数之后可以保证在每一步中x_t 的方差维持在一个合理范围内。当 β_t 很小的时候sqrt(1-β_t) 约等于 1缩放效果微弱当 β_t 接近 1 的时候缩放系数接近 0上一步的信息几乎被完全丢弃只剩下噪声。这个设计保证了整个前向过程在数值上是稳定的不会出现梯度爆炸或者数值溢出。你在实现的时候如果忘了这个缩放项训练几乎必然失败这是新手最容易犯的错误之一。1.3 从单步到任意步重参数化的威力单步公式看着简单但实际用的时候有个大问题如果我要得到 x_500难道真的要一步步迭代 500 次吗那训练的时候每次前向传播都要算 500 次效率太低了。DDPM 的核心数学技巧之一就是重参数化reparameterization它允许你从 x_0 直接一步计算出任意时刻 t 的 x_t不需要迭代。推导过程用到了高斯分布的性质这里直接给结论q(x_t | x_0) N(x_t; sqrt(ᾱ_t) * x_0, (1-ᾱ_t) * I)其中 α_t 1 - β_tᾱ_t α_1 * α_2 * ... * α_t也就是所有 α 的累乘。写成可直接采样的形式x_t sqrt(ᾱ_t) * x_0 sqrt(1-ᾱ_t) * ε其中 ε ~ N(0, I) 是标准高斯噪声。这个公式是整个扩散模型里最重要的公式之一。它的意义在于训练时你可以随机采一个 t然后一步到位算出 x_t完全不需要模拟前面的 t-1 步。这让训练效率提升了几百倍。没有这个重参数化技巧DDPM 根本不可能实用。我当初第一次看到这个推导的时候觉得最妙的地方在于它把逐步加噪这个看似必须串行的过程变成了一个可以并行计算的闭式解。这种数学上的化简是扩散模型能落地的关键。2. 噪声调度β_t 怎么设才合理2.1 线性调度与余弦调度的取舍β_t 的选择直接决定了前向过程的质量进而影响整个模型的生成效果。DDPM 原论文用的是线性调度β_1 1e-4β_T 0.02从 1e-4 到 0.02 均匀递增。但后来 OpenAI 在 Improved DDPM 里指出线性调度有个问题在接近 tT 的时候ᾱ_t 衰减得太快导致图像过早地变成纯噪声后面的很多步其实是在浪费计算。他们提出了余弦调度ᾱ_t cos²((t/T s) / (1 s) * π/2)其中 s 是一个小的偏移量通常取 0.008用来防止 t0 附近 β_t 太小。两种调度的对比如下调度方式优点缺点适用场景线性调度实现简单原论文验证后期信息衰减过快快速实验低分辨率余弦调度信息衰减更均匀实现稍复杂高分辨率高质量生成实测下来如果你做的是 256x256 以上的图像生成余弦调度确实能带来更稳定的训练和更好的最终效果。但如果只是跑个 demo 验证想法线性调度完全够用不用过度纠结。2.2 β_t 的取值范围为什么不能乱设β_t 必须满足 0 β_t 1这是硬性约束。如果 β_t 太大比如设成 0.5那么几步之后图像就完全变成噪声了前向过程太暴力反向去噪的时候模型很难学到有意义的中间状态。如果 β_t 太小比如 1e-6那么 1000 步之后图像可能还保留大量原始信息前向过程没有完成从信号到噪声的转换反向过程也就无从学起。实践中β_1 通常在 1e-4 到 1e-5 之间β_T 在 0.01 到 0.02 之间。这个范围是经过大量实验验证的你如果自己从头设计调度建议先在这个范围内调不要一上来就搞极端值。还有一个容易忽略的点β_t 的累乘 ᾱ_t 在 tT 时必须接近 0。这意味着 sqrt(ᾱ_T) 接近 0x_T 几乎完全由噪声项 sqrt(1-ᾱ_T) * ε 决定也就是纯噪声。如果你发现训练完之后采样出来的图全是噪声第一步就要检查 ᾱ_T 是不是没有衰减到足够小。2.3 代码实现噪声调度的具体写法用 PyTorch 实现线性调度和余弦调度都很简单下面给出可直接用的代码import torch import numpy as np def linear_beta_schedule(timesteps, beta_start1e-4, beta_end0.02): 线性噪声调度DDPM 原论文方案 return torch.linspace(beta_start, beta_end, timesteps) def cosine_beta_schedule(timesteps, s0.008): 余弦噪声调度Improved DDPM 方案 steps timesteps 1 x torch.linspace(0, timesteps, steps) alphas_cumprod torch.cos(((x / timesteps) s) / (1 s) * torch.pi * 0.5) ** 2 alphas_cumprod alphas_cumprod / alphas_cumprod[0] betas 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) return torch.clip(betas, 0.0001, 0.9999) # 使用示例 T 1000 betas linear_beta_schedule(T) alphas 1.0 - betas alphas_cumprod torch.cumprod(alphas, dim0) sqrt_alphas_cumprod torch.sqrt(alphas_cumprod) sqrt_one_minus_alphas_cumprod torch.sqrt(1.0 - alphas_cumprod)这段代码里alphas_cumprod就是 ᾱ_tsqrt_alphas_cumprod和sqrt_one_minus_alphas_cumprod就是重参数化公式里的两个系数。把这两个系数预先算好存起来训练时直接查表比每次重新计算快得多。注意torch.cumprod是累乘操作不是累加。我见过有人写成torch.cumsum结果 ᾱ_t 直接超过 1训练完全跑不通。这个坑很隐蔽因为代码不报错只是结果不对。3. 重参数化公式的代码落地与训练采样3.1 从 x_0 一步采样 x_t 的完整实现有了预计算的系数从 x_0 采样 x_t 就是一行代码的事def extract(a, t, x_shape): 从预计算的系数数组中按时间步 t 提取对应值并 reshape 成可广播的形状 batch_size t.shape[0] out a.gather(-1, t.cpu()) return out.reshape(batch_size, *((1,) * (len(x_shape) - 1))).to(t.device) def q_sample(x_start, t, noiseNone): 前向扩散从 x_0 直接采样 x_t if noise is None: noise torch.randn_like(x_start) sqrt_alphas_cumprod_t extract(sqrt_alphas_cumprod, t, x_start.shape) sqrt_one_minus_alphas_cumprod_t extract(sqrt_one_minus_alphas_cumprod, t, x_start.shape) return sqrt_alphas_cumprod_t * x_start sqrt_one_minus_alphas_cumprod_t * noiseextract这个函数的作用是训练时一个 batch 里每个样本的 t 可能不同需要根据每个样本的 t 取出对应的系数然后 reshape 成 (batch_size, 1, 1, 1) 的形状这样才能和图像张量 (batch_size, C, H, W) 做广播乘法。这个细节如果处理不好会出现维度不匹配的报错或者更糟糕的——广播成了错误的形状但程序不报错结果训练出来的模型完全不能用。3.2 训练循环里前向过程扮演什么角色DDPM 的训练逻辑非常简洁核心就是随机采一个 t用前向公式算出 x_t然后让模型预测加进去的噪声 ε。def train_step(model, x_0, optimizer): batch_size x_0.shape[0] t torch.randint(0, T, (batch_size,), devicex_0.device).long() noise torch.randn_like(x_0) x_t q_sample(x_0, t, noise) predicted_noise model(x_t, t) loss torch.nn.functional.mse_loss(predicted_noise, noise) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()这里有个关键点模型预测的是噪声 ε不是原始图像 x_0。这是 DDPM 的设计选择。为什么预测噪声而不是预测图像因为预测噪声的优化目标更稳定噪声是标准高斯分布尺度统一模型更容易学习。如果直接预测 x_0由于图像像素值范围差异大训练会不稳定。当然后来也有工作尝试预测 x_0 或者预测速度场 v sqrt(ᾱ_t) * ε - sqrt(1-ᾱ_t) * x_0各有优劣。但对于入门来说先理解预测噪声这个范式就够了。3.3 前向过程在采样阶段的反向使用虽然前向过程本身不参与采样但采样阶段会反复用到前向过程定义的系数。反向去噪的公式是x_{t-1} (1/sqrt(α_t)) * (x_t - (β_t/sqrt(1-ᾱ_t)) * ε_θ(x_t, t)) σ_t * z其中 σ_t 是反向过程的方差z 是采样的噪声t0 时不加。这个公式里的 α_t、β_t、ᾱ_t 全部来自前向过程的定义。所以前向过程虽然简单但它是整个扩散模型参数体系的来源反向过程只是在这个体系上做逆运算。我个人的经验是把前向过程的系数表打印出来看一眼对理解整个流程帮助极大。你会直观地看到 sqrt(ᾱ_t) 从 1 逐渐降到接近 0sqrt(1-ᾱ_t) 从 0 逐渐升到接近 1这个此消彼长的过程就是信号逐渐被噪声淹没的数学表达。4. 实操中容易踩的坑与排查思路4.1 系数计算精度问题导致采样出雪花图这是我在实际项目里遇到过的最隐蔽的 bug 之一。用 float32 计算 ᾱ_t 的累乘时当 T1000 且 β_t 较小时累乘结果在数值上可能不够精确导致 sqrt(ᾱ_T) 不是精确的 0 而是一个很小的正数或者 sqrt(1-ᾱ_T) 不是精确的 1。这个误差在训练时可能看不出来但采样时会被放大最终生成的图像出现大量噪点看起来像雪花屏。解决方案是用 float64 计算系数表然后再转成 float32 给模型用betas linear_beta_schedule(T).double() alphas 1.0 - betas alphas_cumprod torch.cumprod(alphas, dim0).float()这个改动很小但能避免很多莫名其妙的采样问题。如果你发现训练 loss 正常下降但采样结果全是噪声第一件事就是检查系数表的数值精度。4.2 时间步 t 的采样策略影响训练效果训练时 t 是从 0 到 T-1 均匀随机采的这是 DDPM 的标准做法。但实践中我发现如果数据集比较小或者图像比较简单均匀采样会导致模型在 t 较大的区域噪声多学得好在 t 较小的区域噪声少学得差因为小 t 区域的噪声幅度小损失值天然就小梯度贡献也小。一个实用的改进是对 t 做重要性采样让模型在损失大的区域多训练。不过这个改动会增加实现复杂度建议先把基础版本跑通确认没问题之后再尝试。对于大多数场景均匀采样已经够用了。还有一个细节t 的类型必须是 longint64不能是 float。我见过有人用torch.rand生成 t 然后直接传给extract函数结果gather操作报错。这个错误信息很明确但新手可能不知道问题出在哪。4.3 图像归一化范围与前向过程的配合前向过程假设输入图像 x_0 的像素值在一个合理的范围内通常是 [-1, 1] 或者 [0, 1]。DDPM 原论文用的是 [-1, 1]。如果你用 [0, 255] 的原始像素值直接做前向扩散噪声的尺度相对于信号来说太小模型需要学很久才能适应。标准做法是把图像归一化到 [-1, 1]# 假设 image 是 [0, 1] 范围的 tensor x_0 image * 2.0 - 1.0采样完成之后再反向变换回 [0, 1]image (x_0 1.0) / 2.0 image torch.clamp(image, 0.0, 1.0)这个归一化步骤看起来不起眼但如果不做训练会明显变慢生成质量也会下降。我在早期实验里偷懒没做归一化结果模型训练了 200 个 epoch 还是出不了清晰的图加上归一化之后 50 个 epoch 就有模有样了。4.4 排查清单前向过程相关问题的快速定位遇到问题时按下面这个清单逐项检查能快速定位大部分前向过程相关的 bug检查项正常表现异常表现与原因ᾱ_T 的值接近 0 1e-4接近 1 说明 β 太小或 T 太小sqrt(ᾱ_t) 曲线从 1 单调降到 0非单调说明 β 计算有误x_T 的分布接近标准高斯有明显结构说明前向不充分系数表精度float64 计算后转 float32全程 float32 可能有累积误差t 的数据类型torch.longfloat 类型会导致 gather 报错图像归一化范围[-1, 1][0, 255] 会导致训练缓慢这张表建议存下来每次新项目开始前过一遍能省下大量调试时间。5. 从 DDPM 到 Stable Diffusion前向过程的继承与变化5.1 Stable Diffusion 里的前向过程有什么不同Stable Diffusion 虽然架构比 DDPM 复杂得多引入了 VAE 和 Cross-Attention但前向扩散过程的数学定义几乎完全一样。区别在于DDPM 是在像素空间做扩散Stable Diffusion 是在 VAE 的潜空间latent space做扩散。这个区别带来的影响是潜空间的维度远小于像素空间比如 512x512x3 的图像VAE 编码后可能变成 64x64x4所以前向过程的计算量大幅降低训练和采样都快得多。但前向过程的公式、噪声调度的设计思路、重参数化技巧全部原封不动地继承了下来。你如果理解了 DDPM 的前向过程去看 Stable Diffusion 的代码会发现q_sample函数几乎一模一样只是输入的 x_0 从图像变成了 latent。这就是为什么我说前向过程是地基——地基打好了上面的楼怎么盖都能看懂。5.2 噪声调度在 Stable Diffusion 中的实际选择Stable Diffusion 系列模型包括 SD 1.5、SD 2.1、SDXL用的都是scaled linear schedule也就是线性调度的变体。具体来说它是对 β_t 做线性插值但插值的起点和终点经过了重新缩放使得 ᾱ_t 在 tT 时更接近 0。这个选择的原因和前面说的余弦调度类似让信息衰减更均匀避免后期浪费步数。但 Stable Diffusion 没有直接用余弦调度可能是因为线性调度的实现更简单而且在潜空间里线性调度的表现已经足够好。实际使用 Stable Diffusion WebUI 或者 Forge 的时候你会在采样器设置里看到 Schedule type 选项可以选 Linear、Karras、Exponential 等。这些调度方式影响的是采样阶段的步长分配但它们的底层逻辑都源于前向过程定义的 ᾱ_t 曲线。理解了前向过程的噪声调度你就能明白为什么 Karras 调度在少步数采样时效果更好——它是在 ᾱ_t 变化快的区域多采样变化慢的区域少采样。5.3 前向过程知识对实际调参的指导意义很多人调 Stable Diffusion 的参数是靠试试多了就有感觉了。但如果你懂前向过程很多参数是可以推理出来的。比如采样步数为什么 20 步和 50 步的效果差异不大但 5 步和 10 步差异很大因为前向过程的 ᾱ_t 曲线在 t 较小的时候变化快在 t 较大的时候变化慢。采样步数少的时候如果步长均匀分配小 t 区域的采样点太少重建质量就差。Karras 调度通过非均匀步长解决了这个问题所以少步数下表现更好。再比如 CFG scale为什么 CFG 太高会导致图像过饱和、颜色失真因为 CFG 放大了条件预测和无条件预测的差异而这个差异在噪声预测空间里被放大后反向过程计算出的 x_{t-1} 可能超出前向过程定义的合理范围导致数值不稳定。理解前向过程的数值范围约束就能明白 CFG 不能无限调高。这些推理不一定百分之百准确但它们给你提供了一个思考框架比盲目试参数高效得多。6. 自己动手验证前向过程的几个实验6.1 可视化不同 t 时刻的加噪图像最直观的验证方式就是把不同 t 时刻的 x_t 画出来。用 MNIST 或者 CIFAR-10 这种小数据集跑一下前向过程把 t0, 100, 200, ..., 900, 1000 的图像排成一行你会看到图像从清晰逐渐变成噪声的过程。import matplotlib.pyplot as plt def visualize_forward_process(x_0, timesteps_to_show[0, 100, 200, 400, 600, 800, 999]): fig, axes plt.subplots(1, len(timesteps_to_show), figsize(15, 3)) for idx, t_val in enumerate(timesteps_to_show): t torch.tensor([t_val], devicex_0.device) x_t q_sample(x_0, t) img (x_t[0].cpu().permute(1, 2, 0).numpy() 1) / 2 img np.clip(img, 0, 1) axes[idx].imshow(img) axes[idx].set_title(ft{t_val}) axes[idx].axis(off) plt.tight_layout() plt.show()这个实验看起来简单但能帮你建立对前向过程的直觉。我第一次做的时候发现 t400 左右图像就已经很难辨认了这让我意识到前向过程的信息衰减比想象中快也理解了为什么反向过程在 t 较大的区域主要是在猜而不是在还原。6.2 检查 ᾱ_t 曲线的形状把 ᾱ_t 随 t 变化的曲线画出来对比线性调度和余弦调度的差异import matplotlib.pyplot as plt T 1000 betas_linear linear_beta_schedule(T) betas_cosine cosine_beta_schedule(T) alphas_cumprod_linear torch.cumprod(1 - betas_linear, dim0) alphas_cumprod_cosine torch.cumprod(1 - betas_cosine, dim0) plt.figure(figsize(10, 5)) plt.plot(alphas_cumprod_linear.numpy(), labelLinear Schedule) plt.plot(alphas_cumprod_cosine.numpy(), labelCosine Schedule) plt.xlabel(Timestep t) plt.ylabel(alpha_bar_t) plt.legend() plt.grid(True) plt.show()你会看到线性调度的 ᾱ_t 在前期下降较慢后期急剧下降余弦调度则更平滑整个过程中下降速度更均匀。这个视觉差异直接对应到训练和采样效果的差异看一次就记住了。6.3 验证重参数化公式的正确性最后一个实验是验证重参数化公式分别用逐步迭代和一步到位两种方式计算 x_t看结果是否一致在数值误差范围内。def q_sample_iterative(x_0, t): 逐步迭代版本仅用于验证 x x_0 for i in range(t): beta betas[i] x torch.sqrt(1 - beta) * x torch.sqrt(beta) * torch.randn_like(x) return x # 对比两种方式 x_0 torch.randn(1, 3, 32, 32) t_val 500 x_t_iter q_sample_iterative(x_0, t_val) x_t_direct q_sample(x_0, torch.tensor([t_val])) # 由于随机噪声不同不能直接比较数值但可以比较统计特性 print(fIterative mean: {x_t_iter.mean():.4f}, std: {x_t_iter.std():.4f}) print(fDirect mean: {x_t_direct.mean():.4f}, std: {x_t_direct.std():.4f})注意由于每次采样的噪声是随机的两次结果不会完全相同但它们的均值和标准差应该接近。如果差异很大说明重参数化公式的实现有问题。这个验证方法虽然简单但能帮你确认代码实现的正确性避免在错误的基础上继续开发。我自己在实现 DDPM 的时候就是靠这几个实验一步步确认前向过程没问题的。尤其是第三个实验当时我发现两种方式的统计特性对不上排查了半天才发现是extract函数里的 reshape 维度写错了导致系数广播到了错误的维度上。这种 bug 不看统计特性根本发现不了因为程序不报错loss 也在下降只是生成质量差。前向扩散过程看起来只是加噪声这么简单一件事但里面的数学设计、数值稳定性、代码实现细节每一样都值得花时间吃透。把这块搞明白了后面看反向过程、看采样器、看各种改进版本都会顺畅很多。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →