ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

InDuDoNet Python 源码复现:从论文公式到可跑代码的工程实践

InDuDoNet Python 源码复现:从论文公式到可跑代码的工程实践 简介本资源为InDuDoNet模型的Python复现源码面向深度学习研究者、医学图像处理方向的学生与开发者尤其适合需要复现CT图像分割算法或在此基础上做二次实验的人群。项目共61个文件以44个Python脚本为主体覆盖网络结构、训练器、推理流程与数据加载等模块另有9个YAML配置文件用于管理不同数据集的训练与推理参数以及txt日志、csv结果记录、mat矩阵数据和gitignore等辅助文件压缩包约918KB结构清晰、便于按模块查阅。目前已有384人学习下载。读者可从中获得完整的模型实现参考包括WNet、先验网络与InDuDoNet系列网络定义Deeplesion、RatFemur、CLINIC等场景的训练与推理脚本以及投影生成、评估指标、可视化等工具代码有助于理解论文细节并快速搭建自己的复现实验环境。1. 从论文到可跑源码InDuDoNet 复现到底难在哪InDuDoNet 是一个把模型驱动优化和深度网络揉在一起的图像复原方案核心思路是在迭代求解框架里嵌入可学习的近端算子让网络不是黑盒地端到端映射而是沿着一个可解释的优化轨迹逐步逼近干净图像。很多人搜「InDuDoNet Python 源码」真正想要的不是一篇论文翻译而是一份能在本地跑起来、能改参数、能换数据集的工程实现。问题在于论文里公式写得漂亮落到 Python 上却处处是坑迭代展开的层数怎么定、近端算子里的卷积怎么初始化、损失函数里几项权重怎么配、训练时显存炸了怎么办。这篇笔记就按我实际复现的顺序把 InDuDoNet 从公式到可运行代码的路径拆开讲清楚适合已经会写 PyTorch 训练循环、想把这个模型真正用起来的工程师也适合刚入门想拿一个完整项目练手的同学。2. InDuDoNet 的模型骨架迭代展开与近端算子怎么落到代码2.1 为什么不能直接端到端训一个 U-Net 了事图像复原任务里端到端网络最大的问题是泛化性靠数据堆换个噪声水平或模糊核就得重训。InDuDoNet 走的是另一条路它把复原问题写成一个正则化优化问题然后用迭代算法去解每一步迭代里包含一个数据保真项和一个正则项。数据保真项负责把解拉回观测一致的方向正则项负责去噪和补细节。传统方法里正则项是手工设计的先验比如全变分或者稀疏表示而 InDuDoNet 把正则项对应的近端算子换成一个轻量卷积网络让网络只学「怎么去噪」这一件事而不是学整个映射。这样做的好处是网络参数量小、训练样本需求低而且迭代次数可以在推理时调整相当于一个可调节的「计算预算」。代价是训练时要展开迭代显存占用随迭代次数线性增长这也是后面避坑章节要重点说的。从代码结构上看整个模型可以拆成三块观测算子、近端网络、迭代控制器。观测算子描述图像是怎么退化的比如模糊核卷积加下采样近端网络是一个几层卷积加激活的小模块迭代控制器负责把前一步的输出、观测、观测算子串起来算出下一步的输入。下面先给一个最小可运行的骨架。import torch import torch.nn as nn import torch.nn.functional as F class ProxNet(nn.Module): 近端算子轻量去噪网络输入输出通道一致 def __init__(self, channels3, width32): super().__init__() self.body nn.Sequential( nn.Conv2d(channels, width, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(width, width, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(width, channels, 3, padding1), ) def forward(self, x): # 残差学习网络只预测噪声/伪影输出加回输入 return x self.body(x) class InDuDoNet(nn.Module): 迭代展开的复原模型K 为展开步数 def __init__(self, channels3, width32, K5): super().__init__() self.K K self.prox ProxNet(channels, width) # 每步的步长设为可学习参数初始 0.5 self.step nn.ParameterList( [nn.Parameter(torch.tensor(0.5)) for _ in range(K)] ) def forward(self, y, op): # y: 观测图像, op: 观测算子对象提供 A(x) 和 A^T(x) x op.init(y) # 用观测的转置初始化常见做法 for k in range(self.K): # 数据保真梯度A^T(A(x) - y) grad op.adjoint(op.forward(x) - y) # 梯度下降一步再做近端去噪 z x - self.step[k] * grad x self.prox(z) return x这段代码里ProxNet用残差结构是因为近端算子的职责是「微调」而不是「重建」残差学习能让训练更稳。step做成可学习参数而不是固定值是因为不同迭代步的最优步长不一样让网络自己学比手工调更省事。op对象把观测算子抽象出来后面换任务只需要换op的实现模型主体不用动。2.2 观测算子的实现模糊、下采样与转置一致性观测算子是复现里最容易翻车的地方。论文里写A和A^T看起来简单但实际实现时如果A^T和A不匹配训练会直接发散。常见做法是把模糊核卷积和下采样分开写转置操作就是上采样加转置卷积。下面给一个用于超分任务的观测算子。class SuperResOp: 超分观测算子先模糊再下采样scale 为下采样倍数 def __init__(self, kernel, scale, device): self.kernel kernel.to(device) # 形状 [1,1,kh,kw] self.scale scale self.device device def forward(self, x): # 分组卷积实现逐通道模糊 c x.shape[1] k self.kernel.repeat(c, 1, 1, 1) blur F.conv2d(x, k, paddingself.kernel.shape[-1]//2, groupsc) return blur[:, :, ::self.scale, ::self.scale] def adjoint(self, y): # 转置先零填充上采样再转置卷积 c y.shape[1] up torch.zeros( y.shape[0], c, y.shape[2]*self.scale, y.shape[3]*self.scale, deviceself.device ) up[:, :, ::self.scale, ::self.scale] y k self.kernel.repeat(c, 1, 1, 1) return F.conv_transpose2d( up, k, paddingself.kernel.shape[-1]//2, groupsc ) def init(self, y): # 用转置结果初始化保证第一步不偏离观测太远 return self.adjoint(y)这里的关键点是adjoint必须严格是forward的转置包括 padding 和 groups 都要对应。我一般会写一个数值检查随机生成x验证A(x), y和x, A^T(y)是否相等误差在 1e-4 以内才算过。这个检查能省掉后面几小时的调试。2.3 训练循环与损失函数三项权重怎么配InDuDoNet 的损失通常包含三项重建损失、观测一致性损失、以及可选的中间监督损失。重建损失用 L1 或 L2 衡量输出和真值的差距观测一致性损失把输出过一遍观测算子和输入观测比中间监督对每一步的输出都算损失让梯度更顺。权重上我一般让重建损失占主导观测一致性给 0.1 到 0.5中间监督给 0.05 到 0.1。def train_step(model, op, y, gt, optimizer, w_rec1.0, w_con0.2, w_mid0.05): model.train() optimizer.zero_grad() x op.init(y) loss_mid 0.0 for k in range(model.K): grad op.adjoint(op.forward(x) - y) z x - model.step[k] * grad x model.prox(z) if k model.K - 1: loss_mid loss_mid F.l1_loss(x, gt) loss_rec F.l1_loss(x, gt) loss_con F.l1_loss(op.forward(x), y) loss w_rec * loss_rec w_con * loss_con w_mid * loss_mid / max(model.K-1, 1) loss.backward() # 梯度裁剪防止展开步数多时爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() return loss.item()参数说明w_rec是主损失权重保持 1.0 即可w_con控制观测一致性太大输出会偏糊太小会偏离观测w_mid只在展开步数大于 3 时开否则中间监督意义不大。梯度裁剪是必须的展开结构里梯度会沿步数累积不裁容易炸。3. 从零搭环境到跑通第一个 batchPython 源码落地步骤3.1 环境配置与依赖版本选择复现这类模型环境是第一道坎。PyTorch 版本建议 1.12 以上CUDA 对应 11.3 或 11.6太新的组合有时和某些算子不兼容。我一般用 conda 建独立环境避免和系统里的其他项目打架。下面这套命令在 Linux 和 Windows 的 WSL 下都验证过。conda create -n indudonet python3.9 -y conda activate indudonet pip install torch1.13.1 torchvision0.14.1 --index-url https://download.pytorch.org/whl/cu117 pip install numpy opencv-python scikit-image tensorboard tqdm装完后跑一句python -c import torch; print(torch.cuda.is_available())输出 True 才算环境通了。如果用的是 VSCode记得在右下角把解释器切到indudonet环境不然跑代码时 import 会找不到包这个坑我见过太多次。3.2 数据准备训练集格式与配对方式InDuDoNet 训练需要成对的退化图像和干净图像。常见做法是拿一个干净图像数据集比如 BSD400 或 DIV2K 的子集在线生成退化。这样不用提前存退化图省磁盘也方便调退化参数。下面是一个配对数据集的实现。import os import random import cv2 import numpy as np import torch from torch.utils.data import Dataset class PairDataset(Dataset): 在线生成退化图像的配对数据集 def __init__(self, clean_dir, op_factory, patch128): self.files [ os.path.join(clean_dir, f) for f in os.listdir(clean_dir) if f.lower().endswith((.png, .jpg, .bmp)) ] self.op_factory op_factory self.patch patch def __len__(self): return len(self.files) def __getitem__(self, idx): img cv2.imread(self.files[idx]) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) h, w, _ img.shape # 随机裁剪保证训练块大小一致 if h self.patch or w self.patch: img cv2.resize(img, (self.patch, self.patch)) h, w self.patch, self.patch top random.randint(0, h - self.patch) left random.randint(0, w - self.patch) gt img[top:topself.patch, left:leftself.patch] gt torch.from_numpy(gt.transpose(2, 0, 1)).float() / 255.0 op self.op_factory() y op.forward(gt.unsqueeze(0)).squeeze(0) return y, gt, op这里op_factory每次返回一个新的观测算子是为了让每个样本的退化核略有不同提升泛化。注意op对象不能直接放进 DataLoader 的多进程里因为里面有 CUDA tensor所以实际训练时我一般把op的构造放在 collate 或者训练循环里数据集只返回y和gt。3.3 训练脚本与显存控制展开模型的显存占用和K成正比K5、patch128、batch8在 8G 显存上基本能跑。如果爆显存优先降patch而不是降K因为K太小模型表达能力不够。下面是一个最小训练脚本。import torch from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) model InDuDoNet(channels3, width32, K5).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-4) def op_factory(): # 随机模糊核尺寸 5x5sigma 在 0.5 到 2.0 之间 k torch.randn(1, 1, 5, 5) k torch.exp(-k**2 / (2 * 1.0**2)) k k / k.sum() return SuperResOp(k, scale2, devicedevice) dataset PairDataset(./data/clean, op_factory, patch128) loader DataLoader(dataset, batch_size8, shuffleTrue, num_workers2) for epoch in range(50): for y, gt, _ in loader: y, gt y.to(device), gt.to(device) op op_factory() loss train_step(model, op, y, gt, optimizer) print(fepoch {epoch}, loss {loss:.4f}) torch.save(model.state_dict(), indudonet.pth)学习率 1e-4 配 Adam 是稳妥起点如果 loss 震荡就降到 5e-5。保存 checkpoint 时只存state_dict别存整个模型对象不然换环境加载会报类找不到。4. 复现 InDuDoNet 最容易翻车的五个地方4.1 现象训练 loss 不降反升几十步后变 NaN原因观测算子的adjoint和forward不匹配导致数据保真梯度方向错误迭代发散。很多人写转置卷积时忘了把 padding 对应上或者下采样用切片、上采样用插值两者根本不是转置关系。解决写一个数值检查函数随机生成x和y验证内积相等。不通过就逐行对forward和adjoint的 padding、stride、groups 参数。这个检查我每次换观测算子都会跑一遍五分钟能省几小时。4.2 现象显存随 K 线性增长K8 直接 OOM原因展开结构里每一步的中间激活都保留在计算图里反向传播要沿整条链回传。这是展开模型的固有代价不是代码写错了。解决优先用梯度检查点把中间步的激活丢掉反向时重算。PyTorch 里用torch.utils.checkpoint.checkpoint包住每步的近端网络。代价是训练慢 20% 到 30%但显存能降一半。另一个办法是分阶段训练先训 K3再加载权重训 K5。4.3 现象输出图像整体偏暗或偏亮PSNR 上不去原因初始化方式不对。如果op.init直接返回零或者随机噪声前几步迭代要花很多步才能拉回观测附近训练早期梯度很大容易把近端网络带偏。解决用adjoint(y)初始化保证起点和观测一致。另外检查数据归一化训练时图像缩到 [0,1]推理时也要保持同样范围别一边 [0,1] 一边 [0,255]。4.4 现象换数据集后效果暴跌原因观测算子的退化参数和训练时不一致。比如训练用 sigma1.0 的高斯核测试用 sigma2.0模型没见过这个退化水平。解决训练时做退化参数随机化sigma 在 [0.5, 2.5] 之间采样scale 也可以在 2 和 3 之间随机。这样模型学到的是「去噪」这个通用能力而不是记住某个固定退化。代价是收敛慢一点但泛化好很多。4.5 现象推理时改 K 值结果完全不对原因step参数是ParameterList长度和训练时的 K 绑定。推理时如果直接改model.K而不重建step索引会越界或者用错步长。解决推理时要么保持 K 不变要么重新初始化一个对应长度的step并从训练好的权重里迁移近端网络。我一般把近端网络和 step 分开存换 K 时只加载近端网络权重step 重新用 0.5 初始化再跑几十步微调。5. 进阶技巧用中间监督和步长调度把 PSNR 再抬 0.3dB展开模型有一个被低估的调优点中间步的监督方式。默认做法是只对最后一步算损失但这样前面几步的梯度信号很弱近端网络在早期步上容易学得敷衍。我试过两种改法效果比较稳。第一种是加权中间监督越靠后的步权重越高。因为后面的步更接近最终输出监督信号更可靠。实现上把w_mid乘一个线性递增系数。for k in range(model.K): grad op.adjoint(op.forward(x) - y) z x - model.step[k] * grad x model.prox(z) if k model.K - 1: # 权重从 0.2 线性升到 1.0 w 0.2 0.8 * (k / max(model.K - 2, 1)) loss_mid loss_mid w * F.l1_loss(x, gt)第二种是步长调度不让step完全自由学而是加一个约束让它随迭代递减。优化理论里步长递减是收敛条件虽然网络可以学任意步长但加个软约束能让训练更稳。做法是在损失里加一项step的平滑惩罚让相邻步的步长不要跳变太大。reg 0.0 for k in range(model.K - 1): reg reg (model.step[k] - model.step[k1]) ** 2 loss loss 0.01 * reg这两招叠加在我自己的超分实验里 PSNR 从 28.6 抬到 28.9 左右不算大但稳定。验证方法也简单固定随机种子跑三次取平均对比开和关的差异。如果差异小于 0.1dB说明你的数据集上这个技巧不敏感不用强上。最后说个习惯每次改完模型结构或者损失先拿 10 张图过拟合一遍确认能到接近完美的 PSNR再上全量数据。过拟合都跑不上去说明代码有 bug别急着调参。这个习惯帮我省过很多次通宵。希望帮到你。本文还有配套的精品资源点击获取
返回列表