ARDM扩散模型图像修复:注意力-残差耦合与掩码条件注入
简介面向图像修复与扩散模型方向的研究人员和技术开发者这份文档围绕基于U-Net的扩散修复方案展开覆盖环境配置、数据预处理、模型定义与训练评估全流程。内容结合PyTorch代码解释重点讨论注意力建模与残差连接的多模块融合并提出Attention-Residual Diffusion ModelARDM在多个baseline数据集上完成对比实验、消融实验及指标评估适合作为新方法探索、论文复现和性能优化的参考。资源包仅1个docx文件约18KB以文字说明与代码片段为主便于按实验环节查阅。目前已有245人学习下载。读者可从中获取完整的实验框架、创新模块设计思路、训练与评估流程以及指标提升原因分析用于理解扩散模型在图像修复中的潜力并迁移到自身课题尤其适合具备一定深度学习基础、希望复现中文核心实验的读者。1. 扩散模型做图像修复ARDM 为什么值得单独拆一遍图像修复image inpainting这块GAN 流派做了很多年PSNR 也刷得不低但掩码面积一旦超过三成、纹理又开始重复判别器压不住的模式坍塌就会以块状伪影的形式冒出来。扩散模型走的是另一条路先把整张图加噪到接近纯噪声再让 U-Net 学会逐步去噪修复只是把条件从文本换成「掩码加已知像素」换句话说它是在把缺失区域重新拉回真实数据的流形附近而不是硬猜像素值。这份实验框架里的 ARDMAttention-Residual Diffusion Model就是在这个基础上把注意力建模与残差连接在 U-Net 瓶颈处做耦合再用对比实验和消融实验把提升归因清楚。手上有单卡、想跑通一套完整中文核心实验的读者可以照着往下拆。2. 前向加噪与 U-Net 反向去噪的工程化落地扩散模型的原理说起来只有两行公式真正卡人的是调度器、时间步嵌入和训练目标这三件事怎么落到 PyTorch 里。原框架给的UNet类只有两层卷积加 ReLU连时间步都没接进去那样训出来的网络根本不是扩散模型只是个自编码器。这一章把缺的部分补上。2.1 β 调度线性与余弦的差别在哪β_t决定每一步加多少噪声也决定反向采样时模型要覆盖的噪声区间。线性调度在 T1000 时后期α_bar衰减过快导致高时间步几乎全是纯噪声模型学不到东西余弦调度把中间段拉长重建质量通常更稳。调度方式表达式适用场景常见问题Linearβ从 1e-4 线性到 2e-2低分辨率、步数 ≤ 500末端信噪比过低细节丢Cosinecos²((t/Ts)/(1s)·π/2)64×64 到 256×256 修复前期步进慢需更多 epochScaled-linear线性 β 起点缩到 1e-5潜空间扩散需配合潜变量方差缩放下面这段是把两种调度统一封装成alpha_bar让前向加噪和反向采样共用同一份系数表避免手写错下标。import math import torch def make_beta_schedule(T1000, schedulecosine, s0.008): 返回 (betas, alpha_bar)alpha_bar[t] prod(1 - beta_0..t) if schedule linear: betas torch.linspace(1e-4, 2e-2, T) elif schedule cosine: steps torch.linspace(0, T, T 1) f torch.cos((steps / T s) / (1 s) * math.pi / 2) ** 2 alpha_bar f / f[0] # 由 alpha_bar 反推 beta并夹紧避免除零 betas torch.clip(1 - alpha_bar[1:] / alpha_bar[:-1], 1e-5, 0.999) else: raise ValueError(funknown schedule: {schedule}) alpha_bar torch.cumprod(1.0 - betas, dim0) return betas, alpha_bar def q_sample(x0, t, alpha_bar, noiseNone): 封闭形式加噪: x_t sqrt(a_bar) * x0 sqrt(1 - a_bar) * eps noise torch.randn_like(x0) if noise is None else noise a alpha_bar[t].view(-1, 1, 1, 1) return a.sqrt() * x0 (1 - a).sqrt() * noise, noiset传进来是形状[B]的长整型张量必须先view成[B,1,1,1]才能广播到[B,C,H,W]noise参数留出外部注入的口子做消融实验时固定随机种子复现同一条加噪轨迹会用到。alpha_bar建议在训练一开始就to(device)缓存好每步重算会明显拖慢吞吐。2.2 U-Net 主干必须补齐的三个组件很多复现跑不出效果的根因就在主干上。一是时间步嵌入用正弦位置编码加两层 MLP再通过 AdaGN 或逐层相加注入每个残差块否则同一个网络要对所有噪声等级给出一致预测必然欠拟合。二是下采样与上采样对称64×64 输入通常做 3 次下采样到 8×8 瓶颈通道数按 64→128→256 递增。三是低分辨率自注意力只在 16×16 和 8×8 两级加全局注意力高分辨率加会直接吃满显存。训练目标本身也有讲究原框架用MSELoss(output, data)让网络重建输入图这是错的。扩散模型预测的是噪声eps或v α·eps - σ·x0x0由预测噪声反算。这套 ARDM 用 eps-prediction因为它在中等噪声区间梯度更平稳。2.3 掩码感知的训练循环修复任务的关键在于损失只应主要作用于缺失区域但完全不管已知区域又会让边界出现接缝。所以这里用「缺失区主损失 已知区辅助损失」的加权形式。def train_step(model, x0, mask, alpha_bar, device, lambda_valid0.1): mask: 1 表示缺失待修复, 0 表示已知 model.train() x0, mask x0.to(device), mask.to(device) t torch.randint(0, alpha_bar.shape[0], (x0.size(0),), devicedevice) x_t, noise q_sample(x0, t, alpha_bar) # 条件输入 加噪图 掩码 已知区域像素缺失处置 0 cond torch.cat([x_t, mask, x0 * (1 - mask)], dim1) pred model(cond, t) se (pred - noise) ** 2 loss_missing (se * mask).sum() / (mask.sum() 1e-8) loss_valid (se * (1 - mask)).sum() / ((1 - mask).sum() 1e-8) return loss_missing lambda_valid * loss_validlambda_valid取 0.1 是经验值调大到 0.3 以上模型会倾向于直接复制已知区域缺失区变糊调到 0 则边界接缝明显。mask.sum()上加1e-8是防止某个 batch 恰好没有缺失像素时出现 NaN。反向传播时记得optimizer.zero_grad()放在loss.backward()之前原框架把它写在了前向之前逻辑上没错但容易在梯度累积场景里踩坑。3. ARDM 的注意力-残差耦合与掩码条件注入命名叫 ARDM 不是把两个模块并排放进去就算创新。审稿人最常问的一句是「这和加个 SE 块有什么区别」所以耦合方式必须能说清楚。3.1 为什么不是简单 AB耦合点选在哪里普通做法是残差块后面接一个注意力块两者串行。串行的问题是注意力拿到的特征是残差输出而残差输出已经被恒等映射拉回了原分布注意力学到的权重会退化成近似均匀分布。ARDM 改成注意力作用于残差分支内部先算出h conv(norm(x))再用通道注意力重标定h最后整体缩放后加回输入。这样注意力的梯度只影响新学到的残差项不会污染主干恒等通路早期训练更稳。另一处是门控系数gamma零初始化。扩散模型早期噪声等级高残差分支输出方差大如果直接相加会让 loss 在前 2k 步剧烈震荡把gamma初始化为 0等价于模型一开始就是恒等映射注意力分支从零开始慢慢长出来。3.2 AttentionResidualBlock 的实现与超参class AttentionResidualBlock(nn.Module): def __init__(self, in_ch, reduction16): super().__init__() # GroupNorm 对小 batch 更友好且不受噪声等级影响 self.norm nn.GroupNorm(8, in_ch) self.conv nn.Sequential( nn.Conv2d(in_ch, in_ch, 3, padding1), nn.SiLU(inplaceTrue), nn.Conv2d(in_ch, in_ch, 3, padding1), ) hidden max(in_ch // reduction, 4) # 低通道层保护防止压缩到 0 self.channel_attn nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(in_ch, hidden, 1), nn.SiLU(inplaceTrue), nn.Conv2d(hidden, in_ch, 1), nn.Sigmoid(), ) self.gamma nn.Parameter(torch.zeros(1)) # 零初始化从恒等映射起步 def forward(self, x): h self.conv(self.norm(x)) h h * self.channel_attn(h) # 注意力在残差分支内部生效 return x self.gamma * hreduction16是通道压缩比64 通道层会压到 4 维如果主干第一层只有 32 通道max(..., 4)能兜住。GroupNorm(8, in_ch)要求通道数能被 8 整除主干通道按 64/128/256 设计就天然满足。训练脚本里建议把gamma单独加进日志它的绝对值能直接反映注意力分支是否真的被激活。3.3 掩码条件注入的三种方式注入方式输入通道优点代价通道拼接3137实现最简单兼容任意主干第一层卷积参数增加部分卷积31参数量小边界过渡自然每层都要维护 mask 更新掩码做门控31显式控制信息流需额外超参调门控温度这套实现选的是通道拼接把加噪图、掩码、已知区域像素拼成 7 通道喂给第一层卷积。选它的理由是消融实验好做只要把拼接通道数改掉就能控制变量不用动主干结构。原框架把AttentionResidualBlock(3)直接接在 3 通道输入上等于掩码信息完全没进去训练出来的是无条件生成模型这一点在复现时务必改掉。4. 对比实验与消融实验的完整流水线实验能不能过审一半看模型一半看对照组和指标写得对不对。4.1 数据集与掩码生成策略MNIST 用来验证流程通不通CIFAR-10 用来验证彩色纹理上的泛化。两者都Resize((64,64))后归一化到[-1,1]归一化参数从(0.5,0.5,0.5)改成按数据集统计量算更稳。掩码不能只用中心方块那样模型会学到「只看边界」的捷径建议按 60% 自由笔画 30% 中等矩形 10% 大面积块状混合。import random import numpy as np def random_free_form_mask(h64, w64, n_strokes8, brush6): 自由笔画掩码返回 [h,w] 的 float321 表示缺失 mask np.zeros((h, w), dtypenp.float32) for _ in range(n_strokes): x, y random.randint(0, w - 1), random.randint(0, h - 1) for _ in range(random.randint(10, 25)): x int(np.clip(x random.randint(-brush, brush), 0, w - 1)) y int(np.clip(y random.randint(-brush, brush), 0, h - 1)) mask[max(0, y - brush):y brush, max(0, x - brush):x brush] 1.0 return mask def build_mask_batch(b, h, w, p_free0.6): 按比例混合三种掩码形态返回 [b,1,h,w] 张量 masks [] for _ in range(b): r random.random() if r p_free: m random_free_form_mask(h, w) elif r p_free 0.3: m np.zeros((h, w), np.float32) m[h // 4:h // 4 h // 2, w // 4:w // 4 w // 2] 1.0 else: m np.zeros((h, w), np.float32) m[: h // 3, :] 1.0 masks.append(m) return np.stack(masks)[:, None]掩码随机种子要和验证集分开固定否则每次评估用的破坏形态不同指标波动会盖过模型差异。训练时每张图每个 epoch 重新采样掩码相当于做了数据增广。4.2 与 GAN baseline 的对齐对比实验最容易翻车的地方是配置不对齐学习率、batch size、训练轮数任意一项不同审稿人就有理由质疑结论。原框架里 GAN 用MSELoss训练生成器缺了判别器对抗损失和感知损失这不是 GAN 而是纯回归。规范做法是生成器用L1 0.1·对抗损失 0.1·感知损失判别器用 hinge loss两边交替更新。对齐清单同一组掩码分布、同一 batch size32、同一优化器Adamlr1e-4betas(0.5,0.999)、同一评估频率。扩散模型因为要采样才能评估建议每 5 个 epoch 评一次采样步数固定 50别用不同步数去比。4.3 PSNR/SSIM 的正确计算姿势原框架的评估函数有三处问题。第一model(data)输入是完整图输出也当完整图去算过程中没有掩码这算的是自编码重建不是修复。第二数据和输出都在[-1,1]值域skimage对 float 默认data_range1.0算出来的 PSNR 会虚高。第三新版skimage已经把multichannel参数改成channel_axis老写法直接报错。import numpy as np import torch from skimage.metrics import peak_signal_noise_ratio as psnr from skimage.metrics import structural_similarity as ssim torch.no_grad() def evaluate(model, sampler, loader, device, mask_fn): sampler 是反向采样函数签名 sampler(model, cond) model.eval() psnrs, ssims [], [] for x0, _ in loader: x0 x0.to(device) b, _, h, w x0.shape mask torch.from_numpy(mask_fn(b, h, w)).to(device) # [B,1,H,W] cond torch.cat([x0 * (1 - mask), mask, x0 * (1 - mask)], dim1) recon sampler(model, cond) # 值域 [-1,1] gt ((x0 1) / 2).clamp(0, 1).cpu().permute(0, 2, 3, 1).numpy() pd ((recon 1) / 2).clamp(0, 1).cpu().permute(0, 2, 3, 1).numpy() for i in range(b): psnrs.append(psnr(gt[i], pd[i], data_range1.0)) ssims.append(ssim(gt[i], pd[i], channel_axis2, data_range1.0)) return float(np.mean(psnrs)), float(np.mean(ssims))指标建议报两组全图指标和缺失区指标。全图指标好看多半是因为已知区域占了大头缺失区指标才反映真实修复能力两个数放一起才说明问题。下表是这套配置在 64×64、50% 掩码下的量级参考用来判断实现有没有跑偏实际数字以你自己跑出来的为准。模型全图 PSNR全图 SSIM缺失区 PSNR缺失区 SSIMGAN baseline24.80.81219.40.703U-Net 自编码25.60.83620.10.726ARDM完整27.30.87622.70.7944.4 消融矩阵怎么排消融要能回答「哪个模块贡献了多少」所以至少四组完整 ARDM、去掉注意力只留残差、去掉残差注意力直接接主干、以及把gamma改回随机初始化。第四组常被忽略但它恰好证明零初始化门控是有效的而不是摆设。每组跑三个随机种子取均值方差单次结果差异小于 0.2 dB 时不要下结论。5. 潜空间加速、重采样一致性与排错清单5.1 潜空间扩散把采样成本压下来64×64 上跑 1000 步 DDPM 采样单张图要几秒做完整测试集评估会等到崩溃。潜在扩散模型Latent Diffusion的做法是先训一个 VAE把图像压到 8 倍下采样的潜空间扩散过程在潜空间里进行采样成本直接降一个数量级。修复场景下要注意掩码也要同步下采样成潜空间分辨率已知区域的约束则在解码后回到像素域施加否则潜变量里的边界会糊掉。# 潜空间条件构造z 为潜变量mask_lat 由 mask 平均池化 8 倍得到 z vae.encode(x0).latent_dist.sample() * 0.18215 mask_lat F.avg_pool2d(mask, kernel_size8, stride8) cond torch.cat([z * (1 - mask_lat), mask_lat, z * (1 - mask_lat)], dim1)5.2 DDIM 与 RePaint 式重采样把采样步数从 1000 降到 50 用 DDIMeta0时过程完全确定同一张输入每次输出一致评估可复现。但纯 DDIM 有个副作用已知区域也会被重新生成出现色偏。RePaint 的思路是每采样若干步把已知区域替换成「加噪到当前时间步的真实像素」这样已知区域始终锚定在原图附近。torch.no_grad() def ddim_repaint(model, cond, alpha_bar, steps50, jump10, devicecuda): T alpha_bar.shape[0] ts torch.linspace(T - 1, 0, steps, devicedevice).long() x torch.randn_like(cond[:, :3]) for i, t in enumerate(ts): eps model(torch.cat([x, cond[:, 3:]], dim1), t.expand(x.size(0))) a_t alpha_bar[t].view(-1, 1, 1, 1) x0_pred ((x - (1 - a_t).sqrt() * eps) / a_t.sqrt()).clamp(-1, 1) a_prev alpha_bar[ts[i 1]].view(-1, 1, 1, 1) if i 1 steps else torch.ones_like(a_t) x a_prev.sqrt() * x0_pred (1 - a_prev).sqrt() * eps # eta0 # 每 jump 步把已知区域重锚定一次 if i % jump 0: known cond[:, 3:4] a_t2 alpha_bar[t].view(-1, 1, 1, 1) noise torch.randn_like(x) x_known a_t2.sqrt() * cond[:, :3] (1 - a_t2).sqrt() * noise x x * known x_known * (1 - known) return xsteps50、jump10是这套配置下的平衡点jump调小如 5一致性更强但纹理多样性下降调大如 20则重新出现色偏。eta想加随机性就设成 0.2 左右但评估时统一用 0 保证可复现。5.3 排错清单现象常见根因定位手段loss 前 1k 步剧烈震荡gamma未零初始化 / 学习率过大打印gamma绝对值降到 5e-5缺失区全灰无纹理eps 预测但损失用了 x0 重建检查train_step的 target 是否为noise边界出现明显接缝只算缺失区损失lambda_valid0提到 0.1 并加 1 像素膨胀的软掩码PSNR 异常高35 dB值域没对齐或没加掩码打印输入输出 min/max确认在 [-1,1]采样后整图色偏已有区域被重新生成上 RePaint 重锚定显存 OOM 在 16×16 注意力层高分辨率层加了全局注意力只保留 8×8/16×16 两级最后补一个验证技巧把训练集里一张图的已知区域和另一张图的缺失区域拼起来当输入看模型是复制原图还是生成新内容。如果输出几乎是拷贝已知区域说明lambda_valid太大或掩码监督漏了如果输出结构合理但纹理重复那是注意力分支没被激活回去看gamma的日志曲线。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →