简介:面向深度学习与医学图像分割初学者的PyTorch U-Net完整实现包,提供了从网络搭建、数据预处理到训练测试的一整套流程,可直接用于训练自定义图像分割数据集。压缩包共27个文件,其中py源码覆盖核心模块(如net.py、train.py、data.py、test.py及评估脚本),png为可视化结果,md为说明文档,整体约602KB,结构清晰便于对照学习。目前已有773人学习下载,尤其适合希望理解编码器-解码器、跳跃连接等关键机制,并动手实践自定义分割任务的读者。资源不仅包含可直接运行的U-Net模型与配套工具脚本,还展示了数据标注与增强的处理思路,可辅助完成从数据集准备到分割效果评估的完整闭环。
1. 为什么拿 PyTorch 自己搭 U-Net:图像分割没你想的那么玄
用 PyTorch 搭建 U-Net 这件事,网上一搜一大把教程,但真到了要训练自己的数据集那一步,翻车率其实很高。我见过太多人下了开源代码,跑通了示例图,换了自己的图片之后 loss 死活不降,或者输出的 mask 全黑。问题大多不在网络结构,而在于数据处理和训练细节。这份 pytorch-UNet 项目正好把训练、测试、评估的完整链路都串起来了,代码量不大,适合做图像分割的入门骨架。这篇文章我会按「网络结构 → 数据集制作 → 训练调参 → 避坑 → 验证部署」的顺序把它拆开讲,适合正在做医学影像分割、卫星图地物提取或者工业质检的开发者。
2. 拆解 U-Net 结构:先从 net.py 看懂编码器、解码器和跳跃连接
U-Net 的网络结构图网上到处都是,对称的 U 型、左边收缩右边扩张,但图和代码是两回事。真正让自己能改、能调,还是要逐行读 net.py。这个文件里的实现是典型的原始 U-Net:双卷积块 + 下采样 + 转置卷积上采样 + 跳跃连接。下面从最小单元开始讲。
2.1 双卷积块是 U-Net 的最小单元
U-Net 的基本构件不是单个卷积层,而是「两次卷积 + 激活 + 归一化」的组合。原论文里没有 BatchNorm,但这个项目的实现加了,实际训练时 BN 对收敛帮助很大。
import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super(DoubleConv, self).__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x)这里的 kernel_size=3、padding=1 是个关键选择,它保证卷积前后特征图的宽高不变。如果不加 padding,每次卷积尺寸都会变小,下采样路径的尺寸对齐会变得很麻烦。BatchNorm 放在卷积和 ReLU 之间,作用是对每个 batch 内的特征做归一化,让激活值分布稳定。ReLU 用inplace=True节省显存,在训练深网络时积少成多。第二个卷积的输出通道保持 out_ch 不变,也就是说 DoubleConv 只做一次通道数改变,空间尺寸全程不变。
2.2 编码器:下采样让通道数翻倍、尺寸减半
编码器路径做的事情很规律:先 MaxPool 把特征图缩小一半,再进入 DoubleConv 把通道数翻倍。项目里这个操作被封装成了 Down 模块。
class Down(nn.Module): def __init__(self, in_ch, out_ch): super(Down, self).__init__() self.mpconv = nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_ch, out_ch) ) def forward(self, x): return self.mpconv(x)nn.MaxPool2d(2)把 H 和 W 各缩小一半,池化窗口是 2×2,步长默认等于窗口大小。每经过一次 Down,特征图的尺寸减半、通道数翻倍。在原论文里,这个通道变化序列是 64 → 128 → 256 → 512 → 1024,但这个项目在最后一个 Down 里做了改动,用 512 → 512 而不是 512 → 1024,这样做的直接好处是参数量减少,对显存更友好。如果你的数据量不大,这个改动反而能抑制过拟合。我在自己项目里试过原版 1024 的配置,显存占用差了将近一倍,但精度提升很有限。
2.3 解码器:转置卷积上采样,拼接编码器特征
解码器的核心是 Up 模块。它先用nn.ConvTranspose2d把特征图放大一倍,然后把来自编码器的跳跃连接特征拼过来,最后再过一次 DoubleConv。
class Up(nn.Module): def __init__(self, in_ch, out_ch): super(Up, self).__init__() self.up = nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride=2) self.conv = DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 = self.up(x1) diffY = x2.size()[2] - x1.size()[2] diffX = x2.size()[3] - x1.size()[3] x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x = torch.cat([x2, x1], dim=1) return self.conv(x)注意 forward 的参数:x1 是来自上一层的上采样特征,x2 是来自编码器同层的跳跃连接。F.pad是处理尺寸不整除情况的保险操作,比如原图尺寸是奇数时,下采样后特征图尺寸对不上,直接拼接会报错。实际中我一般把输入图片统一 resize 成 16 的倍数,这样 pad 分支基本不会触发,但留着它能让网络更健壮。
转置卷积的 kernel_size=2、stride=2 意味着上采样正好放大 2 倍。in_ch // 2是转置卷积的输出通道数,这是为了在拼接后通道数回到期望值。拼接用的是torch.cat而不是相加,这是 U-Net 的原始设计:编码器的浅层特征保存着边缘和纹理信息,解码器的深层特征语义更强,拼接让两部分信息互补,相当于给解码器开了一条「短接线」。
2.4 forward 完整数据流:从输入到输出的逐层走向
把前面几个模块串起来,就是完整的 UNet 类。
class UNet(nn.Module): def __init__(self, n_channels, n_classes): super(UNet, self).__init__() self.inc = DoubleConv(n_channels, 64) self.down1 = Down(64, 128) self.down2 = Down(128, 256) self.down3 = Down(256, 512) self.down4 = Down(512, 512) self.up1 = Up(1024, 256) self.up2 = Up(512, 128) self.up3 = Up(256, 64) self.up4 = Up(128, 64) self.outc = nn.Conv2d(64, n_classes, 1) def forward(self, x): x1 = self.inc(x) x2 = self.down1(x1) x3 = self.down2(x2) x4 = self.down3(x3) x5 = self.down4(x4) x = self.up1(x5, x4) x = self.up2(x, x3) x = self.up3(x, x2) x = self.up4(x, x1) return self.outc(x)forward 里的数据流很清晰:x1 到 x5 是编码器的逐层输出,up1 的输入是 x5 和 x4,输出通道 256。之后每一层 Up 接收的通道数都是上一层的输出加上跳跃连接的通道。最后一个 Up 输出 64 通道,经过 1×1 卷积nn.Conv2d(64, n_classes, 1)变成类别数。
两个构造参数需要重点理解。n_channels是输入图像的通道数,RGB 图填 3,灰度图填 1。n_classes是分割类别数,二分类任务填 1,多分类填实际类别数。最后的 1×1 卷积相当于是把 64 通道的特征映射到类别空间,二分类时输出单通道,后面接 sigmoid。
3. 数据集准备:把原始图片整理成 JPEGImages 和 SegmentationClass
U-Net 的数据集组织方式沿用了 VOC 分割的目录结构:原始图片放 JPEGImages,标注 mask 放 SegmentationClass。这个项目里还带了一个 make_mask_data.py 脚本,专门处理「把标注转成模型能读的 mask」这件事。很多人拿到项目第一步就卡在这里,因为网上能找到的示例数据都是处理好的,自己的标注却五花八门。
3.1 目录约定:VOC 格式的命名与类别映射
项目根目录下有两个关键文件夹:JPEGImages 和 SegmentationClass。前者放原始图片,后者放对应的分割标签图。命名必须一一对应,举个例子,JPEGImages/001.jpg对应的 mask 就是SegmentationClass/001.png。
类别映射是整个数据准备里最容易出问题的环节。VOC 格式的 mask 图里,背景像素是黑色(RGB 值为 0,0,0),目标区域用纯色块填充,比如红色代表类别 1,绿色代表类别 2。模型训练时并不直接读这些 RGB 值,而是要把它们转换成单通道索引图:每个像素存一个整数,0 表示背景,1 表示第一个类别,2 表示第二个类别。
我一般会用下面这个函数来检查自己的 mask 到底有几个类别:
import numpy as np from PIL import Image mask = Image.open("SegmentationClass/001.png") arr = np.array(mask) print("mask shape:", arr.shape) print("unique values:", np.unique(arr))关键要看输出。如果 mask 是单通道灰度图,unique values 应该是一组小整数,比如 [0, 1, 2]。如果 mask 是三通道 RGB,unique values 会是一大串 0 到 255 的值,那就需要先量化成索引。如果出现 255,说明标注里用了白色表示目标,但模型期望的是 0 和 1,这个不处理,训练出来的结果必然不对。
3.2 make_mask_data.py 的 mask 制作逻辑
这个脚本做的事情就是把标注图从「人眼看的格式」转成「模型吃的格式」。核心逻辑是像素级颜色映射:
import os import numpy as np from PIL import Image color_to_class = { (0, 0, 0): 0, # 背景 (255, 0, 0): 1, # 类别1:红色 (0, 255, 0): 2, # 类别2:绿色 } def make_mask(src_dir, dst_dir): os.makedirs(dst_dir, exist_ok=True) for name in os.listdir(src_dir): if not name.lower().endswith((".png", ".jpg", ".jpeg")): continue rgb = Image.open(os.path.join(src_dir, name)).convert("RGB") arr = np.array(rgb) h, w, _ = arr.shape mask = np.zeros((h, w), dtype=np.uint8) for color, cls in color_to_class.items(): match = (arr[..., 0] == color[0]) & \ (arr[..., 1] == color[1]) & \ (arr[..., 2] == color[2]) mask[match] = cls Image.fromarray(mask).save(os.path.join(dst_dir, os.path.splitext(name)[0] + ".png"))这段代码的逻辑很直白:遍历每一张原始图片,把 RGB 像素值和预设的颜色字典比对,命中的像素位置写入对应的类别编号。最后保存成单通道 PNG,PNG 格式支持单通道灰度图,不会丢信息。注意dtype=np.uint8不能省,如果不指定,默认生成的数组是 int64,PIL 保存时会报错或者写出 16 位图。
实际项目里 labels 往往不是纯色块,尤其医学影像里很多标注工具生成的是带边缘羽化的 PNG。这种情况我一般会先把像素 RGB 值转成 HSV,再按色相区间归类,或者直接用标注工具导出的索引 PNG,而不是在脚本里做颜色猜测。make_mask_data.py 适用于「颜色规范、边界清晰」的标注,如果你的标注来源复杂,先把颜色统一再做映射。
3.3 数据增强与归一化:img 和 mask 必须同步变换
数据集类里最容易踩的坑是:图片做了旋转翻转,mask 也跟着变,但很多人写增强时只对 img 做了变换,mask 没动。结果训练时模型看到的是「错位的标签」,loss 高到爆炸还找不到原因。
import os import random from PIL import Image from torch.utils.data import Dataset class VOCDataset(Dataset): def __init__(self, img_dir, mask_dir, size=(512, 512)): self.img_dir = img_dir self.mask_dir = mask_dir self.names = [n for n in os.listdir(img_dir) if n.lower().endswith((".png", ".jpg", ".jpeg"))] self.size = size def __getitem__(self, idx): name = self.names[idx] img = Image.open(os.path.join(self.img_dir, name)).convert("RGB") mask = Image.open(os.path.join(self.mask_dir, os.path.splitext(name)[0] + ".png")) img = img.resize(self.size, Image.BILINEAR) mask = mask.resize(self.size, Image.NEAREST) if random.random() > 0.5: img = img.transpose(Image.FLIP_LEFT_RIGHT) mask = mask.transpose(Image.FLIP_LEFT_RIGHT) if random.random() > 0.5: img = img.transpose(Image.FLIP_TOP_BOTTOM) mask = mask.transpose(Image.FLIP_TOP_BOTTOM) img = np.array(img, dtype=np.float32) / 255.0 mask = np.array(mask, dtype=np.int64) img = img.transpose(2, 0, 1) # HWC -> CHW return torch.from_numpy(img.copy()), torch.from_numpy(mask.copy())注意两个 resize 的插值方式不同。图像用 BILINEAR 双线性插值,保留平滑的边缘;mask 必须用 NEAREST 最近邻插值,否则类别边界会出现「新类别」——比如 0 和 1 之间插值出 0.5,模型就懵了。随机翻转时 img 和 mask 要执行同一种变换,上面代码用的两个独立 if 块其实有问题,应该用一个随机种子或者同时翻转。这是我实际写代码时踩过的坑,正确写法是用seed = random.random()决定是否翻转,img 和 mask 用同一个 seed。
归一化方面,图像除以 255 后取值范围是 0 到 1,足够用了。ImageNet 的 mean/std 归一化在 U-Net 上不一定更好,我试过几次,对分割任务提升不明显,反而多一步换算。mask 保持 0、1、2 这样的整数,不做归一化,交叉熵损失要求的是索引值而不是 one-hot 概率。
3.4 数据量不足:切 patch 和离线增强是补救手段
医学图像分割经常遇到「标注图只有三五十张」的情况。U-Net 虽然比大模型的参数少,但没有预训练权重时,30 张图根本喂不饱。常见做法是切 patch:把 1024×1024 的大图切成 256×256 的 patch,相邻 patch 之间可以设 50% 重叠,相当于把数据量翻了十几倍。训练时额外加在线增强,包括 90 度旋转、缩放、亮度抖动、高斯噪声。如果还是不够,再考虑用更大的切图步长生成更多 patch。但要注意,切 patch 后要让训练集和验证集来自不同的原图,否则模型会「记住」重叠区域,验证指标虚高。
4. 训练自己的数据集:train.py 的参数设置和损失函数选型
数据集就绪之后,训练脚本就是整个项目的引擎。train.py 里做了数据加载、模型初始化、前向传播、损失计算、反向传播和模型保存。这个环节的选型直接决定模型能不能收敛。
4.1 训练入口与数据加载:先确认 PyTorch 环境
开始训练之前,先保证基础环境是干净的。我的习惯是 conda 新建一个独立环境,避免系统 Python 里一堆包的版本冲突。如果你还没搭好 PyTorch 环境,下面是常见做法:
conda create -n unet python=3.8 conda activate unet pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118装完后可以用一行命令验证 CUDA 是否可用:
import torch print(torch.cuda.is_available(), torch.cuda.get_device_name(0))这里有个常见问题:很多人装的是 CPU 版 PyTorch,训练时才发现torch.cuda.is_available()返回 False,白白跑了一晚上。所以我一般会在 train.py 开头强制检查设备状态,而不是靠运气。数据加载部分用 PyTorch 的DataLoader,核心参数是 batch_size、shuffle 和 num_workers:
from torch.utils.data import DataLoader train_dataset = VOCDataset("JPEGImages", "SegmentationClass", size=(512, 512)) train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=2) model = UNet(n_channels=3, n_classes=1) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device)num_workers是数据加载的子进程数,Windows 上设成 0 最稳,Linux 可以设成 2 或 4。shuffle=True 必须在每个 epoch 打乱数据顺序,否则模型会学到数据的排列顺序,而不是语义特征。
4.2 损失函数:BCE 和 Dice Loss 怎么选
二分类分割任务里,这个项目用的是nn.BCEWithLogitsLoss。它的特点是内部先做了 sigmoid 再算交叉熵,数值上比「自己算 sigmoid + BCE」更稳定,因为避免了 log(0) 的情况。多分类任务则用nn.CrossEntropyLoss,它内部做了 softmax,输入的是网络最后的原始 logits。
医学分割场景下,我更推荐把 Dice Loss 也试一下。Dice 系数衡量的是预测和真实 mask 的重叠率,它对前景占比很小的图特别友好,因为交叉熵在小目标上容易被背景淹没。常用的组合是bce + dice,两者加权相加:
import torch.nn.functional as F def bce_dice_loss(pred, target, alpha=0.5): bce = F.binary_cross_entropy_with_logits(pred, target.float()) pred_prob = torch.sigmoid(pred) smooth = 1.0 dice = 1 - (2 * (pred_prob * target).sum() + smooth) / \ (pred_prob.sum() + target.sum() + smooth) return alpha * bce + (1 - alpha) * dice这里target需要是 float 类型,且形状和 pred 一致。smooth 参数是为了防止分子分母都是 0 时除零错误,一般设 1 就够了。alpha 控制两个损失的比例,常见做法是 0.5 对半分。如果你的数据集正负样本比例严重失衡,比如病灶只占整张图的 5%,建议把 alpha 调低到 0.3,让 Dice 占主导。
4.3 优化器、学习率、batch size 和权重初始化
优化器这块没什么悬念,Adam 是 U-Net 训练的首选。它的自适应学习率让新手不用频繁调参,收敛速度也快。
| 超参数 | 推荐值 | 备注 |
|---|---|---|
| optimizer | Adam | 默认 betas=(0.9, 0.999) |
| 初始学习率 | 1e-3 | 如果 loss 震荡,降到 1e-4 |
| batch_size | 4~8 | 取决于显存,512×512 输入通常 4 就够 |
| epochs | 50~200 | 小数据集关注验证集,不要盲目堆 epoch |
| 输入尺寸 | 512×512 | 需要能被 16 整除 |
import torch.optim as optim optimizer = optim.Adam(model.parameters(), lr=1e-3) criterion = nn.BCEWithLogitsLoss() for epoch in range(100): model.train() for imgs, masks in train_loader: imgs, masks = imgs.to(device), masks.to(device).float() preds = model(imgs) loss = criterion(preds, masks.unsqueeze(1)) optimizer.zero_grad() loss.backward() optimizer.step() print(f"Epoch {epoch}, Loss: {loss.item():.4f}")masks.unsqueeze(1)把 [B, H, W] 变成 [B, 1, H, W],因为网络输出是单通道,必须对齐维度。如果忘记做这一步,PyTorch 会报维度不匹配的错误,这是新手最常见的问题。optimizer.zero_grad()每次迭代都要调用,否则梯度会累加。学习率方面,小数据集上我习惯用ReduceLROnPlateau来降低学习率:验证集 loss 连续 5 个 epoch 不降就除以 10。
权重初始化也值得注意。PyTorch 的 Conv2d 默认初始化是 Kaiming 均匀分布,U-Net 结构比较深,直接训练基本没问题。但如果你的数据量很少,建议载入 ImageNet 上预训练的 encoder 部分,再微调解码器。
5. 避坑指南:U-Net 训练自己的数据集常见的 5 个翻车现场
训练 U-Net 的过程,本质上是在和数据格式、维度、显存作斗争。下面这几条都是我多次见过的实际问题,每一条都是「现象 → 原因 → 解决」的记录。
5.1 mask 是三通道 RGB,训练时 loss 死活不降
现象:训练时 loss 维持在 0.69 左右不动,输出图全黑或者全白。
原因:mask 是从标注工具直接导出的 RGB 三通道 PNG,不是单通道索引图。模型输出是单通道,但计算损失时 mask 是 [B, 3, H, W],两个张量对不上,模型只能在三个通道里「猜」。
解决:用 3.2 节的make_mask_data.py把 RGB mask 转成单通道索引图。转换后检查一下np.unique(mask)的输出,确保只有 0、1、2 这种小整数。
5.2 mask 像素值是 0 和 255,不是 0 和 1
现象:训练能跑,loss 也在下降,但验证集准确率低得离谱。
原因:很多标注工具把目标区域标成纯白(255),背景标成纯黑(0)。BCEWithLogitsLoss 期望的目标是 0 到 1 之间的概率,用 255 和 0 去算交叉熵,相当于用一个被放大的错误值去更新梯度。
解决:训练脚本里对 mask 做一次归一化,mask = (mask > 0).float(),把大于 0 的值全部变成 1。注意这一步在 DataLoader 里做,而不是在数据预处理阶段做。
5.3 输入尺寸不是 16 的倍数,训练时跳跃连接报尺寸错误
现象:报错信息形如RuntimeError: Sizes of tensors must match except in dimension 1。
原因:U-Net 有 4 次下采样,特征图尺寸每层减半。如果输入尺寸不是 16 的倍数,经过 4 次池化后,不同分支的特征图尺寸会出现 1 像素的差异,拼接时直接报错。
解决:所有图片统一 resize 到 16 的倍数,比如 256、320、512。如果你不想 resize,也可以像我一样在 Up 模块里保留F.pad补齐,但尽量用统一尺寸,让数据分布更稳定。
5.4 CUDA 显存溢出,batch_size 调到 1 还是不够
现象:CUDA out of memory,程序直接崩。
原因:输入图分辨率太高,比如原始 CT 图是 1024×1024,即使 batch_size=1,前向传播时中间特征图的显存开销也很大。
解决:两个方案。一是把输入尺寸改成 512×512 或 256×256;二是改用 patch 训练,把大图切成 256×256 的小块。医学影像场景里我一般用后者,因为下采样后的特征图能保留更多细节。推理时再用 overlap-tile 策略拼接回原尺寸。
5.5 训练集只有几十张图,模型过拟合严重
现象:训练集 loss 降到 0.01,验证集 Dice 只有 0.5,训练曲线前后差距巨大。
原因:模型参数几百万,数据量太少,网络把训练集的噪声细节全背下来了。
解决:数据增强 + 减少模型容量。U-Net 的通道数可以从 64 起步砍成 32,参数量直接减少 4 倍。还可以加 Dropout,或者用带预训练的 encoder 做迁移学习。如果标注成本可以接受,优先找更多数据,这是最踏实的解法。
6. 验证与部署:test.py 结果可视化,以及模型转 ONNX 的实用技巧
训练完成后,验证和部署是真正检验模型价值的两步。test.py 加载训练好的权重,对新图片做前向推理,输出分割 mask。评估脚本 get_evaluation.py 会计算 Dice 和 IoU 指标,值得在每次训练后都跑一遍。
model = UNet(n_channels=3, n_classes=1) model.load_state_dict(torch.load("model.pt", map_location="cpu")) model.eval() with torch.no_grad(): pred = model(img.unsqueeze(0)) pred = torch.sigmoid(pred).squeeze().numpy() pred_binary = (pred > 0.5).astype(np.uint8) * 255img.unsqueeze(0)是为了把 [3, H, W] 变成 [1, 3, H, W] 的 batch 维度。model.eval()必须显式调用,它会关闭 Dropout 和 BatchNorm 的统计更新,否则同一张图推理两次结果会不一样。阈值 0.5 是二分类的默认选择,如果你的类别严重不平衡,需要根据验证集的 precision-recall 曲线调整。
模型导出 ONNX 是部署到实际工程里的关键一步。常见的做法是导出固定尺寸的 ONNX,再验证输出一致性:
dummy_input = torch.randn(1, 3, 512, 512) torch.onnx.export( model, dummy_input, "unet.onnx", opset_version=11, input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}} )这里只把 batch 维设为动态,输入分辨率保持固定,因为很多部署框架对动态分辨率支持不稳定。导出前模型必须处于 eval 状态,否则 BatchNorm 会被封装进计算图。之后用 onnxruntime 加载,和 PyTorch 推理结果对比,最大误差一般应该小于 1e-5,如果差异明显,多半是预处理方式不一致。
从那以后,我每次训练完都会强制走一遍「test 单张图 → 算 Dice → 导出 ONNX → 对比一致性」这条流程,而不是只看训练 loss。没有这最后一步,模型在验证集上再漂亮,部署时也可能翻车。希望帮到你。
本文还有配套的精品资源,点击获取