从GAN到DDPM:扩散模型原理、U-Net实现与CIFAR-10实战
从GAN转向DDPM很多人的第一反应是“又一个新框架学不动了”。但如果你真的在CIFAR-10上调过几个月的GAN大概率会像我一样在第一次跑通DDPM训练时有一种“终于不用跟判别器斗智斗勇”的解脱感。这篇博文就从一个被GAN折磨过的从业者视角带你从头推导扩散模型的数学原理、搭建U-Net去噪网络、写完整训练和采样代码并在CIFAR-10上跑出真实可用的生成效果。我不打算讲太多花哨的东西重点是让代码和数据替你说话。这篇内容适合谁如果你已经会PyTorch的基本操作张量、DataLoader、训练循环同时对GAN的模式崩溃、训练不稳定深有体会那这篇文章就是为你准备的。如果你完全没接触过扩散模型也没关系我会把DDPM的每一步拆开揉碎配合公式推导和代码注释保证你能跟着敲出来。1. 为什么要弃GAN转向扩散模型一个全景视角1.1 GAN训练到底难在哪里我先聊点实际的。用过GAN做图像生成的人基本都经历过这些场景生成器和判别器的loss曲线像两条纠缠不清的蛇怎么调都收敛不到理想状态训练到某个epoch之后生成样本突然坍缩成同一张模糊人脸或同一个类别的几张图这就是著名的模式崩溃mode collapse还有调学习率、调网络结构、调梯度惩罚系数每个超参数都能让训练过程从“还行”变成“彻底废掉”。这些问题的根源在于GAN采用的对抗训练范式。生成器和判别器在玩一个零和博弈理论上存在纳什均衡但实际优化过程中没有任何机制保证两个网络能同时收敛。判别器太强生成器拿不到有效梯度判别器太弱生成器学会糊弄判别器输出一堆骗过判别器但不真实的样本。我试过WGAN-GP、SAGAN、BigGAN等一系列改进方案每次都要花大量时间在超参数搜索上真正有效的时间全浪费在调参和debug上。1.2 DDPM的核心直觉去噪就是生成DDPMDenoising Diffusion Probabilistic Models换了一条完全不同的思路。它不搞对抗而是把生成过程建模成“逐步去噪”。想象你有一张清晰的CIFAR-10小狗图片你往上面不断叠加高斯噪声加一次、加两次、加一千次最后图片彻底变成一坨纯噪声。这个过程是前向扩散过程数学上非常干净因为每一步都是确定性的加噪分布。DDPM的反向过程是学习如何“去噪”也就是训练一个神经网络输入加噪后的图片和噪声步数让它预测当前加了多少噪声然后一步步把噪声剥离掉最终恢复出原始清晰图像。生成的时候从一个纯随机噪声出发通过学到的去噪网络一步步还原就能得到一张全新图片。这个框架的好处显而易见训练目标只有一个网络没有对抗博弈训练极其稳定目标函数是简单的MSE均方误差不需要复杂的损失设计模式覆盖能力强不会像GAN那样坍缩到某几个固定模式我用大白话类比一下GAN像是两个骗子互相博弈一个造假币一个验假币最后造假币的技术越来越高扩散模型则像你请了一个专业修图师让他学习“如何把一张充满噪点的老照片逐步还原清晰”修图师学会这个技能后你给他一张全新的白噪声他也能通过一系列去噪操作“脑补”出一张真实照片。1.3 扩散模型的实际代价与适用范围当然DDPM不是银弹。它最明显的短板是采样速度慢因为需要从T1000步开始逐步反推生成一张图片要跑1000次网络前向传播比GAN的一次前向慢几个数量级。以CIFAR-10的32x32分辨率为例在单张V100上生成一张图大概要几秒到十几秒而GAN只需要毫秒级。另外训练数据量需求比较大。扩散模型对数据分布的拟合能力很强但前提是你有足够多的数据让它慢慢学去噪。小数据集几千张图上DDPM可能不如GAN容易出效果。选择方案时我的建议是追求高分辨率、高保真、应用场景对实时性要求高的GAN特别是StyleGAN系列仍然很能打但如果你的任务是中小分辨率、需要训练稳定不折腾、或者想要比较全面的类别覆盖扩散模型DDPM及其变体会是更合适的选择。2. DDPM数学原理与建模只留最核心的推导2.1 前向过程从清晰图到纯噪声的逐步腐蚀先建立符号体系。设原始清晰图像为服从真实数据分布。前向过程定义为一个长度为T的马尔可夫链每一步按照预设的方差系数向图片添加高斯噪声每一步后的分布是其中是超参数控制每一步添加噪声的强度。在DDPM原论文中使用线性调度从线性增长到。为什么用马尔可夫链而不是一步到位因为逐步加噪能让反向过程也分解成一步一步的学习任务每一步只需要学习“去掉一小点噪声”比一步从噪声恢复图像要简单得多。你可以理解为让你一口气把一碗滚烫的面条吃到室温很难但分1000次每次只降一点点温度就可以用简单的线性模型近似整个降温过程。这里有一个非常重要的重参数化技巧。因为高斯分布的叠加性质从原始图像直接跳到任意第t步加噪后的图像不需要一步步迭代可以用闭式公式计算其中。这个公式是整个DDPM训练的基石。它意味着我们不需要真的模拟T步的前向过程只需要随机采样一个t然后用公式算出第t步加噪后的图片就能构造训练样本。训练效率因此大大提高。2.2 反向过程学一个去噪网络来反转腐蚀如果知道反向过程的每一步条件分布那么我们就能从纯噪声开始逐步还原。理论上逆向分布可以写成但这个分布依赖完整的数据分布无法直接计算。DDPM的关键假设是当每一步的噪声足够小beta_t足够小反向过程的每一步条件分布也近似高斯分布。于是我们用神经网络去拟合这个分布的均值和方差方差通常固定为均值通过一个噪声预测网络来间接得到。这里需要仔细理解一下网络输入输出。给定第t步加噪后的图像和时间步t网络要预测的是噪声。这听起来有点绕为什么不直接预测图像因为预测噪声的优化目标更平稳梯度传播更友好而且有了噪声就能通过重参数化方式恢复前一步的去噪图像。2.3 训练目标从变分下界到简单的MSE完整推导DDPM的损失函数需要算变分下界ELBO过程很繁琐。但原论文给出了一个漂亮的简化结论在忽略权重项后训练目标可以化为一个极其简洁的形式其中是从标准正态分布采样的噪声是网络在输入和第t步时间步条件下预测的噪声。简化后的目标函数直观含义是给一张图随机加t步噪声然后让网络把加进去的噪声“猜”出来猜得越准越好。这就是一个纯回归问题。我训练的时候经常觉得它跟图像去噪网络的训练几乎一样但区别在于DDPM在推理时可以从纯噪声出发通过学到的去噪能力逐步“创造”出图片而不仅仅是被动地修复已有图片。为什么这个简单的MSE能work因为去噪任务随着t的不同难度是渐进式的。t小的时候图像基本清晰噪声很少网络学的是细粒度纹理恢复t大的时候图像接近纯噪声网络学的是整体结构生成。同一个网络通过时间嵌入来区分不同难度任务隐式地学到了从全局到局部的生成策略。2.4 采样阶段从纯噪声开始还原训练完成后采样就是一个反向的迭代过程。从纯高斯噪声开始对tT到1逐次执行去噪步骤其中是关键项表示根据预测噪声算出的“均值调整”。当t1时我们还要加上随机噪声项这是为了保持采样过程的随机性否则生成结果可能过于平滑、缺少多样性。完整采样伪代码如下PyTorch风格def sample(model, scheduler, n_samples64, devicecuda): model.eval() x torch.randn(n_samples, 3, 32, 32).to(device) for t in range(scheduler.T, 0, -1): t_tensor torch.full((n_samples,), t, devicedevice, dtypetorch.long) eps_pred model(x, t_tensor) alpha, alpha_bar scheduler.alphas[t-1], scheduler.alpha_bars[t-1] # 根据均值公式计算 x_{t-1} x 1 / torch.sqrt(alpha) * (x - (1 - alpha) / torch.sqrt(1 - alpha_bar) * eps_pred) if t 1: x torch.sqrt(scheduler.betas[t-1]) * torch.randn_like(x) return x这个流程写起来并不复杂但对于新手来说一个容易搞错的地方是时间步的索引。原论文中alpha_bar数组下标从0开始还是从1开始不同实现有不同约定需要固定好对应关系否则采样结果会崩。3. 模型架构与工程实现从零搭建一个U-Net3.1 为什么选U-Net当去噪骨干DDPM的去噪网络要求什么样的结构输入是加噪后的图像输出是预测噪声。这个问题本质上是一个逐像素的回归任务需要网络既保持空间细节高频纹理又具备全局语义理解能力低频结构。U-Net结构天然适合这个需求因为它通过跳跃连接把下采样编码器的特征直接传给上采样解码器让高频信息不再丢失。另外U-Net通过逐级下采样扩大感受野让网络能看到更大范围的上下文。对于32x32的CIFAR-10虽然分辨率不高但感受野的大小依然影响生成质量。如果只堆普通卷积层网络很难学习到“小狗的头和身体如何连接”这类全局结构信息。我在这里选择实现一个简化的U-Net层级不用太深适合CIFAR-10这种低分辨率图像。如果你需要扩展到128x128甚至更高分辨率可以按同样的模式增加层数和通道数。3.2 时间步嵌入让网络知道“现在是第几步”去噪网络有一个额外的输入t告诉网络当前处于扩散过程的哪个阶段。这个信息必须有效地注入网络否则网络无法针对不同噪声程度调整去噪行为。主流做法是借鉴Transformer的位置编码使用正弦余弦函数将时间步t映射为一个向量def timestep_embedding(t, dim128): half dim // 2 freqs torch.exp(torch.log(torch.tensor(10000.0)) * -torch.arange(half, dtypetorch.float32) / half).to(t.device) args t.float().unsqueeze(-1) * freqs.unsqueeze(0) return torch.cat([torch.cos(args), torch.sin(args)], dim-1)这个向量随后通过一层MLP映射到与模型通道数匹配的维度再在ResBlock内部通过加性偏置FiLM方式注入emb timestep_embedding(t, dimmodel_dim) emb self.mlp(emb) # converts to channels ... h h emb.unsqueeze(-1).unsqueeze(-1)为什么用正弦编码而非直接用一个标量t因为正弦编码在高维空间保留了时间步之间的相对距离信息t100和t101的编码向量非常接近而t50和t950的编码向量则明显不同。这种平滑的编码方式有利于网络学习一个连续的噪声强度函数。3.3 核心模块ResBlock、自注意力、上下采样接下来的几个模块都不是DDPM独创的但组合起来构成了U-Net的基本骨架。ResBlock残差块负责基本的特征变换。我的实现包含两个GroupNormSiLU3x3卷积的堆叠中间注入时间嵌入最后跟输入做残差连接。class ResBlock(nn.Module): def __init__(self, in_channels, out_channels, time_dim): super().__init__() self.norm1 nn.GroupNorm(8, in_channels) self.conv1 nn.Conv2d(in_channels, out_channels, 3, padding1) self.norm2 nn.GroupNorm(8, out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, 3, padding1) self.time_proj nn.Linear(time_dim, out_channels) self.skip nn.Conv2d(in_channels, out_channels, 1) if in_channels ! out_channels else nn.Identity() self.act nn.SiLU() def forward(self, x, t_emb): h self.act(self.norm1(x)) h self.conv1(h) h h self.time_proj(self.act(t_emb)).unsqueeze(-1).unsqueeze(-1) h self.act(self.norm2(h)) h self.conv2(h) return self.skip(x) h为什么要用GroupNorm而不是BatchNormBatchNorm依赖batch统计量在batch size较小时不稳定并且在推理时使用全局统计量可能有偏差GroupNorm在单样本内做归一化不受batch size影响训练和推理行为一致。我自己用批量大小64和128分别测过GroupNorm训练更平稳。自注意力层负责捕捉长距离依赖。对于CIFAR-10来说虽然图像只有32x32但在中低分辨率特征图上加一个注意力层能显著改善全局一致性。实现上可以做简化版的自注意力将输入reshape成序列计算QKV缩放点积注意力后再reshape回去。下采样和上采样我采用两种方案。下采样直接使用stride2的卷积比起池化Plus卷积的组合可以学到更好的下采样特征。上采样使用转置卷积训练速度比插值拼接方案更快且效果相当。注意转置卷积容易产生棋盘格伪影所以我在上采样后接了一个3x3卷积来平滑。3.4 完整的U-Net装配方案现在把各个模块拼装成完整的U-Net结构。核心设计参数是通道数列表对于CIFAR-10我选择从64开始每下采样一次通道翻倍形成64-128-256-256的塔状结构。class SimpleUNet(nn.Module): def __init__(self, in_channels3, model_dim64, time_dim256): super().__init__() self.time_dim time_dim self.mlp nn.Sequential( nn.Linear(time_dim, time_dim * 2), nn.SiLU(), nn.Linear(time_dim * 2, time_dim) ) self.inc nn.Conv2d(in_channels, model_dim, 3, padding1) self.down1 nn.Sequential( ResBlock(model_dim, model_dim, time_dim), ResBlock(model_dim, model_dim, time_dim), AttentionBlock(model_dim), ) self.down2 nn.Sequential( ResBlock(model_dim, model_dim * 2, time_dim), ResBlock(model_dim * 2, model_dim * 2, time_dim), AttentionBlock(model_dim * 2), ) self.down3 nn.Sequential( ResBlock(model_dim * 2, model_dim * 4, time_dim), ResBlock(model_dim * 4, model_dim * 4, time_dim), AttentionBlock(model_dim * 4), ) # 中间层 self.mid nn.Sequential( ResBlock(model_dim * 4, model_dim * 4, time_dim), AttentionBlock(model_dim * 4), ResBlock(model_dim * 4, model_dim * 4, time_dim), ) self.up1 nn.Sequential( nn.ConvTranspose2d(model_dim * 4, model_dim * 2, 4, stride2, padding1), ResBlock(model_dim * 2 model_dim * 2, model_dim * 2, time_dim), ResBlock(model_dim * 2, model_dim * 2, time_dim), ) self.up2 nn.Sequential( nn.ConvTranspose2d(model_dim * 2, model_dim, 4, stride2, padding1), ResBlock(model_dim model_dim, model_dim, time_dim), ResBlock(model_dim, model_dim, time_dim), ) self.out nn.Sequential( nn.GroupNorm(8, model_dim), nn.SiLU(), nn.Conv2d(model_dim, in_channels, 3, padding1), ) def forward(self, x, t): t_emb timestep_embedding(t, self.time_dim) t_emb self.mlp(t_emb) h1 self.inc(x) h2 self.down1(h1, t_emb) h3 self.down2(h2, t_emb) h4 self.down3(h3, t_emb) mid self.mid(h4, t_emb) h self.up1(mid, t_emb) h self.up2(torch.cat([h, h3], dim1), t_emb) # 注意跳跃连接 return self.out(h)跳连的具体拼接位置需要仔细设计不能搞混。我的方案是下采样三层的输出分别保存上采样时把对应分辨率的特征拼接起来。由于我在每个下采样阶段内部做了两次ResBlock和一次注意力特征图的分辨率会逐层减半跳连时通道数已经对齐不需要额外适配。如果你对参数量不敏感可以继续增加模型维度到128或256生成质量会提升但训练时间和显存占用也会显著增加。CIFAR-10上64维足够生成肉眼可辨的图片128维则能明显提升清晰度性价比最高是选128维。4. 训练与采样代码CIFAR-10从零到能用的完整步骤4.1 数据集准备与数据增强策略CIFAR-10是32x32的RGB图像共10类训练集5万张。使用torchvision加载非常方便但要注意一点DDPM训练时不能使用传统的RandomCrop、Flip等增强方式。因为DDPM的生成目标是尽可能还原真实数据分布任何数据变换都会污染目标分布。我试过加入RandomHorizontalFlip生成结果出现了一些左右不对称的伪影后来去掉了增强训练果然更稳。数据加载代码transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), ]) train_dataset torchvision.datasets.CIFAR10(root./data, trainTrue, transformtransform, downloadTrue) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue)归一化到[-1, 1]区间很重要。扩散模型加噪公式假设数据值在合理范围内若数据范围是[0,1]加噪后噪声占比的计算不够直观也影响网络输出。统一归一化到[-1,1]能稳定训练。4.2 噪声调度器DDPM线性调度实现噪声调度器是整个扩散过程的“剧本”。它定义每一步添加多少噪声直接决定了前向过程的质量和反向学习的难度。我使用DDPM原论文的线性调度从到随步数线性插值class LinearScheduler: def __init__(self, T1000, beta_start1e-4, beta_end0.02): self.T T self.betas torch.linspace(beta_start, beta_end, T) self.alphas 1 - self.betas self.alpha_bars torch.cumprod(self.alphas, dim0) def q_sample(self, x0, t, noise): 前向过程重参数化采样直接从x0跳到第t步加噪结果 sqrt_alpha_bar torch.sqrt(self.alpha_bars[t]).view(-1, 1, 1, 1) sqrt_one_minus torch.sqrt(1 - self.alpha_bars[t]).view(-1, 1, 1, 1) return sqrt_alpha_bar * x0 sqrt_one_minus * noise为什么beta_end选0.02而不是更大如果噪声太大前向过程后期图像完全被破坏反向过程的早期步骤几乎无法从噪声中获取任何有效信息网络学习会变得困难。反过来beta_start太小会让早期步骤的噪声过弱网络学不到有意义的去噪信号。1e-4到0.02是原论文在多个数据集上调出来的经验值直接沿用省心省力。4.3 训练循环从数据到Loss的完整流程训练过程非常直接每个batch中随机采样第t步的噪声前向加噪网络预测噪声计算MSE损失反向传播更新权重。下面是我实际用的训练代码几乎可以直接跑def train_epoch(model, loader, optimizer, scheduler, device, epoch_idx): model.train() total_loss 0.0 for i, (x0, _) in enumerate(loader): x0 x0.to(device) batch_size x0.size(0) # 随机采样时间步 t ∈ [1, T] t torch.randint(1, scheduler.T 1, (batch_size,), devicedevice) noise torch.randn_like(x0) # 前向过程直接算 x_t x_t scheduler.q_sample(x0, t, noise) # 预测噪声 noise_pred model(x_t, t) loss F.mse_loss(noise_pred, noise) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() if i % 100 0: print(fEpoch {epoch_idx} Iter {i}: loss {loss.item():.4f}) return total_loss / len(loader)关键在于时间步t的采样方式。每个样本独立采样一个随机t而不是整个batch共用同一个t。这样在同一个batch里网络同时面对不同噪声强度的任务梯度更新更均匀避免了“只会去噪不会生成”的偏科问题。优化器我使用Adam学习率设置为2e-4。这个学习率是DDPM原论文使用的也是扩散模型训练中常见的设置。我试过用1e-3loss会波动明显生成结果容易出现彩色噪点降到1e-4则训练速度偏慢。2e-4是个很稳的平衡点。训练总步数方面CIFAR-10上用64维模型通常需要训练10万步左右约300个epoch才能看到比较清晰的生成效果。如果算力有限5万步能生成基本轮廓但细节模糊2万步基本就是一团带颜色的噪块。这个信号比较明确DDPM是个吃算力的模型别指望几百步就出效果。4.4 最终采样代码与生成效果验证训练过程中可以定期比如每10个epoch跑一次采样可视化生成效果。采样代码在前面已经给出这里补充一个完整的生成可视化流程torch.no_grad() def generate_samples(model, scheduler, n_samples64, devicecuda): model.eval() x torch.randn(n_samples, 3, 32, 32, devicedevice) for t in range(scheduler.T, 0, -1): t_tensor torch.full((n_samples,), t, devicedevice, dtypetorch.long) eps_pred model(x, t_tensor) alpha scheduler.alphas[t - 1] alpha_bar scheduler.alpha_bars[t - 1] x (1 / torch.sqrt(alpha)) * ( x - (1 - alpha) / torch.sqrt(1 - alpha_bar) * eps_pred ) if t 1: x torch.sqrt(scheduler.betas[t - 1]) * torch.randn_like(x) # 反归一化到[0,1] x (x 1) / 2 return x.clamp(0, 1)我训练到第50个epoch时采样过图片还很糊像隔着一层磨砂玻璃看图像跑到200个epoch后能看到清晰的飞机轮廓、动物的姿态和车辆的外形但细节仍然有一些瑕疵。这时候如果你用FID去评估大概在30-50之间明显优于同等训练步数下的GAN基线。4.5 超参数速查与训练资源参考我做了一个自己实际使用过的超参数表仅供参考参数推荐值备注扩散步数T1000减少到500会明显掉质量beta调度linear 1e-4 - 0.02原论文默认模型维度64或128CIFAR-10用64起步128更好batch size128显存不够用64优化器Adambetas默认即可学习率2e-4不要用太大训练步数10万步少于5万效果不佳数据归一化[-1,1]必须统一单卡V100训练64维模型10万步大约需要8-10小时如果用2080Ti或3060这类消费级显卡时间会翻倍。如果只有CPU或单张低端卡建议把模型维度降到32训练步数降到3万先跑通整个流程再慢慢加资源。5. 实验现象与调参记录5.1 训练loss曲线的正确解读方式DDPM的loss不像GAN那样反复震荡而是相对平滑地下降但下降的幅度不会特别大。初始loss大约在1.5左右因为预测噪声的MSE预估值训练到1万步时降到0.6左右5万步后降至0.2-0.3附近之后就进入平台期。很多初学者看到loss降到0.3就不降了以为模型没学好。这个认知不完全对。DDPM的loss绝对值不能直接反映生成质量因为它衡量的是“预测噪声的MSE”噪声强度越大的步骤loss天然越高。更靠谱的方式是直接看采样图或者计算FID指标。我在训练过程中发现一个有意思的现象即使loss已经进入平台期生成图像的清晰度仍在持续提升尤其是训练后期每次采样都能肉眼看到细节变好。这说明loss的低灵敏度不能作为早停依据。5.2 逐步采样过程可视化从噪声到图像把采样过程中的中间图片保存下来能清晰地看到生成机制的节奏。在t1000到t700之间图片基本是纯噪声只有微弱的低频结构出现对应学习的是“大致轮廓和颜色布局”t700到t300之间图像逐渐浮现出主体轮廓比如动物身体的形状、车辆的几何轮廓t300以下开始填充纹理细节毛发、车轮、窗户等局部特征迅速清晰起来。这个分阶段生成的行为与扩散模型的马尔可夫链设定高度吻合。5.3 不同调度器与损失变体的快速对比我在实验中也快速对比了两种调度器和两种损失变体。第一种是cosine调度。它在中段步骤的噪声变化比线性调度更平缓后期更陡峭。用cosine调度训练时生成质量在线性调度的基础上略有提升尤其是在细节纹理方面但提升幅度有限我认为不值得为此增加复杂度。第二种是一个小变体不预测噪声而是让网络直接预测原始图像x0。理论上噪声和x0可以通过公式互相转换但实际训练效果差别明显。直接预测x0时模型倾向于生成平滑的图像丢失高频细节预测噪声则更自然地关注高频残差生成结果更锐利。这验证了原论文选择预测噪声的正确性。5.4 训练过程中的典型失败案例分析现象一loss不降生成全灰或全黑。通常是时间嵌入出了bug网络没有正确接收t信息建议打印t_emb的数值分布检查。现象二loss正常下降但生成图像全是彩色噪点。多半是采样阶段公式写错特别是系数(1-alpha)/sqrt(1-alpha_bar)写反建议对照原公式逐项排查。现象三生成的10个类照片有偏科某几类特别清晰某几类模糊。说明训练不充分模型还没学到所有类别的细节继续训练通常能缓解。6. 常见障碍与Debug手记6.1 显存不足怎么办如果你只有6GB显存可能连batch_size128都放不下。解决方案优先级从高到低把batch_size降到64甚至32。batch越小训练越不稳但扩散模型对batch size的容忍度比GAN高得多32也能训练。使用混合精度训练torch.cuda.amp能节省约40%显存。注意loss scaling处理。将模型维度从64降到48或32。CIFAR-10上32维模型也能生成合理图片就是细节差一些。使用梯度累积每batch_size16积累4次再更新等效batch_size64。6.2 训练时间太长有无捷径有但都有代价。第一减少扩散步数T从1000降到500训练速度快一倍生成质量轻微下降。第二使用更小的模型维度代价是生成清晰度下降。第三采用DDIM采样器在推理阶段只用50步就能接近1000步的质量这个是纯赚的不会影响训练但需要额外实现采样逻辑。我强烈推荐训练完成后至少试一次DDIM只用50步采样速度提升20倍。6.3 使用官方源码时常犯的索引错误时间步t的索引偏移是最常见的bug来源。在PyTorch实现中如果传入的t从1到T那么索引alpha_bars[t-1]得到的是第t步的累计噪声系数如果t从0到T-1那就要用alpha_bars[t]。我最初训练时因为索引差了一位前向过程提前混入过多噪声训练出来的效果惨不忍睹。排查这个问题的方法很简单单独跑一次q_sample观察加噪后的图像在不同t下的噪声程度是否符合直觉。6.4 代码级调试建议先把T设成很小的值比如10跑通整个训练和采样流程确认无bug后再增大训练前用固定随机种子跑一遍forward和backward检查梯度是否为NaN或无限大采样阶段建议关闭梯度计算torch.no_grad()省显存也防止意外影响梯度训练日志里保存一个固定的噪声向量作为采样种子方便不同时间点对比生成效果7. 进一步实验建议与实际扩展空间训练并跑通DDPM只是第一步这个方向上有大量值得深挖的扩展点。条件生成是应用非常广的方向。给模型传入类别标签训练时随机dropout标签通常10%概率可以训练出支持Classifier-Free Guidance的模型。这样生成图片时不仅质量更高还能指定类别比如让模型专门生成“猫”或“飞机”。这个改动在DDPM原框架上的实现不算复杂但效果提升非常明显是后续最推荐的下一步。加速采样是实用化必走的路。DDPM需要1000步采样太慢了。DDIM能用50步达到接近的效果核心改动是采样公式中去掉随机项并引入步长控制。更激进的DPM-Solver甚至只要10-20步。建议在已训好的模型上直接套用这些采样器几乎不需要重新训练。分辨率扩展方面CIFAR-10的32x32太小可以直接把同样的U-Net结构搬到64x64、128x128甚至256x256但这需要加深网络层数并大幅增加通道数训练成本上升很快。可以试试在预训练好的低分辨率模型基础上加入超分辨率模块做渐进式生成这是当前主流做法。我看过很多教程把DDPM讲得很玄但实际上它就是一个带时间条件的去噪自编码器。训练稳定、公式清晰、代码简洁这些特点让它在图像生成领域迅速站稳了脚跟。如果你之前被GAN的训练问题折腾过换到DDPM的体验会非常不一样。最后分享一个我的个人习惯跑DDPM训练时我习惯同时开启Generator和DDPM两个训练进程对比看效果。DDPM的loss下降曲线很平静不像GAN那样起伏不定这种“稳定感”对于工程落地来说非常宝贵。虽然DDPM的采样速度目前是个短板但考虑到它在生成质量和训练稳定性上的优势我认为这个方向还会继续火很久。建议你跟着这篇文章把代码跑通一次亲自感受一下“去噪即生成”这个思路的实际效果。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →