ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

PyTorch 扩散模型图像修复:从原理到 DDIM 采样与重贴

PyTorch 扩散模型图像修复:从原理到 DDIM 采样与重贴 简介这份资源面向图像处理与深度学习方向的研究者和开发者聚焦扩散模型在图像修复任务中的创新应用帮助读者理解如何从零搭建训练流程并完成一套可复现的实验分析。内容以PyTorch为基础涵盖环境配置、多数据集加载与归一化预处理、U-Net网络定义以及将注意力建模与残差连接融合而成的“Attention-Residual Diffusion Model (ARDM)”并完整展开训练、对比实验、消融实验与指标评估环节。资源为1个docx文档压缩包约18KB体量轻便可当作实验框架与代码讲解的参考索引便于按章节查阅扩散与修复的对应实现。文中还讨论了创新模块的设计思路及性能提升原因有助于读者把握模型改进的动机与验证方式。目前已有245人学习下载适合希望快速进入扩散模型图像修复方向、需要对照代码理解方法细节的读者参考。1. 从一张缺了半张脸的老照片说起一张被水渍啃掉半张脸的老合影传统 inpainting 给出来的往往是一团糊状纹理——它知道这里该有皮肤却不知道这张脸该长什么样。扩散模型不直接预测缺失像素而是在“整张照片看起来像真实照片”这个约束下从纯噪声里一步步推出一个合理答案。容易被低估的一点是修复的目标不是像素级准确而是语义与纹理同时可信。PSNR 高不等于肉眼好看大面积缺失、结构断裂的场景里回归式模型倾向输出均值化的模糊块扩散采样能保住高频细节结果被拉回自然图像的流形附近。下面先拆原理再落到一份能直接跑的最小实现最后给实验设计、参数表和失败模式的排查路径。会写 PyTorch、想系统跑一遍修复任务的人可以顺着往下看。2. 扩散模型做图像修复的原理拆解从去噪公式到掩码条件注入2.1 前向加噪与反向去噪两条公式先立住扩散模型的骨架非常简洁一条前向链把图像逐步污染成高斯噪声一条反向链学着把噪声逐步还原。前向过程是固定的马尔可夫链不需要学习反向过程才是要训练的部分。前向每一步加一点噪声q(x_t | x_{t-1}) N(x_t; sqrt(1-β_t) · x_{t-1}, β_t · I)这个式子没法直接采样到第 t 步得迭代 t 次。但因为高斯分布可加作者用重参数化把任意步写成闭式x_t sqrt(ᾱ_t) · x_0 sqrt(1 - ᾱ_t) · ε, ε ~ N(0, I)其中ᾱ_t ∏(1-β_i)从 1 单调衰减到 0。这一行是整个训练能高效跑起来的关键随便抽一个 t一次采样就能构造出对应的含噪图像不用循环 t 次。反向过程参数化成高斯p_θ(x_{t-1} | x_t) N(μ_θ(x_t, t), Σ_θ(x_t, t))工程上普遍把方差 Σ 固定成 β_t 或它的插值只让网络预测均值。再往下推会得到一种更稳定的等价形式不预测 x_0而是直接预测这一步注入的噪声 ε。损失就是预测噪声和真实噪声之间的 MSE。import torch def q_sample(x0, t, alphas_cumprod): x0: 干净图像 [B, 3, H, W] t: 时间步索引 [B]取值 0..T-1 alphas_cumprod: 预计算的 ᾱ 表形状 [T] 返回加噪后的 x_t 以及本步实际使用的噪声 a_bar alphas_cumprod[t].view(-1, 1, 1, 1) # 每个样本取自己的 ᾱ_t noise torch.randn_like(x0) # ε ~ N(0, I) x_t a_bar.sqrt() * x0 (1 - a_bar).sqrt() * noise # 闭式重参数化 return x_t, noise逻辑说明alphas_cumprod在训练开始前一次性算好避免每个 batch 重复计算。view(-1,1,1,1)是为了让标量广播到[B,3,H,W]这一步漏掉会直接报维度不匹配。参数说明t一般用均匀分布随机采样也可以用重要性采样偏向中间步alphas_cumprod的调度方式线性、余弦、sigmoid会明显影响细节还原后面第 4 章会给对照。2.2 图像修复为什么天然适配扩散范式修复任务的定义是给定一张有缺失区域的图输出一张在缺失区域内容合理、在已知区域尽量不变的完整图。这本质上是一个条件生成问题——已知区域是条件缺失区域要从图像先验里采样。回归式模型U-Net、GAN 的 L1 分支在这个问题上的通病是均值化。损失函数每个像素独立算误差当缺失区域存在多种合理解释时最优解就是它们之间的加权平均视觉结果是模糊。扩散采样不取平均它沿着得分函数的方向走天然偏向流形上概率高的点。“流形”这个词在这里不是装饰。自然图像在高维像素空间里只占极薄的一层低维结构随机噪声几乎不可能落在上面。扩散的反向过程每步都在往概率密度高的方向推累积 T 步之后结果被拉回自然图像的流形附近而不是悬在流形之外的“中间态”。这解释了为什么扩散修复的纹理更真实也解释了为什么采样步数太少时结果会显得生硬——步数不够没走完这段回推路径。另一个好处是已知区域的硬约束很容易施加。修复任务里有一个天然可用的操作每一步采样后把已知区域的像素替换回原始值只让缺失区域继续去噪。这个动作叫重贴已知区域重投影实现成本几乎为零却是扩散修复效果明显优于纯生成的关键。2.3 条件注入的三种主流接法2.3.1 掩码拼接最省事也最稳把掩码当成额外通道和噪声图拼在一起送进网络输入从 3 通道变 4 通道。优点是改动量小、不需要改注意力结构训练稳定代价是网络要自己学会把掩码当作硬约束如果掩码表达方式不好比如用 0/1 而不是 -1/1有时会忽略它。常见做法是把已知区域也做同样的加噪再乘上(1-mask)让网络看到的信息是自洽的。import torch.nn as nn class MaskedUNet(nn.Module): def __init__(self, base_unet): super().__init__() self.unet base_unet def forward(self, x_t, t, mask): # x_t: [B,3,H,W] 当前含噪图mask: [B,1,H,W]1 表示待修复区域 inp torch.cat([x_t * (1 - mask), mask], dim1) # 缺失区域置零掩码单独成通道 return self.unet(inp, t) # 输出预测噪声 [B,3,H,W] 逻辑说明x_t * (1 - mask) 把待修复区域的噪声抹掉避免网络从这块“无效信息”里学捷径掩码通道单独保留告诉网络哪些位置需要生成、哪些位置需要保留。参数说明如果换成无掩码训练网络仍然能学去噪但推理时无法指定修复位置只能做整图生成。 #### 2.3.2 注意力重加权 在 U-Net 的 self-attention 层里对已知区域的 key/value 做加权让缺失区域的 query 更多地从已知区域借信息。实现上是在注意力矩阵上加一个由掩码导出的偏置项。它比掩码拼接能利用更远的上下文适合大缺口、跨区域结构延续的场景代价是要改网络结构调参空间也更大。 #### 2.3.3 潜在空间融合 先把图像编码到 VAE 的潜空间在低维潜空间上做扩散掩码也同步下采样。这是潜在扩散模型LDM的思路显存和速度优势明显代价是 VAE 的重建误差会限制修复精度的上限。选它还是选像素空间取决于算力预算和细节要求。 ### 2.4 像素空间扩散 vs 潜在扩散模型选型对照 | 维度 | 像素空间扩散 | 潜在扩散模型 | | --- | --- | --- | | 表示分辨率 | H×W×3 | (H/8)×(W/8)×CC 一般 4 | | 显存占用 | 高512×512 以上吃紧 | 低同分辨率下约为像素方案的四分之一到八分之一 | | 单步采样成本 | 高 | 低可跑到更多步数 | | 高频细节上限 | 高不受 VAE 重建限制 | 受 VAE 重建误差压制细纹理易被磨平 | | 训练数据需求 | 大小数据集易过拟合 | 相对小潜空间本身是压缩先验 | | 适用场景 | 老照片、人脸、医学影像等细节敏感任务 | 大规模自然图像修复、需要快速迭代的场景 | 选型建议按任务走如果缺失区域以结构和大块纹理为主且希望单卡就能跑选潜在扩散如果目标是像素级细节、后续要放大或做打印输出选像素空间或者用潜在扩散出一版粗结果再用像素模型细化。 ## 3. 用 PyTorch 手写一个最小可跑的图像修复扩散模型 ### 3.1 环境确认与依赖清单 先确认 CUDA 可用不然训练一轮的时间会劝退。依赖不复杂注意 einops 用来写张量重排会很省事。 bash python -m venv venv source venv/bin/activate pip install torch torchvision einops pillow tqdm python -c import torch; print(cuda:, torch.cuda.is_available(), torch.cuda.get_device_name(0))逻辑说明虚拟环境避免污染全局依赖最后一行用来验证 PyTorch 版本和 GPU 是否匹配。参数说明如果cuda打印 False先检查驱动和 CUDA 版本对应关系而不是直接改代码混合精度依赖的torch.cuda.amp在较新版本里已经和 Tensor Core 绑定不匹配会静默降速。3.2 噪声调度与时间步嵌入调度决定训练难度。线性调度在前向过程末端会过度破坏信息余弦调度在低噪声区更密集细节还原更稳。import math import torch def cosine_schedule(T1000, s0.008): 余弦噪声调度返回 betas [T] 和 alphas_cumprod [T] steps torch.arange(T 1, dtypetorch.float64) f torch.cos((steps / T s) / (1 s) * math.pi / 2) ** 2 alphas_cumprod f / f[0] # ᾱ_t首项归一化为 1 betas 1 - alphas_cumprod[1:] / alphas_cumprod[:-1] return betas.clamp(1e-8, 0.999).float(), alphas_cumprod[1:].float() def timestep_embedding(t, dim): 正弦位置嵌入把整数时间步映射成连续向量 half dim // 2 freqs torch.exp(-math.log(10000) * torch.arange(half, devicet.device) / half) args t[:, None].float() * freqs[None] return torch.cat([torch.sin(args), torch.cos(args)], dim-1)逻辑说明cosine_schedule从余弦函数反推 β保证 ᾱ_T 接近 0clamp防止端点出现 0 或 1 导致除零或完全丢失信号。参数说明s是一个很小的偏移用来避免 t0 附近噪声过小dim一般取 128 或 256越大对时间步越敏感但也更容易过拟合。3.3 带掩码条件的 U-Net 改造最小实现用残差块堆一个下采样—上采样结构即可重点是输入输出通道和跳跃连接。张量形状含义x_t[B, 3, H, W]当前含噪图mask[B, 1, H, W]1 表示待修复t_emb[B, dim]时间步嵌入cond[B, 4, H, W]拼接后的网络输入eps_pred[B, 3, H, W]预测噪声import torch import torch.nn as nn from einops import rearrange class ResBlock(nn.Module): def __init__(self, in_ch, out_ch, t_dim): super().__init__() self.norm1 nn.GroupNorm(8, in_ch) self.conv1 nn.Conv2d(in_ch, out_ch, 3, padding1) self.t_proj nn.Linear(t_dim, out_ch) # 时间步嵌入投影到通道维 self.norm2 nn.GroupNorm(8, out_ch) self.conv2 nn.Conv2d(out_ch, out_ch, 3, padding1) self.skip nn.Conv2d(in_ch, out_ch, 1) if in_ch ! out_ch else nn.Identity() def forward(self, x, t_emb): h self.conv1(torch.nn.functional.silu(self.norm1(x))) h h self.t_proj(t_emb)[:, :, None, None] # 把时间信息广播进空间特征 h self.conv2(torch.nn.functional.silu(self.norm2(h))) return h self.skip(x)逻辑说明时间步嵌入不做广播的话网络无法区分自己处在去噪的哪个阶段训练几乎不会收敛GroupNorm 在 batch 较小时比 BatchNorm 稳。参数说明in_ch对第一层是 43 通道图加 1 通道掩码后续可以逐级加到 128、256残差跳连用于保留低频结构信息。3.4 训练循环里必须盯的四个量for step, (x0, mask) in enumerate(loader): x0, mask x0.cuda(), mask.cuda() t torch.randint(0, T, (x0.size(0),), devicex0.device) x_t, noise q_sample(x0, t, alphas_cumprod.cuda()) pred model(x_t, timestep_embedding(t, 128), mask) loss torch.nn.functional.mse_loss(pred, noise) optimizer.zero_grad() loss.backward() grad_norm torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() ema.update(model) # 指数滑动平均要盯的四个量训练 loss 是否平滑下降而不是抖动梯度范数是否长期贴着裁剪阈值EMA 权重和原始权重的采样差异固定验证集上的 PSNR 是否在 loss 下降时同步改善。参数说明clip_grad_norm_的 1.0 是常用的保守值梯度爆炸时先降学习率再看EMA 的衰减率一般取 0.999太小起不到平滑作用太大会让采样长期滞后于最新权重。3.5 采样阶段把已知区域逐步贴回去torch.no_grad() def inpaint_sample(model, x_known, mask, alphas_cumprod, steps50): DDIM 风格加速采样每步把已知区域重贴回去 b, c, h, w x_known.shape x torch.randn_like(x_known) # 从纯噪声起步 ts torch.linspace(len(alphas_cumprod) - 1, 0, steps).long().cuda() for i in range(len(ts)): t ts[i].repeat(b) pred_eps model(x, timestep_embedding(t, 128), mask) a_bar alphas_cumprod[t].view(-1, 1, 1, 1) x0_pred (x - (1 - a_bar).sqrt() * pred_eps) / a_bar.sqrt() x0_pred x0_pred.clamp(-1, 1) if i len(ts) - 1: # 还没到最后一步回加噪声 a_next alphas_cumprod[ts[i 1]].view(-1, 1, 1, 1) x a_next.sqrt() * x0_pred (1 - a_next).sqrt() * pred_eps else: x x0_pred x x * mask x_known * (1 - mask) # 硬约束已知区域强制还原 return x逻辑说明DDIM 用非马尔可夫推断把 1000 步压到 50 步左右代价是随机性降低、多样性略减每步之后重贴已知区域保证输出和原图在已知部分完全一致。参数说明steps低于 20 时结构容易崩高于 200 时收益递减x0_pred.clamp(-1,1)防止数值外溢导致下一步输入失真。4. 实验设计与参数调优让修复结果从“能看”到“能用”4.1 数据集、掩码生成与评价指标数据集典型用途掩码类型常用指标CelebA人脸修复中心矩形、随机笔刷PSNR、SSIM、LPIPSPlaces2通用场景不规则自由形状FID、LPIPS自有老照片划痕、水渍细长条、块状PSNR、人眼评分指标要看组合。PSNR 对模糊不敏感一块灰色填进去也可能拿高分LPIPS 和 FID 更贴近人眼但需要足够样本量才稳定。掩码生成不要只用中心矩形那会让模型学出“只能修中间”的偏置训练时混合多种掩码形状验证时也要分类型汇报否则实验结论站不住。4.2 关键超参数表T、β 调度、引导强度、采样步数参数常用取值影响训练总步 T1000太小则前向破坏不足太大则训练慢β 调度余弦 / 线性余弦在细节还原上更稳采样步数50200小于 20 易崩大于 200 收益递减引导强度 w1.07.5越大越贴条件过大会出现色斑学习率1e-42e-4配合 warmup 更稳批大小416受显存限制小 batch 用 GroupNorm引导强度是最值得单独调的参数。分类器无关引导在修复里表现为w 偏小缺失区域更自由但可能偏离上下文w 偏大结果更服从已知区域的约束但容易出现饱和度异常和块状伪影。我的习惯是先固定 w1.0 跑通全流程再以 0.5 为步长向上扫到 5.0看验证集 LPIPS 的最低点。4.3 消融实验怎么设计才算数消融要一次只动一个变量并且固定随机种子。下面是一个可以落地的配置脚本。EXPERIMENTS { full: dict(use_mask_condTrue, repaintTrue, samples100), no_mask_cond: dict(use_mask_condFalse, repaintTrue, samples100), no_repaint: dict(use_mask_condTrue, repaintFalse, samples100), fewer_steps: dict(use_mask_condTrue, repaintTrue, samples25), } # 所有实验共用同一份掩码、同一个初始噪声种子否则差异无法归因逻辑说明no_mask_cond用来验证掩码拼接是否真的在起作用很多时候指标差异会比你预期的大no_repaint用来量化重贴操作带来的增益这个增益通常在结构一致性指标上最明显fewer_steps用来确认步数压缩是否已经在牺牲质量。参数说明samples保持一致样本量太小会让指标波动盖过真实差异。4.4 常见失败模式与排查路径输出整体发灰、像蒙了一层雾多半是采样末端x0_pred没有裁剪或者 β 调度在线性下末端过大。先检查 clamp再换余弦调度。缺失区域边缘有明显接缝重贴步骤里掩码没有做羽化硬边界把两侧的噪声分布切开。把掩码边缘做 35 像素的模糊过渡。训练 loss 下降但采样图是噪声时间步嵌入没有广播到空间维或者alphas_cumprod的设备不一致。修复结果和上下文语义不搭引导强度偏低或者掩码拼接时把已知区域也置零了导致网络丢失上下文。排查前先确认张量设备一致这一步能省掉大量无用调试。python -c import torch from model import MaskedUNet m MaskedUNet().cuda().eval() x torch.randn(2,3,64,64).cuda(); mask torch.ones(2,1,64,64).cuda() t torch.zeros(2, dtypetorch.long).cuda() print(m(x, t, mask).shape) 逻辑说明这段脚本用固定输入验证前后向形状能在训练前暴露通道数、设备、时间步类型这几类高频错误。5. 进阶技巧潜空间压缩与重跳采样5.1 潜空间压缩带来的显存与速度收益把扩散过程搬到 VAE 潜空间等于在 1/8 分辨率上做去噪。512×512 的图在潜空间里只有 64×64注意力层的显存占用按平方级下降单步采样成本通常能降一个数量级。代价是 VAE 的重建误差会形成一个精度天花板如果 VAE 本身在细纹理上有损失扩散过程再怎么优化也补不回来。判断是否值得切换可以先单独跑一遍 VAE 的重建看 PSNR 上限落在哪里再对比像素空间模型的实际表现。5.2 RePaint 式重跳采样在贴合与自然之间反复协商标准采样里已知区域只在每步末尾被贴回一次。RePaint 的做法是在反向过程中引入“跳回”——先按正常流程去噪到某一步再人为加噪回到更早的时间步重复若干次。这样缺失区域会多次经过从粗到细的协商纹理连贯性明显改善尤其适合大面积、长条状缺口。策略重跳次数适用缺口采样耗时标准 DDIM0小面积、规则缺口1×轻量重跳35中等块状缺口23×强重跳10 以上大面积、跨结构缺口5× 以上实现上只需在采样循环里加一个外层重复把t序列按resample次数回退。注意重跳次数和引导强度会互相影响两者同时调高容易出现过度锐化的边缘。5.3 验证语义可信度的几个动作指标之外建议固定做三件事把修复图和原图并排看轮廓延续性尤其是建筑边缘、发丝、文字用另一个分类或检测模型跑一遍修复后的区域看语义标签是否和上下文一致把同一张图用不同随机种子采样 5 次如果结果差异巨大说明引导强度不足或模型欠拟合。这三个动作不需要额外训练几分钟就能筛掉大部分“指标好看但不可用”的结果。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进