尧图精选

从离散到连续:扩散模型SDE与概率流ODE统一框架解析

🕒 发布时间:2026/9/19 4:22:17 📁 来源:尧图网络
1. 从离散到连续为什么我们需要重新理解扩散模型如果你最近在折腾生成模型大概率已经被各种Diffusion Model的变体绕晕了。DDPM、DDIM、Score-based、SDE、ODE、Flow Matching……名字一个比一个唬人但真正把这条线串起来的人其实不多。我自己最开始学的时候也是东一榔头西一棒子直到把随机微分方程SDE和概率流ODE这两块硬骨头啃下来才算是真正看懂了扩散模型连续化建模的底层逻辑。这篇内容就是把我踩过的坑、推导过的公式、以及实际写代码时遇到的工程细节完整地梳理一遍。核心目标很明确帮你建立起从离散扩散过程到连续时间SDE再到确定性概率流ODE的完整认知链条。适合已经跑过DDPM、DDIM想进一步理解Score-based模型和连续化生成建模的读者。如果你还没接触过扩散模型的基础概念建议先去补一下DDPM的前向加噪和反向去噪流程否则后面会有点吃力。先说结论扩散模型的本质是在学习一个概率分布随时间的演化过程。离散版本DDPM把这个过程切成1000步连续版本SDE把它变成无穷小步长而概率流ODE则是在这个连续框架下找到了一条确定性的演化路径。理解了这个主线后面所有的变体你都能自己推导出来。2. 离散扩散的瓶颈与连续化的动机2.1 DDPM的离散框架到底在做什么DDPM的前向过程定义得很简单每一步加一点高斯噪声经过T步之后原始数据被完全破坏成标准正态分布。数学上写成q(x_t | x_{t-1}) N(x_t; sqrt(1-β_t) x_{t-1}, β_t I)反向过程则是学习一个网络来预测每一步的噪声然后逐步去噪。这套框架work得很好但它有几个让人不舒服的地方。第一个问题是步数固定。你训练的时候用了1000步采样的时候就必须走1000步想加速就得用DDIM那种跳步策略但跳步又引入了近似误差。第二个问题是离散化误差。每一步的噪声方差β_t是人为设定的不同的schedulelinear、cosine、sqrt会带来不同的效果但理论上并没有一个“最优”的离散schedule。第三个问题是理论分析困难。你想证明采样收敛性、想分析似然下界在离散框架下做起来非常繁琐。这些问题的根源其实是一个扩散过程本质上是连续的我们人为把它切成了离散的。就像你想描述一条曲线用折线去逼近当然可以但如果你直接用微分方程来描述很多性质就自然浮现了。2.2 连续化之后我们得到了什么把离散的步长推到无穷小前向过程就变成了一个随机微分方程SDEdx f(x, t) dt g(t) dw这里f是漂移项driftg是扩散项diffusion coefficientw是标准维纳过程。这个SDE描述了一个连续时间的随机过程它的概率密度p_t(x)随时间的演化由Fokker-Planck方程控制。连续化的好处是立竿见影的。首先你可以用任意数值求解器来采样步长可以自适应调整不再受限于固定的T。其次SDE的理论工具非常成熟你可以直接借用随机分析里的结论来分析收敛性和稳定性。最后也是最重要的一点同一个前向SDE对应着无穷多条反向演化路径其中有一条是确定性的这就是概率流ODE。注意连续化不是简单的“把步数调大”而是从建模思路上把离散的马尔可夫链替换成连续时间的微分方程。这个视角转换才是关键。2.3 从Score Matching到SDE的桥梁理解连续化扩散模型绕不开Score Matching这条线。Song Yang等人在2021年的那篇Score-Based Generative Modeling through SDEs里把DDPM和Score Matching统一到了一个框架下。核心洞察是这样的反向SDE的漂移项里有一个关键量叫score function也就是∇_x log p_t(x)即对数概率密度对输入的梯度。这个量告诉你“在当前时间点往哪个方向走能让概率密度增大”。如果你能估计出每个时间点的score function你就能把前向SDE反转过来从噪声生成数据。DDPM里网络预测的噪声ε其实和score function有一个简单的线性关系score -ε / σ_t其中σ_t是当前时间步的噪声标准差。这就是为什么DDPM的网络可以直接拿来做Score-based建模——它们本质上在学同一个东西只是参数化方式不同。这个统一视角的意义在于你不再需要死守DDPM那套离散推导而是可以在连续框架下自由设计前向SDE的形式只要你能估计对应的score function就能做生成。3. SDE框架的核心细节与实操要点3.1 前向SDE的设计空间在连续框架下前向SDE的形式是dx f(x, t) dt g(t) dwf和g的选择决定了整个生成过程的性质。实际中最常用的有两种VP-SDEVariance Preserving也叫OU过程f(x, t) -0.5 β(t) x g(t) sqrt(β(t))这个形式保证前向过程的方差有界最终收敛到标准正态分布。DDPM的连续版本就是VP-SDE。VE-SDEVariance Explodingf(x, t) 0 g(t) sqrt(d[σ²(t)]/dt)这个形式下方差随时间爆炸式增长最终也收敛到正态分布但路径完全不同。Score-based模型里的NCSN就是VE-SDE的离散版本。选择哪种SDE取决于你的数据特性和采样需求。VP-SDE的方差有界数值求解时更稳定VE-SDE在低噪声区域的行为更接近恒等映射对某些图像任务效果更好。我自己的经验是图像生成任务优先试VP-SDE音频和连续信号任务可以试试VE-SDE。3.2 反向SDE的推导与实现给定前向SDE反向过程也是一个SDEdx [f(x, t) - g(t)² ∇_x log p_t(x)] dt g(t) dw̄这里dw̄是反向时间的维纳过程。这个公式是整个框架的核心它告诉你只要你能估计出score function就能从纯噪声反向演化出数据。实际实现的时候网络输出的参数化方式很关键。最常见的做法是让网络预测噪声ε_θ(x_t, t)然后通过前面的关系式转换成score。但这里有一个容易踩的坑不同的SDE对应不同的噪声尺度定义你不能直接把DDPM的噪声预测网络搬到VE-SDE上用因为σ_t的定义不一样。我自己写代码的时候习惯把score function作为网络的直接输出然后在loss里做转换。这样切换SDE类型的时候只需要改前向过程的参数网络结构不用动。3.3 采样器的选择与步长控制反向SDE的数值求解可以用Euler-Maruyama方法也可以用更高阶的Milstein方法。Euler-Maruyama最简单for t in reversed(timesteps): dt t - t_next drift f(x, t) - g(t)**2 * score(x, t) diffusion g(t) * torch.randn_like(x) * torch.sqrt(-dt) x x - drift * dt diffusion但实际用的时候有几个细节要注意。第一时间步的离散化策略很重要。均匀步长在低噪声区域可能不够精细建议用非均匀步长在噪声变化剧烈的区域多采几步。第二随机项的缩放要小心sqrt(-dt)里的符号容易搞错。第三最后一步的处理有些实现会在最后加一个去噪步骤来提升质量。实操心得如果你用Euler-Maruyama采样发现结果模糊或者有噪声残留先检查时间步的离散化是不是太粗了。把步数从100加到500试试如果质量明显提升说明是离散化误差的问题。4. 概率流ODE确定性采样的数学本质4.1 从SDE到ODE的推导概率流ODE的推导其实很直观。Fokker-Planck方程描述了概率密度随时间的演化∂p_t/∂t -∇·(f p_t) 0.5 g² Δp_t这个方程可以改写成∂p_t/∂t -∇·([f - 0.5 g² ∇ log p_t] p_t)右边这个形式就是一个连续性方程对应的是一个确定性ODEdx [f(x, t) - 0.5 g(t)² ∇_x log p_t(x)] dt这就是概率流ODE。它和反向SDE的区别在于没有随机项而且漂移项里的系数是0.5 g²而不是g²。这个ODE有一个非常重要的性质它和反向SDE产生相同的边缘概率密度p_t(x)。也就是说如果你从同一个初始分布出发沿着ODE演化得到的样本分布和沿着SDE演化得到的分布是一样的。但ODE是确定性的给定初始点轨迹完全确定。4.2 概率流ODE的三大优势第一确定性采样。同样的初始噪声每次采样得到的结果完全一样。这在需要可复现性的场景下非常重要比如科研实验或者生产环境。第二可以用高阶求解器。因为是ODE不是SDE你可以直接用Runge-Kutta、Dormand-Prince这些成熟的ODE求解器用很少的步数就能达到很高的精度。实际测试下来概率流ODE用50步左右就能达到反向SDE用1000步的质量。第三支持精确似然计算。ODE的连续性方程让你可以用瞬时变量变换公式精确计算数据的对数似然这在密度估计任务里非常有用。4.3 用概率流ODE做图像编辑概率流ODE还有一个很实用的性质它在隐空间里保持了语义结构。具体来说如果你把两张图片编码到噪声空间然后在噪声空间做插值再沿着ODE解码回来得到的中间图像会有平滑的语义过渡。这个性质被用在很多图像编辑任务里。比如你想把一张猫的图片变成狗的图片可以先分别编码得到两个噪声向量然后做球面插值再解码。中间过程会自然地经过“猫→猫狗混合→狗”的语义路径。我自己试过用这个做风格迁移效果比直接在像素空间插值好很多。关键是插值要在噪声空间做而且要用球面插值而不是线性插值因为噪声空间是高斯分布线性插值会偏离高概率区域。5. 实操流程从零实现一个连续化扩散模型5.1 环境准备与依赖我用的环境是PyTorch 2.0 CUDA 11.8主要依赖就是torch和torchvision。不需要额外的扩散模型库因为我们要自己实现核心逻辑。如果你想省事可以用HuggingFace的diffusers库做参考但我建议至少自己写一遍采样器不然很多细节理解不透。pip install torch torchvision numpy matplotlib tqdm数据集我用的是CIFAR-1032x32的图片训练和调试都比较快。如果你想做更高分辨率的建议先用小数据集把流程跑通再换。5.2 前向SDE的实现以VP-SDE为例前向过程的离散化实现import torch def vp_sde_forward(x0, t): VP-SDE的前向过程给定x0和时间t直接采样x_t 闭式解x_t sqrt(α_t) x0 sqrt(1-α_t) ε 其中α_t exp(-∫β(s)ds) beta_min, beta_max 0.1, 20.0 # 积分β(s)从0到t integral_beta beta_min * t 0.5 * (beta_max - beta_min) * t**2 alpha_t torch.exp(-integral_beta) mean torch.sqrt(alpha_t) * x0 std torch.sqrt(1 - alpha_t) noise torch.randn_like(x0) return mean std * noise, noise这里β(t)用的是线性scheduleβ(t) β_min (β_max - β_min) * t。这个schedule的选择会影响生成质量cosine schedule在低噪声区域更平滑但实现起来稍微复杂一点。5.3 Score网络的训练网络结构我用的是简化的U-Net和DDPM里的一样。关键是loss函数def compute_loss(model, x0): batch_size x0.shape[0] t torch.rand(batch_size, devicex0.device) x_t, noise vp_sde_forward(x0, t) # 网络预测噪声 noise_pred model(x_t, t) # 转换成score alpha_t torch.exp(-(0.1 * t 0.5 * 19.9 * t**2)) std torch.sqrt(1 - alpha_t) score_pred -noise_pred / std.view(-1, 1, 1, 1) # Score matching loss score_target -noise / std.view(-1, 1, 1, 1) loss ((score_pred - score_target) ** 2).mean() return loss这里有一个细节score的尺度在不同时间步差异很大。在t接近0的时候std很小score的数值会非常大。实际训练的时候需要对loss做加权常见的做法是乘以std²或者用似然加权。我试过不加权直接训结果在低噪声区域完全学不动。注意如果你发现训练loss下降但采样质量很差大概率是score的尺度问题。检查一下不同时间步的loss量级如果差异超过两个数量级就需要加权。5.4 概率流ODE采样器实现采样器我用的是Heun方法二阶Runge-Kutta比Euler精度高很多torch.no_grad() def probability_flow_ode_sample(model, shape, num_steps50): device next(model.parameters()).device x torch.randn(shape, devicedevice) timesteps torch.linspace(1.0, 0.0, num_steps 1, devicedevice) for i in range(num_steps): t_current timesteps[i] t_next timesteps[i 1] dt t_next - t_current # 计算当前点的漂移 score_current model(x, t_current) drift_current compute_drift(x, t_current, score_current) # Heun方法先走一步Euler再修正 x_euler x drift_current * dt score_euler model(x_euler, t_next) drift_euler compute_drift(x_euler, t_next, score_euler) x x 0.5 * (drift_current drift_euler) * dt return x def compute_drift(x, t, score): beta_min, beta_max 0.1, 20.0 beta_t beta_min (beta_max - beta_min) * t f -0.5 * beta_t * x g_sq beta_t drift f - 0.5 * g_sq * score return drift这个采样器用50步就能出不错的结果。如果你想更快可以试试DPM-Solver它专门为扩散模型ODE设计20步左右就能达到很好的质量。5.5 训练与采样的完整流程训练循环大概长这样model UNet(in_channels3, out_channels3) optimizer torch.optim.Adam(model.parameters(), lr2e-4) for epoch in range(num_epochs): for x0, _ in dataloader: x0 x0.to(device) loss compute_loss(model, x0) optimizer.zero_grad() loss.backward() optimizer.step() # 每几个epoch采样一次看看效果 if epoch % 10 0: samples probability_flow_ode_sample(model, (16, 3, 32, 32)) save_image(samples, fsamples_epoch_{epoch}.png)CIFAR-10上大概训200个epoch能出比较清晰的样本。如果你用更大的数据集需要相应增加训练时间和模型容量。6. 常见问题与排查技巧实录6.1 采样结果模糊或者有噪声残留这是最常见的问题。排查思路按优先级来问题现象可能原因解决方法整体模糊采样步数太少增加步数到100-200局部有噪点最后几步的离散化误差在低噪声区域加密时间步颜色偏移score尺度估计错误检查loss加权和网络输出范围样本多样性差模式崩溃检查训练数据分布和网络容量我自己的经验是80%的采样质量问题都出在时间步离散化上。先用均匀步长跑如果质量不行换成非均匀步长在t接近0的区域多采几步。6.2 训练不收敛或者loss震荡Score matching的训练比普通监督学习要敏感。几个关键点第一学习率要小。我用2e-4比一般的图像分类任务小一个数量级。第二batch size要够大。score matching的梯度方差比较大batch size小于64的时候训练很不稳定。第三时间步采样策略。均匀采样t在低噪声区域样本太少建议用重要性采样在score变化剧烈的区域多采。实操心得如果你发现loss在前几个epoch下降很快然后卡住大概率是网络容量不够或者时间嵌入维度太低。把时间嵌入从128维加到256维试试。6.3 概率流ODE和反向SDE的结果不一致理论上两者应该产生相同的分布但实际实现中可能有差异。主要原因有两个一是数值误差。ODE和SDE的离散化误差不同步数少的时候差异明显。二是随机项的缺失。SDE的随机项在某些情况下会帮助样本跳出低概率区域而ODE是确定性的可能会卡在局部模式。如果你发现ODE的结果比SDE差先增加ODE的步数。如果还是不行检查一下ODE的漂移项系数是不是0.5 g²而不是g²这个系数搞错的话结果会完全不对。6.4 如何选择SDE类型和采样器这个问题没有标准答案但可以根据任务特点来选图像生成追求质量VP-SDE 概率流ODE Heun求解器50-100步图像生成追求速度VP-SDE DPM-Solver15-20步密度估计VP-SDE 概率流ODE需要精确似然音频/连续信号VE-SDE 反向SDE保留随机性可复现性要求高概率流ODE确定性采样我自己的项目里大部分情况用VP-SDE 概率流ODE就够了。VE-SDE在图像任务上优势不明显但在一些科学计算任务里表现更好。6.5 显存不够怎么办连续化扩散模型的显存开销主要来自两个方面网络本身和采样时的中间激活。几个实用的优化技巧用混合精度训练显存直接省一半采样的时候用梯度检查点牺牲一点速度换显存如果只是采样用**torch.no_grad()**包起来减小batch size但记得同步调整学习率我在16GB显存的卡上训CIFAR-10的U-Netbatch size开到128没问题。如果你做更高分辨率的建议先用小分辨率把流程跑通再放大。7. 连续化建模的扩展方向把SDE和概率流ODE这套框架吃透之后你会发现很多新的模型都可以从这个视角去理解。比如Flow Matching它本质上是在概率流ODE的框架下直接回归一个条件速度场而不是通过score function间接得到漂移项。Rectified Flow则是在概率流ODE的基础上通过迭代拉直轨迹来加速采样。还有一个很有意思的方向是薛定谔桥它把扩散模型推广到了两个分布之间的传输问题可以做分布到分布的转换。这些内容展开讲又是另一大块但核心思想都是一样的在连续时间框架下设计概率密度的演化路径。我自己在实际项目里最常用的还是VP-SDE 概率流ODE这套组合稳定、可控、理论清晰。如果你刚开始接触连续化扩散模型建议先把这套跑通再去探索其他变体。踩过的坑告诉我基础框架理解透了后面学什么都快。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →