新闻详情

新闻详情

首页 / 资讯中心 / 详情

扩散模型原理与PyTorch实现:从加噪到去噪的完整指南

发布时间:2026/9/30 16:27:37来源:尧图网络
扩散模型原理与PyTorch实现:从加噪到去噪的完整指南
1. 从一张图说起扩散模型到底在干什么第一次看到扩散模型Diffusion Model这个词很多人会以为它和热力学里的扩散现象有什么直接关系。其实名字的来源确实借用了物理里“粒子从高浓度向低浓度扩散”的直觉——在模型里它指的是把一张清晰的图像一步步“加噪”直到变成纯噪声再训练一个网络把这个过程反过来从纯噪声一步步“去噪”还原出图像。听起来绕但核心思想就这么朴素。我最早接触这块是在做图像生成相关项目的时候。当时主流的生成方案还是生成对抗网络GAN训练不稳定、模式坍塌这些问题让人头疼。扩散模型刚出来时采样慢得离谱一张图要跑上千步但生成质量确实惊艳。后来DDIM、潜在扩散Latent Diffusion这些改进出来采样步数压到几十步甚至几步才真正具备了落地价值。Stable Diffusion就是潜在扩散模型的典型代表它把扩散过程放到VAE压缩后的潜空间里做算力需求直接降了一个数量级。这篇文章我打算把扩散模型的原理和实现从头到尾捋一遍。不是那种只贴公式的论文解读而是从“一个从业者要动手实现它需要知道什么”的角度来写。内容包括前向加噪和反向去噪的数学形式、训练目标为什么是预测噪声、U-Net骨干网络的结构设计、时间步嵌入怎么做、采样加速的几种思路最后给一份可以跑起来的PyTorch实现。适合有一定深度学习基础、想搞明白扩散模型内部机制、或者想自己动手训一个小模型的人。如果你只是想知道怎么用现成工具出图那这篇可能偏底层了但看完你会对提示词为什么有效、采样步数怎么选这些问题有更本质的理解。2. 扩散模型的核心原理拆解2.1 前向过程把图像一步步“溶解”成噪声前向过程Forward Process也叫扩散过程是一个固定的马尔可夫链。给定一张图像 (x_0)我们定义一系列时间步 (t 1, 2, ..., T)每一步都往图像里加一点高斯噪声[ q(x_t | x_{t-1}) \mathcal{N}(x_t; \sqrt{1-\beta_t} x_{t-1}, \beta_t I) ]这里的 (\beta_t) 是每一步的噪声方差通常从 (10^{-4}) 线性增加到 (0.02)总共 (T1000) 步。(\sqrt{1-\beta_t}) 这个系数是为了保持方差稳定——如果不乘这个系数加噪过程中像素值的方差会越来越大最后数值爆炸。这个式子看起来要一步步算但实际上有个闭式解可以直接从 (x_0) 跳到任意 (x_t)[ q(x_t | x_0) \mathcal{N}(x_t; \sqrt{\bar{\alpha}_t} x_0, (1-\bar{\alpha}_t) I) ]其中 (\alpha_t 1 - \beta_t)(\bar{\alpha}t \prod{s1}^{t} \alpha_s)。这个闭式解是训练能高效进行的关键——我们不需要真的模拟1000步加噪直接采样一个 (t)用公式一步算出 (x_t) 就行。用重参数化技巧写出来就是[ x_t \sqrt{\bar{\alpha}_t} x_0 \sqrt{1-\bar{\alpha}_t} \epsilon, \quad \epsilon \sim \mathcal{N}(0, I) ]我习惯把这个式子理解成“信号和噪声的加权混合”。当 (t) 很小的时候(\sqrt{\bar{\alpha}_t}) 接近1图像基本还是原样当 (t) 接近 (T) 时(\sqrt{\bar{\alpha}_t}) 接近0图像就变成了纯高斯噪声。这个从信号到噪声的渐变过程就是扩散模型名字的由来。注意(\beta_t) 的调度策略noise schedule对生成质量影响很大。早期用线性调度后来余弦调度cosine schedule被证明效果更好因为它让 (\bar{\alpha}_t) 在中间时间段下降得更平缓避免图像信息过早被破坏。2.2 反向过程训练一个网络学会“去噪”反向过程Reverse Process才是真正需要学习的地方。我们希望从纯噪声 (x_T \sim \mathcal{N}(0, I)) 出发一步步去噪最终得到一张清晰的图像 (x_0)。理论上如果知道真实的反向分布 (q(x_{t-1}|x_t))就能精确还原。但这个分布依赖于整个数据集没法直接算。于是我们用神经网络 (p_\theta) 来近似它[ p_\theta(x_{t-1} | x_t) \mathcal{N}(x_{t-1}; \mu_\theta(x_t, t), \Sigma_\theta(x_t, t)) ]关键洞察来了当 (\beta_t) 足够小的时候反向过程也可以近似为高斯分布。所以我们只需要让网络预测这个高斯分布的均值 (\mu_\theta) 和方差 (\Sigma_\theta)。进一步推导可以发现均值 (\mu_\theta) 可以写成[ \mu_\theta(x_t, t) \frac{1}{\sqrt{\alpha_t}} \left( x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}t}} \epsilon\theta(x_t, t) \right) ]也就是说网络真正需要预测的是噪声 (\epsilon_\theta(x_t, t))。这就是为什么扩散模型的训练目标通常是“预测噪声”——不是直接预测图像而是预测当前步加入的噪声然后用它反推出干净的图像。这个设计非常巧妙。直接预测 (x_0) 的话网络要学的东西跨度太大从纯噪声到清晰图像难度很高。而预测噪声相当于让网络专注于“这一步的噪声长什么样”任务更局部、更稳定。我实测下来预测噪声的收敛速度和最终质量都明显优于直接预测 (x_0)。2.3 训练目标一个简化到极致的损失函数有了上面的推导训练目标就变得非常简洁。原始论文从变分下界ELBO出发推导最后化简成一个均方误差[ L_{\text{simple}} \mathbb{E}{t, x_0, \epsilon} \left[ | \epsilon - \epsilon\theta(x_t, t) |^2 \right] ]训练流程就是从数据集采样一张图像 (x_0)随机采样一个时间步 (t \sim \text{Uniform}(1, T))采样噪声 (\epsilon \sim \mathcal{N}(0, I))计算 (x_t \sqrt{\bar{\alpha}_t} x_0 \sqrt{1-\bar{\alpha}_t} \epsilon)让网络预测 (\epsilon_\theta(x_t, t))计算MSE损失反向传播就这么简单。没有对抗训练没有复杂的损失平衡就是一个回归问题。这也是扩散模型比GAN好训的根本原因——它把生成问题转化成了一个去噪自编码器的训练问题。实操心得虽然理论上 (t) 是均匀采样但实际训练中可以对 (t) 做重要性采样让模型更关注那些“去噪难度大”的时间步。我试过在中间时间段(t) 在300到700之间加大采样权重FID指标有轻微改善但提升有限不如把精力放在网络结构和采样策略上。3. 网络架构与关键组件实现3.1 U-Net骨干为什么是它而不是Transformer扩散模型的去噪网络 (\epsilon_\theta(x_t, t)) 需要满足两个要求输入输出尺寸一致且能捕捉多尺度特征。U-Net天然符合这两个条件。U-Net的结构是编码器-解码器加跳跃连接。编码器逐层下采样提取从细粒度到粗粒度的特征解码器逐层上采样恢复空间分辨率跳跃连接把编码器同层的特征直接拼到解码器保留细节信息。对于去噪任务来说这个设计非常合适——低层特征帮助恢复纹理细节高层特征帮助理解全局结构。具体到扩散模型用的U-Net和原始医学图像分割用的U-Net有几个区别时间步嵌入每个残差块都要注入时间步信息让网络知道当前是第几步去噪自注意力层在低分辨率层加入自注意力捕捉长距离依赖组归一化用GroupNorm而不是BatchNorm因为训练时batch size通常很小我实现的时候用的配置是基础通道数64通道倍数(1, 2, 4, 8)每个分辨率层2个残差块注意力分辨率设在16x16和8x8。这个配置在256x256图像上大约有50M参数单卡24G显存可以跑batch size 16左右。3.2 时间步嵌入让网络知道“现在是第几步”时间步 (t) 是一个标量但网络需要它来调节每一层的特征。做法和Transformer的位置编码类似用正弦函数生成一个高维向量import torch import math def timestep_embedding(timesteps, dim, max_period10000): half dim // 2 freqs torch.exp( -math.log(max_period) * torch.arange(half, dtypetorch.float32) / half ).to(timesteps.device) args timesteps[:, None].float() * freqs[None] embedding torch.cat([torch.cos(args), torch.sin(args)], dim-1) if dim % 2: embedding torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim-1) return embedding这个嵌入向量经过两层MLP后加到每个残差块的特征上。为什么用正弦编码而不是直接用一个可学习的embedding因为正弦编码有外推性训练时见过 (t1) 到 (T)推理时如果要用不同的步数调度编码仍然合理。可学习embedding就没这个好处。注意时间步嵌入的维度要和残差块的通道数匹配。我一般设成基础通道数的4倍比如基础通道64嵌入维度就是256。太小了信息容量不够太大了浪费参数。3.3 残差块与注意力去噪网络的基本单元每个残差块的结构是GroupNorm → SiLU激活 → 卷积 → 注入时间步嵌入 → GroupNorm → SiLU → 卷积 → 残差连接。时间步嵌入通过一个线性层投影后加到第一个卷积的输出上。class ResBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.norm1 nn.GroupNorm(32, in_channels) self.conv1 nn.Conv2d(in_channels, out_channels, 3, padding1) self.time_mlp nn.Linear(time_emb_dim, out_channels) self.norm2 nn.GroupNorm(32, out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, 3, padding1) self.skip nn.Conv2d(in_channels, out_channels, 1) if in_channels ! out_channels else nn.Identity() def forward(self, x, t_emb): h self.conv1(F.silu(self.norm1(x))) h h self.time_mlp(F.silu(t_emb))[:, :, None, None] h self.conv2(F.silu(self.norm2(h))) return h self.skip(x)自注意力层用在16x16和8x8分辨率上。注意力计算是标准的QKV形式但加了一个残差连接和归一化。我试过在更高分辨率也加注意力显存直接爆了而且收益不明显——高分辨率层的卷积已经能捕捉足够的局部信息。3.4 潜在扩散把计算搬到潜空间直接在像素空间做扩散256x256的图就要处理196608维的数据训练和采样都慢。潜在扩散Latent Diffusion的思路是先用一个VAE把图像压缩到潜空间比如8倍下采样后变成32x32x4维度降到4096计算量减少几十倍。VAE的编码器把图像 (x) 映射到潜变量 (z \mathcal{E}(x))解码器从潜变量重建图像 (\hat{x} \mathcal{D}(z))。扩散过程在 (z) 上进行训练目标不变只是 (x_0) 换成了 (z_0)。采样时从噪声生成 (z_0)再解码成图像。这个设计的关键是VAE的重建质量要足够好否则扩散模型生成的东西会被VAE的瓶颈限制。Stable Diffusion用的VAE下采样8倍潜空间通道4重建质量在感知上几乎无损。我自己训VAE的时候发现KL正则的权重很关键——太大导致重建模糊太小导致潜空间分布太散扩散模型学起来困难。一般设 (10^{-6}) 到 (10^{-4}) 之间比较合适。4. 完整实现与训练流程4.1 环境准备与依赖我用的环境是Python 3.10 PyTorch 2.0 CUDA 11.8。依赖不多pip install torch torchvision einops accelerate tensorboardeinops用来做张量重排比原生permute可读性好很多。accelerate处理混合精度和分布式训练。数据我用的是CIFAR-10和CelebA-64做实验前者32x32后者64x64单卡就能训。4.2 数据加载与预处理数据预处理很简单归一化到[-1, 1]就行。扩散模型对数据增强不敏感因为加噪过程本身就是一种强增强。我试过加随机翻转效果没有明显变化。from torchvision import datasets, transforms transform transforms.Compose([ transforms.Resize(64), transforms.CenterCrop(64), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize([0.5]*3, [0.5]*3) ]) dataset datasets.CelebA(root./data, splittrain, transformtransform, downloadTrue) dataloader torch.utils.data.DataLoader(dataset, batch_size64, shuffleTrue, num_workers4)4.3 噪声调度与扩散参数噪声调度我用余弦调度比线性调度在低噪声区域更平缓def cosine_beta_schedule(timesteps, s0.008): steps timesteps 1 x torch.linspace(0, timesteps, steps) alphas_cumprod torch.cos(((x / timesteps) s) / (1 s) * math.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)预计算好 (\bar{\alpha}_t)、(\sqrt{\bar{\alpha}_t})、(\sqrt{1-\bar{\alpha}_t}) 这些系数训练时直接查表避免重复计算。4.4 训练循环与关键参数训练循环的核心就是前面说的五步。我用混合精度训练显存占用减少约40%速度提升30%左右。scaler torch.cuda.amp.GradScaler() for epoch in range(num_epochs): for batch in dataloader: x0 batch[0].cuda() t torch.randint(0, T, (x0.shape[0],), devicex0.device) noise torch.randn_like(x0) xt sqrt_alphas_cumprod[t][:, None, None, None] * x0 \ sqrt_one_minus_alphas_cumprod[t][:, None, None, None] * noise with torch.cuda.amp.autocast(): noise_pred model(xt, t) loss F.mse_loss(noise_pred, noise) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()关键参数学习率用2e-4AdamW优化器weight decay 0.01EMA衰减率0.9999。EMA对生成质量影响很大我试过不用EMAFID直接差一截。batch size 64训练了大概200k步在单张A100上跑了约3天。实操心得训练初期loss下降很快但别高兴太早那只是模型学会了预测均值。真正决定生成质量的是中后期loss下降很慢但FID在持续改善。我一般每10k步采样一批图看看效果比只看loss曲线靠谱。4.5 采样从噪声生成图像训练完之后采样就是从 (x_T \sim \mathcal{N}(0, I)) 出发逐步去噪。标准DDPM采样要跑1000步太慢。实际用DDIM采样50步就能出不错的结果torch.no_grad() def ddim_sample(model, shape, steps50, eta0.0): x torch.randn(shape).cuda() timesteps torch.linspace(T-1, 0, steps).long().cuda() for i in range(steps): t timesteps[i] prev_t timesteps[i1] if i1 steps else -1 noise_pred model(x, t.unsqueeze(0)) alpha_t alphas_cumprod[t] alpha_prev alphas_cumprod[prev_t] if prev_t 0 else torch.tensor(1.0) x0_pred (x - torch.sqrt(1-alpha_t) * noise_pred) / torch.sqrt(alpha_t) x0_pred x0_pred.clamp(-1, 1) sigma eta * torch.sqrt((1-alpha_prev)/(1-alpha_t) * (1-alpha_t/alpha_prev)) c torch.sqrt(1-alpha_prev-sigma**2) x torch.sqrt(alpha_prev) * x0_pred c * noise_pred if eta 0: x x sigma * torch.randn_like(x) return xeta0就是确定性DDIMeta1退化成DDPM。我一般用eta050步生成一张64x64的图在A100上约0.5秒。如果追求更高质量可以加到100步但边际收益递减明显。5. 常见问题与排查实录5.1 生成图像模糊或颜色失真这是最常见的问题。排查顺序现象可能原因排查方法解决整体模糊训练不充分看FID是否还在下降继续训练颜色偏灰归一化问题检查数据预处理确认归一化到[-1,1]局部模糊VAE重建瓶颈单独测VAE重建降低KL权重或增大潜空间网格状伪影上采样方式检查U-Net上采样用最近邻卷积替代转置卷积我踩过最坑的一次是颜色整体偏绿查了半天发现是数据加载时通道顺序搞错了RGB当BGR用了。这种低级错误反而最难发现因为loss曲线看起来完全正常。5.2 训练loss不下降或震荡先检查学习率是不是太大。扩散模型对学习率比较敏感2e-4是个比较稳的值超过5e-4容易震荡。如果loss一直不降检查时间步嵌入有没有正确注入——我见过有人忘了把时间步嵌入加到残差块里网络根本不知道自己在第几步loss卡在0.5左右下不去。另一个常见问题是EMA没开或者衰减率设错。EMA衰减率一般设0.999到0.9999太小了起不到平滑效果太大了更新太慢。我习惯用0.9999训练步数超过100k之后效果明显。5.3 采样步数与质量的关系很多人问采样步数怎么选。我的经验是20步以下细节丢失明显适合快速预览50步质量和速度的平衡点日常用这个100步质量提升有限除非做对比实验200步以上基本没有可见提升DDIM的步数不需要和训练步数一致这是它比DDPM灵活的地方。训练用1000步采样用50步完全没问题。但要注意采样步数太少时eta0的确定性采样可能陷入局部最优适当加一点随机性eta0.2左右有时能改善多样性。5.4 显存不够怎么办扩散模型训练显存占用主要来自激活值。几个有效的优化混合精度训练省40%左右梯度检查点省60%但慢30%减小batch size最直接但影响BN统计不过我们用GroupNorm影响小潜在扩散省一个数量级我最早在2080Ti上训64x64的模型11G显存batch size只能开到8。后来换用梯度检查点开到16训练时间从5天缩到4天还算划算。6. 几个值得深挖的扩展方向扩散模型这块发展太快我列几个自己关注的方向。一是采样加速除了DDIM还有DPM-Solver、一致性模型Consistency Models后者能做到一步生成虽然质量还有差距但进步很快。二是条件生成classifier-free guidance是目前的主流做法训练时随机丢掉条件采样时用引导系数控制条件强度系数设7到10之间比较常见。三是与Transformer的结合DiTDiffusion Transformer用Transformer替换U-Net在ImageNet上已经刷到了很好的FID而且扩展性更好模型越大效果越好。我自己最近在试的是把扩散模型用到非图像领域比如音频生成和分子结构生成。核心思路是一样的只是数据形式和网络结构要调整。音频用1D卷积或Transformer分子用图神经网络。踩过的坑是不同领域的数据分布差异很大噪声调度需要重新调不能直接套图像的那套参数。最后分享一个小技巧如果你只是想快速验证一个想法不用从头训。拿预训练的Stable Diffusion冻结VAE和文本编码器只微调U-Net的注意力层用LoRA几张图就能出效果。我试过用20张图微调一个特定风格在A100上跑了15分钟就有模有样了。这个方法适合做风格迁移和小样本适配比从头训划算太多。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

广东省PVC镭射手提袋制造厂家实力分析:欧强塑料包装厂广受信赖 2026/9/30 17:24:47

广东省PVC镭射手提袋制造厂家实力分析:欧强塑料包装厂广受信赖

走进广东商业包装市场,PVC镭射手提袋凭借自带的彩虹炫彩视觉效果,已经成为美妆礼盒、展会伴手礼、明星应援、服装购物等多个场景的热门包装选择。越来越多中小商家开始寻找能承接小批量定制、品质稳定的生产厂家,在众多供应商中,苍…

阅读更多 →
GEO优化选购技巧:众量引擎本地化搜索排名与外贸询盘获取双效策略 2026/9/30 17:24:47

GEO优化选购技巧:众量引擎本地化搜索排名与外贸询盘获取双效策略

随着生成式AI技术的快速普及,大众获取信息的方式正在发生底层变化:越来越多的用户习惯通过ChatGPT、豆包、文心一言等生成式AI引擎提问,获取品牌、产品、服务相关的答案,企业的流量入口已经从传统搜索引擎拓展到了AI生态。 在这样…

阅读更多 →
智能学术搜索:助力科研高效获取精准学术资源的专业检索工具 2026/9/30 17:24:21

智能学术搜索:助力科研高效获取精准学术资源的专业检索工具

导师一句“做AIXX”,很多研究生其实卡在第一步:不知道从哪开始连。 不是不努力,而是跨学科的本质,从来不是多学一点,而是找到——两个领域之间真正能对接的“接口”。 问题在于,这些接口往往是隐形的&…

阅读更多 →
国庆限时招募|给你的 Coding Agent 装上长期记忆 2026/9/30 17:24:00

国庆限时招募|给你的 Coding Agent 装上长期记忆

AI 编程工具已经能够帮助开发者完成代码编写、Bug 排查和功能迭代,但当项目进入长期维护阶段,新的问题也会逐渐出现:项目背景需要反复解释,代码结构需要重新说明,之前排查过的 Bug 和失败方案可能再次被尝试。随着项目…

阅读更多 →
【雷达系统学习笔记 07】PRF 追问专场:测速测的是谁、距离与速度的模糊矛盾 2026/9/30 17:23:54

【雷达系统学习笔记 07】PRF 追问专场:测速测的是谁、距离与速度的模糊矛盾

第 3 章我连追了三个问题,全部问到了 PD 雷达的物理本质上:PRF 是什么?测速测的是谁的速度?为什么 PRF 越高测速越清晰? 核心一句话:PRF 就是采样率——距离和速度,一个要发得慢(等回…

阅读更多 →
我的 Claude Code 最佳实践  7 个经过验证的工作流技巧 2026/9/30 17:23:47

我的 Claude Code 最佳实践 7 个经过验证的工作流技巧

用 Claude Code(下称 CC)半年多,累计消耗数十亿 tokens。踩过不少坑,也读了大量官方文档和实践者分享,慢慢沉淀出一套自己的工作流。这篇不追求面面俱到,只讲经过长期验证、确实能提升交付质量的 7 个实践。…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉