拓冰建站拓冰建站
首页 / 资讯中心 / 正文

从零实现DDPM:PyTorch实战扩散模型核心原理与UNet架构

1. 项目概述从零构建DDPM的动机与价值最近在复现一些经典的生成模型发现Denoising Diffusion Probabilistic ModelsDDPM虽然论文公式看着有点唬人但当你真正动手把它从零搭出来会发现其背后的思想异常清晰和优雅。很多朋友可能通过Stable Diffusion等应用已经体验了扩散模型的强大但对其核心引擎——UNet网络以及时间嵌入、扩散调度等机制的理解可能还停留在“黑箱”调用层面。这次我们就用PyTorch不依赖任何高级扩散模型库彻底拆解并复现一个完整的DDPM。这个过程不仅能让你深刻理解“噪声如何一步步变成图像”更能让你掌握自定义扩散模型、调整生成过程的能力比如修改采样步数、尝试不同的噪声调度器甚至为UNet加入注意力机制。无论你是想深入AIGC领域还是单纯对概率模型和深度学习结合感兴趣这个从零开始的搭建之旅都会让你收获颇丰。2. 核心理论拆解DDPM的前向与逆向过程要搭建模型必须先吃透它的工作原理。DDPM包含两个核心过程前向扩散过程和逆向去噪过程。2.1 前向扩散过程逐步添加噪声前向过程是一个固定的马尔可夫链它逐步向一张原始图片x0添加高斯噪声。这个过程是预先定义好的不包含任何可学习的参数。在每一步t从1到TT是总步数比如1000我们根据上一步的数据x_{t-1}得到当前加噪后的数据x_t。其数学形式如下q(x_t | x_{t-1}) N(x_t; sqrt(1 - β_t) * x_{t-1}, β_t * I)这里的β_t是一个在0到1之间预先定义好的序列称为噪声调度表。它决定了每一步添加的噪声量通常随着t增大而增大。N表示高斯分布。这个公式的意思是x_t的均值是sqrt(1 - β_t) * x_{t-1}方差是β_t。由于每一步都只依赖前一步我们可以推导出一个非常实用的性质可以从原始图像x0直接计算出任意中间时刻t的加噪图像x_t而不需要一步步迭代。q(x_t | x_0) N(x_t; sqrt(ᾱ_t) * x_0, (1 - ᾱ_t) * I)其中α_t 1 - β_tᾱ_t Π_{s1}^{t} α_s。这个性质是DDPM实现高效训练的关键。在代码中这意味着我们可以在训练时随机选择一个时间步t然后直接用这个公式采样出对应的x_t。注意β_t序列的设计至关重要。通常使用线性或余弦调度。线性调度可能导致在过程早期或晚期噪声变化过于剧烈而余弦调度如Improved DDPM中提出的能让噪声添加更平滑往往能带来更好的生成效果。我们复现时会实现这两种。2.2 逆向去噪过程神经网络学习“反扩散”如果前向过程是把一幅画慢慢涂成纯噪声那么逆向过程就是试图从纯噪声中一步步还原出那幅画。这是一个从x_T纯高斯噪声到x_0的生成过程。其每一步也定义为一个高斯分布p_θ(x_{t-1} | x_t) N(x_{t-1}; μ_θ(x_t, t), Σ_θ(x_t, t))这里的μ_θ和Σ_θ是由神经网络参数化的均值和方差。DDPM原文做了一个简化将方差Σ_θ固定为与β_t相关的常数只让神经网络学习均值μ_θ。那么神经网络学的是什么呢经过一番推导这里不展开复杂公式可以发现预测均值μ_θ等价于预测在x_t和t的条件下前向过程中所添加的噪声ε。因此DDPM的训练目标变得极其简洁训练一个噪声预测网络ε_θ。给定任意时间步t和对应的加噪图像x_t网络的目标是预测出添加到x_0上从而得到x_t的那个噪声ε。损失函数就是预测噪声和真实噪声之间的均方误差MSE。L E_{x_0, t, ε} [ || ε - ε_θ(x_t, t) ||^2 ]其中ε是从标准高斯分布中采样的随机噪声。这个简单的目标使得训练非常稳定。2.3 训练与采样算法理解了上述理论训练和采样的伪代码就一目了然训练循环从数据集中采样一个干净图像x_0。从{1, ..., T}中均匀采样一个时间步t。从标准高斯分布采样噪声ε。根据公式x_t sqrt(ᾱ_t) * x_0 sqrt(1 - ᾱ_t) * ε计算加噪图像。将x_t和t输入噪声预测网络ε_θ得到预测的噪声ε_θ。计算ε和ε_θ之间的MSE损失反向传播更新网络参数。采样生成循环从标准高斯分布采样一个随机噪声x_T。从t T到t 1循环 a. 将当前的x_t和t输入网络ε_θ得到预测噪声。 b. 根据预测噪声和公式计算x_{t-1}的均值。 c. 根据固定方差采样一些额外噪声用于随机性。 d. 计算得到x_{t-1}。循环结束后x_0即为生成的图像。3. 核心模块一时间步嵌入在DDPM中时间步t是一个标量但我们需要将它转化为网络能够利用的条件信息。直接输入标量t效果很差因此需要将其嵌入到一个高维向量空间。这里我们采用Transformer中提出的正弦位置编码的变体。3.1 正弦位置编码原理其思想是为每个时间步t生成一个唯一的高维向量并且这个向量能反映时间的顺序关系即t和t1的嵌入向量是相似的。公式如下对于嵌入向量的第i个维度emb(t)[i] sin(ω_i * t)如果i是偶数emb(t)[i] cos(ω_i * t)如果i是奇数其中ω_i 1 / (10000^(2i / d))d是嵌入向量的总维度。这种编码方式能确保不同时间步的嵌入具有区分度同时其内积能反映时间步的接近程度。3.2 PyTorch实现与集成在我们的UNet中时间步t是一个整数例如250。我们首先通过一个nn.Embedding层将其映射为一个初始向量然后通过一个由线性层和SiLU激活函数组成的小型MLP将其投影到与UNet中间特征图通道数相匹配的维度。这个最终的时间条件向量会被加到UNet的各个残差块的特征上通常是通过特征图的通道维度相加或自适应组归一化AdaGN来实现。import torch import torch.nn as nn import math class SinusoidalPositionEmbeddings(nn.Module): def __init__(self, dim): super().__init__() self.dim dim def forward(self, time): device time.device half_dim self.dim // 2 embeddings math.log(10000) / (half_dim - 1) embeddings torch.exp(torch.arange(half_dim, devicedevice) * -embeddings) embeddings time[:, None] * embeddings[None, :] embeddings torch.cat((embeddings.sin(), embeddings.cos()), dim-1) return embeddings class TimeEmbedding(nn.Module): def __init__(self, time_dim, projection_dim): super().__init__() self.time_mlp nn.Sequential( SinusoidalPositionEmbeddings(time_dim), nn.Linear(time_dim, projection_dim), nn.SiLU(), nn.Linear(projection_dim, projection_dim), ) def forward(self, t): return self.time_mlp(t)实操心得嵌入维度time_dim通常设置为256或512就足够了。projection_dim需要与UNet中应用时间条件的特征图通道数对齐。在实际添加时我更喜欢使用“自适应组归一化”AdaGN它将时间嵌入向量通过线性层映射为组归一化GroupNorm的缩放因子gamma和偏移因子beta然后应用于归一化后的特征上。这种方式比简单相加的条件注入方式更强大能更有效地指导网络在不同时间步的行为。4. 核心模块二UNet网络架构UNet是DDPM的“心脏”负责根据带噪图像x_t和时间步t预测噪声ε。它是一个编码器-解码器结构带有跳跃连接。4.1 基础构建块残差块与注意力块我们的UNet由两种基本块堆叠而成残差块和注意力块。残差块每个残差块包含两个卷积层中间有组归一化和SiLU激活函数。时间条件信息来自时间嵌入通过AdaGN注入到第一个归一化层之后。跳跃连接确保梯度流动。注意力块为了提升模型对图像全局结构的建模能力我们在UNet的底层特征图分辨率较低时插入自注意力或交叉注意力块。这里我们实现一个简单的单头自注意力机制。由于注意力机制的计算复杂度与特征图尺寸的平方成正比因此只在下采样后的低分辨率特征上使用是计算可行的。class ResidualBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.time_mlp nn.Linear(time_emb_dim, out_channels * 2) # 输出gamma和beta self.block1 nn.Sequential( nn.GroupNorm(8, in_channels), nn.SiLU(), nn.Conv2d(in_channels, out_channels, 3, padding1), ) self.block2 nn.Sequential( nn.GroupNorm(8, out_channels), nn.SiLU(), nn.Conv2d(out_channels, out_channels, 3, padding1), ) self.residual_conv nn.Conv2d(in_channels, out_channels, 1) if in_channels ! out_channels else nn.Identity() def forward(self, x, t_emb): # 第一部分 h self.block1(x) # 自适应组归一化 gamma, beta self.time_mlp(t_emb).chunk(2, dim1) h h * (gamma[:, :, None, None] 1) beta[:, :, None, None] # 第二部分 h self.block2(h) # 残差连接 return h self.residual_conv(x) class AttentionBlock(nn.Module): def __init__(self, channels): super().__init__() self.norm nn.GroupNorm(8, channels) self.qkv nn.Conv2d(channels, channels * 3, 1) self.proj_out nn.Conv2d(channels, channels, 1) def forward(self, x): b, c, h, w x.shape q, k, v self.qkv(self.norm(x)).chunk(3, dim1) # 重塑为 (b, c, h*w) 并转置k q q.view(b, c, -1).transpose(1, 2) # (b, h*w, c) k k.view(b, c, -1) # (b, c, h*w) v v.view(b, c, -1).transpose(1, 2) # (b, h*w, c) # 注意力分数 attn torch.bmm(q, k) * (c ** -0.5) # (b, h*w, h*w) attn F.softmax(attn, dim-1) # 加权求和 out torch.bmm(attn, v) # (b, h*w, c) out out.transpose(1, 2).view(b, c, h, w) # 恢复形状 return x self.proj_out(out) # 残差连接4.2 完整的UNet组装完整的UNet由下采样路径编码器和上采样路径解码器组成中间有跳跃连接。下采样通过步长为2的卷积或池化实现上采样通过转置卷积或最近邻插值卷积实现。时间嵌入向量在每个分辨率级别的残差块中注入。class UNet(nn.Module): def __init__(self, in_channels3, out_channels3, base_channels64, time_emb_dim256): super().__init__() self.time_embedding TimeEmbedding(time_emb_dim, time_emb_dim*4) # 下采样 self.down1 ResidualBlock(in_channels, base_channels, time_emb_dim) self.down2 nn.Sequential( nn.Conv2d(base_channels, base_channels, 3, stride2, padding1), # 下采样 ResidualBlock(base_channels, base_channels*2, time_emb_dim), AttentionBlock(base_channels*2), # 在低分辨率特征上加注意力 ) # ... 可以继续添加更多下采样层 # 中间层 self.mid nn.Sequential( ResidualBlock(base_channels*4, base_channels*4, time_emb_dim), AttentionBlock(base_channels*4), ResidualBlock(base_channels*4, base_channels*4, time_emb_dim), ) # 上采样 # ... 上采样层与下采样对称包含转置卷积和跳跃连接 self.up1 nn.Sequential( ResidualBlock(base_channels*4 base_channels*2, base_channels*2, time_emb_dim), # 跳跃连接拼接通道 AttentionBlock(base_channels*2), nn.ConvTranspose2d(base_channels*2, base_channels, 2, stride2), # 上采样 ) self.up2 ResidualBlock(base_channels*2, base_channels, time_emb_dim) # 再次拼接跳跃连接 self.out nn.Sequential( nn.GroupNorm(8, base_channels), nn.SiLU(), nn.Conv2d(base_channels, out_channels, 3, padding1), ) def forward(self, x, t): t_emb self.time_embedding(t) # 下采样并保存特征用于跳跃连接 h1 self.down1(x, t_emb) h2 self.down2(h1, t_emb) # ... 中间层 h_mid self.mid(h2, t_emb) # 上采样并拼接跳跃连接 h self.up1(torch.cat([h_mid, h2], dim1), t_emb) h self.up2(torch.cat([h, h1], dim1), t_emb) return self.out(h)注意事项UNet的通道数配置如base_channels64需要根据你的计算资源和图像分辨率调整。对于64x64的图片上述简化结构可能足够对于256x256或更高分辨率需要更深的网络和更多的通道数。跳跃连接是UNet的关键它帮助解码器恢复在编码器中丢失的空间细节信息。5. 核心模块三扩散调度器扩散调度器定义了前向过程中β_t序列以及与之相关的α_t和ᾱ_t序列。它不参与训练但在训练计算x_t和采样计算x_{t-1}时被频繁使用。5.1 线性调度与余弦调度线性调度这是DDPM原论文使用的方案。β_t从β_start如0.0001线性增长到β_end如0.02。β_t β_start (t/T) * (β_end - β_start)余弦调度由Improved DDPM论文提出旨在改善线性调度在过程两端变化过快的问题。它直接定义ᾱ_tᾱ_t f(t) / f(0), 其中f(t) cos((t/T s) / (1s) * π/2)^2这里s是一个小偏移如0.008防止t接近T时ᾱ_t过小导致数值不稳定。余弦调度通常能产生视觉质量更高、更平滑的生成样本。5.2 调度器的实现与缓存由于α_t,ᾱ_t等序列在训练和采样中需要反复使用我们应在初始化时预先计算并缓存它们避免重复计算。import torch import numpy as np class DDPMScheduler: def __init__(self, num_timesteps1000, beta_start1e-4, beta_end0.02, schedulelinear): self.num_timesteps num_timesteps self.schedule schedule if schedule linear: self.betas torch.linspace(beta_start, beta_end, num_timesteps) elif schedule cosine: steps num_timesteps 1 x torch.linspace(0, num_timesteps, steps) alphas_cumprod torch.cos(((x / num_timesteps) 0.008) / 1.008 * torch.pi * 0.5) ** 2 alphas_cumprod alphas_cumprod / alphas_cumprod[0] betas 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) self.betas torch.clip(betas, 0.0001, 0.9999) else: raise NotImplementedError self.alphas 1. - self.betas self.alphas_cumprod torch.cumprod(self.alphas, dim0) # ᾱ_t self.sqrt_alphas_cumprod torch.sqrt(self.alphas_cumprod) self.sqrt_one_minus_alphas_cumprod torch.sqrt(1. - self.alphas_cumprod) # 为采样过程计算参数 self.sqrt_recip_alphas torch.sqrt(1.0 / self.alphas) self.posterior_variance self.betas * (1. - self.alphas_cumprod[:-1]) / (1. - self.alphas_cumprod[1:]) def add_noise(self, original_samples, noise, timesteps): # 根据公式 x_t sqrt(ᾱ_t) * x_0 sqrt(1-ᾱ_t) * ε 添加噪声 sqrt_alpha_prod self.sqrt_alphas_cumprod[timesteps].view(-1, 1, 1, 1) sqrt_one_minus_alpha_prod self.sqrt_one_minus_alphas_cumprod[timesteps].view(-1, 1, 1, 1) noisy_samples sqrt_alpha_prod * original_samples sqrt_one_minus_alpha_prod * noise return noisy_samples def step(self, model_output, timestep, sample): # 根据预测的噪声 ε_θ计算 x_{t-1} t timestep beta_t self.betas[t].view(-1, 1, 1, 1) sqrt_one_minus_alpha_cumprod_t self.sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1, 1) sqrt_recip_alpha_t self.sqrt_recip_alphas[t].view(-1, 1, 1, 1) # 公式x_{t-1}的均值 1/sqrt(α_t) * (x_t - β_t/sqrt(1-ᾱ_t) * ε_θ) pred_original_sample sqrt_recip_alpha_t * (sample - beta_t * model_output / sqrt_one_minus_alpha_cumprod_t) mean pred_original_sample if t 0: noise torch.randn_like(sample) variance (1 - self.alphas_cumprod[t-1]) / (1 - self.alphas_cumprod[t]) * self.betas[t] std torch.sqrt(variance).view(-1, 1, 1, 1) else: std 0. noise 0. prev_sample mean std * noise return prev_sample实操心得add_noise函数用于训练时构造输入x_t。step函数用于采样时根据网络预测的噪声从x_t反推x_{t-1}。注意在t0时方差应为0因为此时应得到确定的x_0。缓存所有张量到设备CPU/GPU上能显著加速训练和采样循环。6. 训练流程完整实现将上述所有模块组合起来就构成了完整的训练流程。我们以在CIFAR-1032x32数据集上训练为例。6.1 数据准备与加载import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms from torchvision.transforms import ToTensor, Lambda, Compose import torch.nn.functional as F # 数据预处理 transform Compose([ transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) # 将图像归一化到[-1, 1] ]) train_dataset datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4)6.2 训练循环代码device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, out_channels3, base_channels64, time_emb_dim256).to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-4) scheduler DDPMScheduler(num_timesteps1000, schedulecosine) mse_loss nn.MSELoss() num_epochs 200 gradient_accumulation_steps 2 # 梯度累积模拟更大batch size optimizer.zero_grad() for epoch in range(num_epochs): model.train() total_loss 0 for step, (clean_images, _) in enumerate(train_loader): clean_images clean_images.to(device) batch_size clean_images.shape[0] # 1. 采样随机时间步和噪声 timesteps torch.randint(0, scheduler.num_timesteps, (batch_size,), devicedevice).long() noise torch.randn_like(clean_images) # 2. 根据时间步和噪声对干净图像加噪得到 x_t noisy_images scheduler.add_noise(clean_images, noise, timesteps) # 3. 模型预测噪声 noise_pred model(noisy_images, timesteps) # 4. 计算损失预测噪声与真实噪声的MSE loss mse_loss(noise_pred, noise) loss loss / gradient_accumulation_steps # 梯度累积 loss.backward() # 5. 梯度累积步骤完成后更新参数 if (step 1) % gradient_accumulation_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # 梯度裁剪防止爆炸 optimizer.step() optimizer.zero_grad() total_loss loss.item() * gradient_accumulation_steps if step % 100 0: print(fEpoch {epoch}, Step {step}, Loss: {loss.item() * gradient_accumulation_steps:.4f}) avg_loss total_loss / len(train_loader) print(fEpoch {epoch} finished. Average Loss: {avg_loss:.4f}) # 可选每个epoch结束后保存一次模型检查点 if epoch % 10 0: torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), loss: avg_loss, }, fddpm_checkpoint_epoch_{epoch}.pth)注意事项归一化输入图像被归一化到[-1, 1]模型输出的噪声预测也在同一范围。在最终生成图像时需要反归一化到[0, 1]。学习率1e-4是一个比较安全的起点。可以使用学习率预热Warmup和余弦衰减Cosine Annealing来优化训练。梯度累积当GPU内存不足以支撑大的batch size时梯度累积是有效的技巧。它通过多次前向传播累积梯度再一次性更新参数等效于增大了batch size。梯度裁剪扩散模型训练通常比较稳定但梯度裁剪可以作为一个额外的安全措施防止训练后期出现梯度爆炸。7. 采样与图像生成训练好模型后我们就可以从随机噪声开始运行逆向过程来生成图像。7.1 采样循环实现torch.no_grad() def sample(model, scheduler, image_size, batch_size16, channels3, devicecuda): 从随机噪声生成图像 model.eval() # 1. 初始化随机噪声 x_T img torch.randn((batch_size, channels, image_size, image_size), devicedevice) # 2. 从 T 到 1 循环采样 for t in reversed(range(scheduler.num_timesteps)): # 创建当前时间步的张量形状为 (batch_size,) timesteps torch.full((batch_size,), t, devicedevice, dtypetorch.long) # 3. 预测噪声 predicted_noise model(img, timesteps) # 4. 使用调度器计算前一步的 x_{t-1} img scheduler.step(predicted_noise, t, img) # 可选显示中间过程例如每100步保存一次 # if t % 100 0: # save_image(img, fsample_step_{t}.png) # 5. 将生成的图像从 [-1, 1] 反归一化到 [0, 1] img (img.clamp(-1, 1) 1) / 2.0 return img # 使用示例 generated_images sample(model, scheduler, image_size32, batch_size16, devicedevice) # 保存或显示图像 from torchvision.utils import save_image save_image(generated_images, generated_samples.png, nrow4)7.2 加速采样技巧DDIM上述采样过程需要迭代完整的T步如1000步这很耗时。Denoising Diffusion Implicit Models (DDIM) 提出了一种在保持生成质量的同时大幅减少采样步数的方法。其核心思想是定义一个非马尔科夫的逆向过程允许跳过一些中间步骤。DDIM的采样公式与DDPM不同它允许我们定义一个子序列{τ_1, τ_2, ..., τ_S}其中S可以远小于T。采样时我们只在这些子时间步上运行模型。在我们的代码中只需实现一个DDIM调度器的step函数即可替换原来的采样循环。class DDIMScheduler(DDPMScheduler): def step(self, model_output, timestep, sample, eta0.0): # eta0 对应DDIM确定性采样eta1 对应DDPM随机采样 t timestep prev_t t - self.num_timesteps // self.num_inference_steps # 假设我们定义了推理步数 alpha_prod_t self.alphas_cumprod[t] alpha_prod_t_prev self.alphas_cumprod[prev_t] if prev_t 0 else torch.tensor(1.0) beta_prod_t 1 - alpha_prod_t beta_prod_t_prev 1 - alpha_prod_t_prev # 预测 x_0 pred_original_sample (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5 # 计算 x_{t-1} 的方向 pred_sample_direction (1 - alpha_prod_t_prev - eta ** 2 * beta_prod_t_prev) ** 0.5 * model_output # 计算 x_{t-1} prev_sample alpha_prod_t_prev ** 0.5 * pred_original_sample pred_sample_direction if eta 0: noise torch.randn_like(model_output) variance (1 - alpha_prod_t_prev) / (1 - alpha_prod_t) * beta_prod_t std eta * variance ** 0.5 prev_sample prev_sample std * noise return prev_sample使用DDIM我们可以用50步甚至20步就获得与1000步DDPM采样相媲美的质量极大提升了生成效率。8. 常见问题、调试技巧与效果优化在实际搭建和训练过程中你肯定会遇到各种问题。下面是我踩过的一些坑和总结的经验。8.1 训练不稳定或损失不下降检查数据归一化确保输入图像和模型输出在预期的范围内通常是[-1,1]。一个常见的错误是输入了[0,1]的图像但没做归一化或者输出层用了错误的激活函数如Sigmoid。检查时间嵌入确保时间步t被正确嵌入并注入到UNet的每一层。可以打印中间特征图看时间条件是否有效改变了特征分布。学习率过高扩散模型对学习率比较敏感。尝试从较低的学习率如1e-5开始配合Warmup。梯度爆炸/消失使用梯度裁剪clip_grad_norm_和检查模型初始化。UNet中的卷积层可以使用He初始化或Xavier初始化。损失值范围MSE损失在训练初期应该在0.9左右因为预测随机噪声然后缓慢下降。如果损失从一开始就非常大或非常小可能是计算x_t的公式有误。8.2 生成的图像模糊或有噪声训练不充分扩散模型需要很长的训练时间才能收敛。在CIFAR-10上可能需要200-500个epoch才能看到清晰的图像。确保训练了足够的轮数。噪声调度问题尝试从线性调度切换到余弦调度。余弦调度通常能产生更清晰、细节更丰富的图像。模型容量不足对于更大分辨率如128x128的图像基础的64通道UNet可能不够。尝试增加通道数如128或加深网络层数。采样步数不足如果使用DDPM采样确保步数足够如1000步。如果使用DDIM加速可以尝试增加推理步数如100步或调整eta参数eta0为确定性采样通常更清晰eta1更随机。8.3 计算资源与性能优化混合精度训练使用torch.cuda.amp进行自动混合精度训练可以显著减少GPU内存占用并加快训练速度尤其对于大型UNet模型。梯度检查点如果GPU内存严重不足可以在UNet的某些层使用torch.utils.checkpoint以时间换空间。多GPU训练使用nn.DataParallel或nn.DistributedDataParallel进行多卡训练可以加快数据吞吐。8.4 可视化与监控监控损失曲线使用TensorBoard或WandB记录训练损失。一个健康的训练曲线应该是平滑下降的。定期采样每隔一定训练步数或epoch运行一次采样函数将生成的图像保存下来。这是判断模型是否在学习的最直观方式。你可以观察到图像从噪声逐渐变得清晰的过程。检查点管理定期保存模型检查点不仅保存模型参数也保存优化器状态和当前epoch方便从中断处恢复训练或选择不同阶段的模型进行采样比较。从零搭建DDPM是一个系统工程涉及理论理解、模块实现、训练调试和效果优化多个环节。当你看到第一张由自己编写的代码生成的、清晰的图片时那种成就感是无与伦比的。这个过程中积累的对扩散模型每个细节的掌控力是直接调用高级API无法比拟的。希望这份详细的指南能帮助你顺利走完这段旅程并为你后续探索更复杂的扩散模型如条件生成、Latent Diffusion等打下坚实的基础。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门