简介:图像去噪方向常用的Noise2Void(N2V)模型,现有基于U-Net的完整Pytorch复现代码,适合具备一定Python基础的深度学习初学者,以及需要处理真实图像噪声的研究与工程人员。N2V的核心特点是无需成对干净图像即可训练,利用噪声本身作为监督信号,解决了真实场景中难以获取干净参考图的问题。代码包共20个文件,以6个Python脚本为主线,分别实现U-Net模型搭建、数据集封装、训练主流程、测试调用与指标曲线绘制;同时提供2个pth格式预训练权重,覆盖灰度图和RGB图场景,下载后可直接运行验证。资源目录规划清晰,data、datasets、weights、Plt等文件夹分别对应数据、训练集、权重与可视化结果,便于按需取用。整体仅5.32MB,便于快速获取。已有995人学习使用。配套讲解文章从N2V盲点去噪原理到U-Net实现细节均有清晰注释,还给出训练过程中的Loss、PSNR、SSIM变化曲线绘制方法,并梳理了完整复现思路,帮助使用者理解模型机制,快速上手训练与测试。
1. 从一个没有干净真值的图像去噪问题说起
做图像去噪的人早晚会遇到一个尴尬:手里一堆带噪图,却拿不出成对的干净真值。Noise2Void(N2V)在 2019 年给出了一条不用真值的路,而把它在 Pytorch 里复现,最顺手的骨架就是 U-Net。荧光显微镜、冷冻电镜这类场景拍图贵、活体样本还在动,想单独拍一张干净真值几乎不现实,传统监督去噪的训练集在这里直接断供,这正好是 N2V 的主场。
N2V 的思路有点绕:训练时随机把一部分像素的输入值替换成邻居值,再让网络预测这些位置的原值。网络看不到被遮像素的真实输入,只能靠四周上下文猜,去噪被迫学会。这个盲点思想后来被 LAN 这类噪声自适应方法继承,但 N2V 的工程实现最简洁,适合作为理解盲点去噪的第一站。
下面按"原理 → 网络 → 训练 → 评估 → 排错"的顺序展开,所有代码都可以直接拷进工程跑起来。
2. Blind-Spot 原理与 N2V 掩码策略
2.1 为什么普通自监督会塌缩成恒等映射
把带噪图 x 直接喂给网络,用 MSE 让输出逼近 x,最优解是恒等映射:输出等于输入,loss 归零,但噪声一点没去掉。这不是优化失败,而是目标函数本身就允许作弊。想让它"不得不去噪",就要切断网络读取中心像素真实值的通道,同时仍然要求它输出中心像素的估计值,网络唯一能用的信息只剩周围像素。
N2V 的做法分三步:随机选一部分像素打掩码;把掩码位置的输入值替换成邻域内随机一个像素的值;只在掩码位置计算预测值与原始值的 MSE。掩码位置的原值被替换后,输入里没有"正确答案"可以抄,网络必须从上下文推断。如果噪声逐像素独立、信号在空间上有结构,那么给定邻域时对原始带噪中心的 MSE 最优估计,恰好就是干净信号的条件期望——这正是 MMSE 去噪器要学的东西。
先对比几条常见去噪路线,方便理解 N2V 的位置:
| 方法 | 训练数据要求 | 还原本质 |
|---|---|---|
| 监督去噪(U-Net 成对图) | 噪声/干净成对图 | 直接回归干净图 |
| Noise2Noise | 同一场景两组独立带噪图 | 用独立噪声的期望代替真值 |
| Noise2Void | 单张带噪图 | 盲点掩码,逼网络用邻域推断中心 |
| LAN 等自适应方法 | 单张带噪图 | 显式估计噪声并联合去噪 |
N2V 不需要成对数据,也不需要同一场景拍两次,这是它和 Noise2Noise 最大的区别,代价是多了"噪声逐像素独立"这个假设,后面所有参数都在为这个假设服务。
2.2 掩码生成与邻居替换的实现
掩码逻辑是整条训练管线的核心,建议单独抽成函数而不是塞进训练循环。下面这份实现按"先选掩码、再从原图取邻居值"的顺序处理:
import numpy as np def generate_mask(shape, p=0.1): """随机生成布尔掩码, True 表示该位置要被替换""" return np.random.random(shape) < p def apply_blind_mask(patch, mask, radius=8): """ 把 mask 位置的原值替换为窗口内随机一个邻居的值。 patch: (H, W) float32; mask: (H, W) bool。 """ out = patch.copy() ys, xs = np.nonzero(mask) for y, x in zip(ys, xs): while True: ny = min(max(y + np.random.randint(-radius, radius + 1), 0), patch.shape[0] - 1) nx = min(max(x + np.random.randint(-radius, radius + 1), 0), patch.shape[1] - 1) if (ny, nx) != (y, x): break out[y, x] = patch[ny, nx] # 必须从原始 patch 取 return out这段代码逐像素循环,patch 只有 64×64,掩码数量在几百个量级,Python 循环耗时完全可接受,刻意向量化反而让代码变难读。radius控制采样窗口半径,窗口太小(比如 1)时邻居与中心太像,掩码形同虚设;窗口太大则替换值可能来自结构完全不同的区域,引入额外方差。论文里一般取 4~8,要和感受野配合着看(下一节)。
注意:替换必须从原始 patch 取邻居,不能从已经改过的
out里取,否则两个相邻掩码会互相污染。
掩码比例p同样敏感:调大训练信号变多,但输入中被篡改的像素变多,输入分布偏离真实分布太远;调小每一步能回传梯度的位置太少。我一般从 0.1 起步,观察验证集输出再往 0.05 或 0.15 调。
2.3 掩码半径、感受野与信息泄漏
掩码能防住"直接抄中心值",但防不住所有捷径。卷积是滑动窗口,中心像素的值会通过 3×3 卷积进入周围位置的感受野;如果结构里有捷径,被遮像素的信息可能绕到输出上。N2V 能工作的关键是它只在掩码位置计算 loss——周围预测里即使泄漏了中心值,只要中心位置的预测没吃到真值,就抄不成。
真正要算的是感受野。四层下采样的 U-Net,每层两个 3×3 卷积加一次 2×2 池化,最后有效感受野超过 100×100,覆盖 64×64 的 patch 绰绰有余。只要掩码半径远小于这个数值,被遮位置四周就有大量未篡改像素提供上下文;如果网络很浅、感受野只有 9×9,掩码半径取 8 时整个上下文都被污染,训练必然发散。调结构按这个顺序检查:patch 大小 ≥ 2×感受野,掩码半径 ≤ 感受野的 1/4。
3. 搭建 U-Net 骨架的 Pytorch 实现
3.1 网络结构与通道约定
N2V 对网络主干没有硬性要求,U-Net 成为事实标准是因为特征复用效率高、参数相对少,跳跃连接还能保住边缘细节。这里用 pytorch 基础框架实现单输入单输出的 U-Net:输入输出都是 1 通道灰度图,值域 [0,1] 的 float32。想做 RGB 就把 in_ch/out_ch 改成 3,掩码与替换按通道独立做。
整体是经典的非对称编解码:编码器四层下采样,每层通道翻倍;瓶颈保持 512 通道;解码器用转置卷积上采样并与同层编码特征拼接。每个卷积块是"3×3 卷积 + BatchNorm + ReLU"两次堆叠,卷积核、步长、填充统一为 3/1/1。先想清楚通道流向再写代码,避免后面 concat 对不上维度。
3.2 从 DoubleConv 到完整前向代码
下面这份代码可以直接存成models.py。整体参数量在千万级(base=64 时约 1300 万),8GB 显存跑 batch 16 的 64×64 patch 没有压力。
import torch import torch.nn as nn class DoubleConv(nn.Module): """两次 3x3 卷积, 每次后面跟 BatchNorm 和 ReLU""" def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x) class Down(nn.Module): """最大池化下采样 + DoubleConv""" def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_ch, out_ch), ) def forward(self, x): return self.block(x) class Up(nn.Module): """转置卷积上采样 + 跳跃连接拼接 + DoubleConv""" def __init__(self, up_ch, skip_ch, out_ch): super().__init__() self.up = nn.ConvTranspose2d(up_ch, up_ch // 2, kernel_size=2, stride=2) self.conv = DoubleConv(up_ch // 2 + skip_ch, out_ch) def forward(self, x, skip): x = self.up(x) dy = skip.size(2) - x.size(2) # 尺寸可能差 1, 中心裁剪对齐 dx = skip.size(3) - x.size(3) skip = skip[:, :, dy // 2: dy // 2 + x.size(2), dx // 2: dx // 2 + x.size(3)] return self.conv(torch.cat([skip, x], dim=1)) class N2VUNet(nn.Module): def __init__(self, in_ch=1, out_ch=1, base=64): super().__init__() self.inc = DoubleConv(in_ch, base) # 1 -> 64 self.down1 = Down(base, base * 2) # 64 -> 128 self.down2 = Down(base * 2, base * 4) # 128 -> 256 self.down3 = Down(base * 4, base * 8) # 256 -> 512 self.down4 = Down(base * 8, base * 8) # 512 -> 512 瓶颈不翻倍 self.up1 = Up(base * 8, base * 8, base * 4) # 512+512 -> 256 self.up2 = Up(base * 4, base * 4, base * 2) # 256+256 -> 128 self.up3 = Up(base * 2, base * 2, base) # 128+128 -> 64 self.up4 = Up(base, base, base) # 64+64 -> 64 self.outc = nn.Conv2d(base, out_ch, 1) def forward(self, x): s1 = self.inc(x) s2 = self.down1(s1) s3 = self.down2(s2) s4 = self.down3(s3) x = self.down4(s4) x = self.up1(x, s4) x = self.up2(x, s3) x = self.up3(x, s2) x = self.up4(x, s1) return self.outc(x)几个值得说明的点:Conv2d的bias=False是因为后面紧跟 BatchNorm,BN 自带偏置,两个偏置同时存在只会浪费参数;Down里先池化再卷积,和先卷积再池化感受野相同,但计算量小一半;Up里的中心裁剪是为了处理奇数尺寸输入,训练时 patch 是偶数,这段逻辑主要保证推理时任意尺寸的图都能过。base=64是最常用的通道基数,显存吃紧时降到 32,效果通常只掉一点点。
3.3 张量形状与参数量核对
写 U-Net 最常见的翻车点是 skip connection 通道数对不上。以 256×256 输入为例,各层输出形状如下,可以用print(x.shape)逐层核对:
| 层 | 输出形状 | 通道变化 |
|---|---|---|
| inc | 64×256×256 | 1 → 64 |
| down1 | 128×128×128 | 第一次池化 |
| down2 | 256×64×64 | 第二次池化 |
| down3 | 512×32×32 | 第三次池化 |
| down4 | 512×16×16 | 瓶颈 |
| up1 | 256×32×32 | 512+512 → 256 |
| up4 | 64×256×256 | 恢复输入尺寸 |
| outc | 1×256×256 | 1×1 卷积输出 |
四层下采样意味着输入尺寸至少 16,否则越池化越小,上采样补不回来。环境搭建上,推荐用 anaconda 建独立环境再装 pytorch,注意 torch wheel 与 CUDA 版本的配套关系(cu121 对应 CUDA 12.1),torch.cuda.is_available()返回 False 时先查驱动和安装版本,别急着怀疑代码。
4. 训练流程:数据管道、损失与关键参数
4.1 数据集类:把掩码做进采样管道
训练的输入输出都是同一张图,区别只在掩码替换。所以数据集每次取 patch 时一并生成掩码和替换后的版本,训练循环只做张量搬运,代码会干净很多。
import os import torch import numpy as np from torch.utils.data import Dataset, DataLoader class N2VDataset(Dataset): def __init__(self, folder, patch_size=64, patches_per_image=100, mask_prob=0.1, radius=8): self.folder = folder self.files = [f for f in os.listdir(folder) if f.endswith((".npy", ".tif"))] self.patch_size = patch_size self.n = patches_per_image self.p = mask_prob self.radius = radius def __len__(self): return len(self.files) * self.n def __getitem__(self, idx): img = np.load(os.path.join(self.folder, self.files[idx // self.n])) if img.ndim == 3: # 多通道图先取单通道 img = img[:, :, 0] h, w = img.shape y = np.random.randint(0, h - self.patch_size + 1) x = np.random.randint(0, w - self.patch_size + 1) patch = img[y:y + self.patch_size, x:x + self.patch_size].astype(np.float32) mask = generate_mask(patch.shape, self.p) x_in = apply_blind_mask(patch, mask, self.radius) to_t = lambda a: torch.from_numpy(a).unsqueeze(0) return to_t(x_in), to_t(patch), to_t(mask)__getitem__的三个返回值分别是替换后的输入、原始 patch、掩码。apply_blind_mask取邻居用的是原始 patch,一旦改成取out,相邻掩码会互相搬值,输入分布慢慢被污染,这是个隐蔽 bug。patches_per_image控制每张图每个 epoch 采多少 patch,图多设 50,图少提到 200。
4.2 损失函数:为什么只在掩码位置算 MSE
如果对整张图所有像素都算 MSE,网络会把大量梯度花在"复制输入"这个简单任务上,掩码位置的监督信号被稀释。正确做法是只在掩码处算:
def n2v_loss(pred, target, mask): """pred/target/mask: (B,1,H,W), mask 为布尔张量""" se = (pred - target) ** 2 return se[mask].mean()选 MSE 而不是 L1 的理由要从估计量角度说:给定邻域上下文,MSE 的最优解是条件期望 E[y | x_context],而噪声零均值且逐像素独立时,这个条件期望正等于干净信号。L1 的最优解是条件中位数,对重尾噪声(椒盐、热像素)更稳健,对高斯噪声则收敛慢。如果图里是泊松-高斯混合噪声,可以先做 Anscombe 变换稳定方差,再套 MSE。训练时 loss 的绝对数值没有意义,看它是否稳步下降、输出是否肉眼可见地干净了。
4.3 训练循环与 pytorch 实战参数
import torch.optim as optim device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = N2VUNet(in_ch=1, out_ch=1, base=64).to(device) opt = optim.Adam(model.parameters(), lr=1e-3) sched = optim.lr_scheduler.StepLR(opt, step_size=20, gamma=0.5) ds = N2VDataset("data/train", patch_size=64, patches_per_image=100) loader = DataLoader(ds, batch_size=16, shuffle=True, num_workers=4, pin_memory=True, drop_last=True) for epoch in range(80): model.train() losses = [] for x_in, x, mask in loader: x_in, x, mask = x_in.to(device), x.to(device), mask.to(device) loss = n2v_loss(model(x_in), x, mask) opt.zero_grad() loss.backward() opt.step() losses.append(loss.item()) sched.step() if epoch % 10 == 0: torch.save(model.state_dict(), f"n2v_epoch{epoch:03d}.pt") print(f"epoch {epoch}, loss {np.mean(losses):.4f}")drop_last=True是给 BatchNorm 兜底的:最后一个 batch 只有一两张图时,BN 统计量剧烈抖动。num_workers取 CPU 核数一半左右,再高磁盘 IO 变瓶颈。pin_memory=True配合to(device, non_blocking=True)减少 H2D 拷贝时间,两个参数要一起用。
4.3.1 第一次跑通的最小验证
正式训练前做一次冒烟测试:数据集只放 3 张图,epoch 设 2,batch 设 4。loss 在几十步内明显下降,说明掩码、损失、反传这条链路是通的;loss 纹丝不动,优先检查掩码是不是全是 False(mask_prob写成了 0),以及apply_blind_mask是否真的改了值——打印几个被遮位置的输入和原始值对比,比盯着 loss 猜快得多。
训练参数按下面的表起步,再根据收敛情况微调:
| 参数 | 推荐值 | 调参方向 |
|---|---|---|
| patch_size | 64 | 细节密集时上调到 128 |
| mask_prob | 0.1 | 残差明显时下调,细节丢失时上调 |
| radius | 8 | 噪声颗粒大时加到 16 |
| batch_size | 16 | 显存不足减半,同步调低 lr |
| lr | 1e-3 | 不收敛降到 3e-4 |
| epoch | 80 | 以验证 loss 平台为准 |
5. 加载训练好的模型做推理与评估
5.1 模型保存与导入(pytorch 模型导入)
训练好的模型以state_dict保存,文件名一般为n2v_unet_best.pt。加载时最关键的一点:模型结构必须和保存时完全一致,in_ch、out_ch、base差一个,load_state_dict都会报 missing 或 unexpected key。
model = N2VUNet(in_ch=1, out_ch=1, base=64) state = torch.load("n2v_unet_best.pt", map_location="cpu") model.load_state_dict(state) model.eval().to(device)map_location="cpu"让权重先落在 CPU 再搬运,避免跨机器迁移时 CUDA 设备号对不上报错。加载后必须model.eval(),把 BN 和 dropout 切到推理模式,否则同一张图两次前向结果不同。上面的代码只保存了网络权重;需要断点续训时,把opt.state_dict()和当前 epoch 一起存成 checkpoint。
5.2 推理时要不要继续用掩码
训练时掩码是必须的,推理时两条路线都可以走:
- 直接前向:带噪图原封不动喂进网络,输出即结果。速度快,多数场景够用,因为网络早已学会依赖上下文而非中心值。
- Monte Carlo 推理:每轮重新采样掩码并替换,T 次前向取平均。掩码是随机采的,单次输出某些位置上下文不足,平均能压掉这部分方差。
官方实现偏向后者。实践中建议先用直接前向看效果,输出在局部出现异常亮点时再切 MC:
def mc_denoise(model, x, T=16, mask_prob=0.1, radius=8): """x: (1,1,H,W), 返回 T 次掩码前向的平均""" outs = [] model.eval() with torch.no_grad(): for _ in range(T): m = torch.rand_like(x) < mask_prob x_in = apply_blind_mask(x[0, 0].cpu().numpy(), m[0, 0].cpu().numpy(), radius) x_in = torch.from_numpy(x_in).unsqueeze(0).unsqueeze(0).to(x.device) outs.append(model(x_in)) return torch.stack(outs).mean(0)| 推理方式 | 速度 | 效果 | 适用场景 |
|---|---|---|---|
| 直接前向 | 快 | 稳定,偶有残差 | 批量离线处理 |
| MC 平均 | 慢 T 倍 | 方差更小 | 关键图、大 radius |
无论哪种方式,输入值域必须与训练一致:训练用 [0,1] 就喂 [0,1],训练用原始计数就喂原始计数,归一化不一致是"输出发灰"最常见的原因。
5.3 PSNR、SSIM 评估与无真值时的替代方案
有干净真值(仿真数据)时用 PSNR/SSIM 评估:
from skimage.metrics import peak_signal_noise_ratio as psnr from skimage.metrics import structural_similarity as ssim def evaluate(ref, out): # ref/out: HxW float32, 值域 [0,1] p = psnr(ref, out, data_range=1.0) s = ssim(ref, out, data_range=1.0) return p, s真实场景没有真值,有两个替代方案。一是拍一段静止场景连续帧,时间平均当伪真值,注意被摄物不能动。二是看残差:带噪图 - 输出图如果均匀且无结构,说明噪声被干净分离;如果残差里看到轮廓,说明模型把信号也抹掉了,此时降低掩码比例或改用 L1。参考:σ=25 的高斯噪声仿真图上,N2V 复现一般能把 PSNR 从 20 dB 附近拉到 26~30 dB 区间,低于 24 dB 先怀疑掩码链路而不是换网络。
6. 复现时最容易踩的 5 个坑与快速验证
6.1 现象、原因与对策对照表
| 现象 | 原因 | 对策 |
|---|---|---|
| 输出几乎等于输入 | 掩码没生效或 loss 在全像素算 | 打印掩码均值,检查 se[mask] |
| 训练 loss 不降 | lr 过大或数据没归一化 | lr 降到 3e-4,输入减均值除标准差 |
| 输出发灰、对比度低 | 推理输入值域与训练不一致 | 统一 [0,1] 或统一原始强度 |
| 棋盘格/网格伪影 | 转置卷积叠太多层 | 上采样改 interpolate + conv |
| 细节糊成一片 | patch 太小或 radius 太大 | patch 提到 128,radius 降到 4 |
6.2 用掩码比例做"真去噪"对照实验
最有效的验证是控制实验:把 mask_prob 设为 0,其余不变,训相同 epoch。如果输出仍然是去噪结果,说明网络在走捷径,可能有信息泄漏到被遮位置;如果输出变成输入本身,恰好证明掩码机制真的阻止了恒等映射,模型只能学上下文。这个实验同时验证了损失是否只在掩码位置计算——设 0 后如果 loss 还能正常下降,那一定是在全像素上算了。
6.3 断点续训与随机性控制
断点续训时只加载网络权重,优化器的动量和学习率状态会丢失,曲线比从头训还难看。保存完整 checkpoint 时带上opt.state_dict(),恢复后手动把sched.last_epoch设回保存的 epoch,StepLR 才能从正确位置继续衰减。N2V 的掩码采样完全依赖随机数,同代码两次训练结果差异会很大,排查问题时很难判断是改动造成的还是随机性造成的。训练脚本开头加一行:torch.manual_seed(42); np.random.seed(42),复现就从一个随机事件变成一个可回放的过程。
本文还有配套的精品资源,点击获取