FEATURED · 精选文章

PyTorch从零实现U-Net图像分割:模型搭建、数据集制作与训练全流程

发布时间 / 2026/9/15 12:45:20
来源 / 创域科博编辑部
栏目 / 资讯中心
PyTorch从零实现U-Net图像分割:模型搭建、数据集制作与训练全流程 U-Net这个结构搞图像分割的人应该都不陌生。从医学影像到遥感图像从缺陷检测到抠图几乎各行各业凡是涉及到“把像素分类”的活儿都能看到它的身影。我之前也复现过不少次但说实话网上很多教程要么只贴了一堆代码让你自己跑要么把数据集直接给你准备好你下载下来一训练完事儿——等到真需要处理自己的业务数据时反而不知道从哪儿下手。这篇文章就围绕两件事展开第一从零复现一个PyTorch版本的U-Net把每一层结构掰开揉碎讲清楚第二把你自己的图片变成能训练的数据集完成从标注到训练的完整闭环。无论你是刚入门分割任务的学生还是工作中需要做图像检测的工程师这篇文章都值得花二十分钟看完我尽量把每一步的“为什么这么做”也讲明白而不是只给结论。1. 整体思路拆解为什么还要手动复现U-Net直接调库不行吗做深度学习快的人可能觉得用现成的库不就行了比如分割模型库MMSegmentation里面就有U-Net装好环境改个配置就能跑为什么还要自己动手写一遍我个人的理解是复现和调库完全不是一个学习量级。U-Net算是最适合用来“解剖”的模型之一它结构对称、逻辑清晰、没有特别复杂的奇技淫巧代码量不大但包含了现代CNN分割模型的核心思路——编码器提取特征、解码器恢复分辨率、跳跃连接融合多层信息。你手写一遍对卷积、上采样、特征融合这些概念的理解比看十篇博客都有用。从工程角度看自己复现的代码也更灵活。比如你后续想在这个基础结构上加一个注意力模块或者把骨干网络换成ResNet直接在源码上改比去改别人的框架参数要方便得多。而且自己写的数据加载和训练代码出问题时你能快速定位到是数据的问题、模型的问题还是训练策略的问题不会像黑盒一样摸不着头脑。1.1 U-Net的核心原理U型结构、跳跃连接到底在干什么U-Net的结构用一个词概括就是“对称”。它分成左边一条编码路径Contracting Path和右边一条解码路径Expanding Path中间通过“跳跃连接”Skip Connection把对应分辨率层的特征拼起来整个网络结构看起来像个字母U所以叫U-Net。编码路径做的事本质上是不断“压缩”图像。每一层都是两次卷积加一次下采样通常是最大池化图像尺寸越来越小但通道数越来越多。这个过程让网络能够看到越来越大的感受野——也就是说后面的层能“看到”原图中更大的区域从而理解物体的上下文信息。但是下采样是有代价的图像细节会丢失边缘信息会模糊。解码路径则反过来把编码器压缩后的特征图逐步恢复尺寸。它用上采样通常是转置卷积或双线性插值把特征图放大一倍然后和编码路径对应层的特征图在通道维度上拼接起来。为什么这么干因为上采样恢复的是空间分辨率但放大的过程会引入噪声、丢失细节而跳跃连接带来的高分辨率特征图正好能把下采样时保留的精细边界信息“补”回来。说一个我自己的类比编码器像一个写摘要的人把整篇文章浓缩成了一段话这段话概括了主旨但丢失了细节解码器就是根据这段话重新扩写成完整文章而跳跃连接相当于把原稿里的关键句子直接抄给扩写的人让扩写出来的内容不至于跑偏太多。分割任务恰恰非常依赖细节因为你要把每个像素归类边界的精确度直接影响最终效果。1.2 方案选型PyTorch版本、损失函数、评估指标怎么选PyTorch在学术界和工业界都是目前最主流的选择之一。它的动态图机制让调试变得很直观你可以随时print中间变量的大小和值不像静态图那样要先构图再执行。U-Net本身结构相对简洁用PyTorch的nn.Module写起来非常顺手而且社区资料多遇到问题基本都能搜到解法。损失函数方面图像分割最常用的是交叉熵损失CrossEntropyLoss。但实际做二类分割任务时正负样本比例经常严重失衡——比如一张图上缺陷只占几个像素大部分都是背景。这种情况下直接算交叉熵模型会倾向于把所有像素都预测成背景因为这样loss已经很低了。我一般会根据任务情况配合使用Dice Loss或者Focal Loss来解决样本不平衡这部分后面我会细讲。评估指标不要只看准确率Accuracy。如果目标区域占比很小全部预测成背景准确率也很高这样会严重误导判断。我更推荐同时关注IoU交并比和Dice系数。IoU是预测区域和真实区域的交集除以并集能从几何重叠的角度衡量分割效果Dice系数的含义与IoU类似但对小目标的敏感度更高两者结合着看能更准确地评估模型好坏。2. 环境准备与PyTorch基础工欲善其事必先利其器开始写代码之前先把环境准备好。这一部分看起来简单但我在网上帮人解决问题的时候发现很大比例的问题都出在环境配置上——比如CUDA版本不匹配、PyTorch装了CPU版本、包冲突等等训练时才发现GPU根本不工作或者疯狂报错。2.1 Anaconda创建虚拟环境给项目一个干净的家我的习惯是一个项目建一个独立的Anaconda虚拟环境因为不同项目依赖的包版本可能冲突。比如你另一个项目要用TensorFlow它要求的Python版本可能跟PyTorch不同塞在同一个环境里很容易出问题。虚拟环境相当于给每个项目一个独立的小房间互不干扰。创建环境的命令很简单在终端里执行conda create -n unet python3.9 conda activate unetPython版本我推荐3.8到3.10之间太高的版本有些库可能还没有适配好太低的版本又会有兼容性问题。激活环境后接下来所有的安装操作都在这个环境里进行。2.2 安装PyTorch和必要依赖CPU版还是GPU版安装PyTorch之前先确认你的机器有没有NVIDIA显卡。如果你有显卡还想显卡加速训练那一定要安装CUDA版本的PyTorch。判断方法是在终端输入nvidia-smi如果有输出就能看到你的显卡支持的最高CUDA版本。然后去PyTorch官网按页面上的提示选择你的操作系统、包管理工具和CUDA版本官网会自动生成安装命令。这里分享一个经验不要直接用pip install torch这样默认装的是CPU版本会导致后面训练时慢得让人崩溃而且你还不容易察觉。装完以后可以用下面这段代码验证GPU是否可用import torch print(torch.__version__) print(torch.cuda.is_available())如果torch.cuda.is_available()返回True说明GPU能正常使用。如果返回False很不幸你需要检查一下安装命令里有没有选对CUDA版本或者显卡驱动是不是有问题。其他依赖就比较简单了用pip install装一下就行主要是后面处理图像要用的opencv-python、numpy以及训练进度显示用的tqdmpip install opencv-python numpy tqdm pillow matplotlib有一点我踩过坑提醒大家注意就是opencv-python和opencv-contrib-python不要同时装会有冲突。用pip uninstall把另一个清理干净再装你需要的那个。3. U-Net核心代码实现从零手写每一个模块环境搞定了接下来就是重头戏——用PyTorch实现U-Net。我不会把一大坨代码一次性丢给你让你去复制而是拆成几个模块来讲这样你看着不累理解也更透彻。等你真正理解了每个模块的作用组合起来自然水到渠成。3.1 基础卷积块两次卷积操作U-Net里最基础的单元是“两次卷积”在它的编码器和解码器的每一层都会用到。它的结构是卷积 - 批归一化 - ReLU激活 - 卷积 - 批归一化 - ReLU激活。为什么要连续做两次卷积而不是一次因为两次卷积堆叠能增大感受野让模型提取到更复杂的特征。单层卷积只能看局部小区域两层卷积堆叠之后实际上能够组合出更抽象的模式。批归一化BatchNorm在这里的作用是加速收敛同时对初始化的敏感性降低训练会更稳定。代码实现如下import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super(DoubleConv, self).__init__() self.conv nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x)这里需要注意padding1这个参数它的作用是保持卷积后特征图尺寸不变。因为3x3卷积核如果不加padding输出尺寸会比输入小2个像素经过多层卷积后尺寸就会越缩越小给后面的连接造成麻烦。设置padding1后输出尺寸和输入尺寸保持一致这样在U-Net中不同层之间的特征图大小就完全对得上了。3.2 编码器路径下采样与特征提取编码器部分做的事情简单说就是“尺寸减半通道数加倍”。每一次操作先经过一个DoubleConv提取特征然后通过最大池化把尺寸缩小一半。为什么用最大池化而不是卷积做下采样最大池化的好处是它不引入额外的可学习参数而且取区域最大值的方式对微小的空间位移有一定鲁棒性。虽然现在有些架构喜欢用stride2的卷积下采样但U-Net的经典设计就是最大池化按照经典来就好效果也不差。class Down(nn.Module): def __init__(self, in_channels, out_channels): super(Down, self).__init__() self.maxpool_conv nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_channels, out_channels) ) def forward(self, x): return self.maxpool_conv(x)输入一张3通道的图片经过第一层Down后尺寸变为原来的一半通道数变为64再经过一层Down尺寸变为原来的四分之一通道数变为128依此类推。整个编码器把输入图像从大尺寸、少通道变成了小尺寸、多通道的特征图。3.3 解码器路径上采样与跳跃连接解码器的核心操作是“上采样 跳跃连接拼接”。上采样用转置卷积实现将特征图尺寸翻倍。然后从编码器的对应层取出之前保存的特征图在通道维度上直接拼接cat。class Up(nn.Module): def __init__(self, in_channels, out_channels): super(Up, self).__init__() self.up nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size2, stride2) self.conv DoubleConv(in_channels, out_channels) def forward(self, x1, x2): x1 self.up(x1) # 跳跃连接将编码器对应的特征图和上采样后的特征图拼接 x torch.cat([x2, x1], dim1) return self.conv(x)这里有个细节要特别留意torch.cat拼接是在通道维度上也就是dim1PyTorch中张量的格式是NCHWN是批次大小C是通道数H是高度W是宽度。拼接后通道数变成了两者之和所以DoubleConv的第一个输入通道数in_channels必须是拼接后的总通道数。这也是为什么定义Up模块时传入的in_channels实际上是两边通道数之和。在我上面写的代码里ConvTranspose2d(in_channels, in_channels // 2, kernel_size2, stride2)的意思是输入in_channels个通道输出in_channels // 2个通道尺寸翻倍。然后in_channels // 2正好等于编码器对应层的通道数所以拼接后总通道数又变回了in_channels传给下一层DoubleConv刚刚好。这个配对关系在U-Net的对称结构里是设计好的写的时候不要搞混。3.4 完整网络组装输入输出与输出层设计有了编码器、解码器和基础卷积块就可以组装出完整的U-Net了。以最经典的输入单通道灰度图、二分类分割为例class UNet(nn.Module): def __init__(self, n_channels, n_classes): super(UNet, self).__init__() self.n_channels n_channels self.n_classes n_classes 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, 1024) self.up1 Up(1024, 512) self.up2 Up(512, 256) self.up3 Up(256, 128) self.up4 Up(128, 64) self.outc nn.Conv2d(64, n_classes, kernel_size1) def forward(self, x): x1 self.inc(x) # 64, H, W x2 self.down1(x1) # 128, H/2, W/2 x3 self.down2(x2) # 256, H/4, W/4 x4 self.down3(x3) # 512, H/8, W/8 x5 self.down4(x4) # 1024, H/16, W/16 x self.up1(x5, x4) # 512, H/8, W/8 x self.up2(x, x3) # 256, H/4, W/4 x self.up3(x, x2) # 128, H/2, W/2 x self.up4(x, x1) # 64, H, W logits self.outc(x) # n_classes, H, W return logits这里有几个要点值得说。第一最后的输出层用的是1x1卷积它的作用不是提取特征而是把64个通道映射到n_classes个通道上。如果做二分类n_classes1输出的每个像素点是一个数值经过Sigmoid函数映射到0到1之间表示属于前景的概率。如果做多分类n_classesK输出的每个像素点有K个数值经过Softmax得到各个类别的概率分布。第二输入图像尺寸最好设计成16的倍数。因为整个网络一共做了4次下采样每次尺寸减半所以输入尺寸必须是2的4次方也就是16的倍数才能保证所有特征图的尺寸都是整数。如果输入尺寸不合适到后期上采样拼接时会因为尺寸对不上直接报错。我一般会先把图片统一缩放或裁剪成例如256x256或512x512的尺寸再送进网络。第三经典的U-Net原论文里是4层下采样通道数是64到1024。你可以根据自己算力情况调整比如减少通道数让模型更小跑得更快代价是精度可能会下降。我的建议是先按经典结构跑通再根据结果去优化。4. 制作自己的数据集从原始图片到可直接训练模型代码写完了但现在还是“无米之炊”——你得有数据才能训练。市面上公开的分割数据集很多比如做医学影像的有很多开源数据集做遥感的有一些地物分类数据集但实际项目中特别是工业缺陷检测这种场景往往没有现成的数据集得自己造。制作数据集的核心流程是收集图片 - 标注 - 划分数据集 - 数据增强 - 写数据加载类。每一步都有不少坑我逐个讲。4.1 图像标注用Labelme制作分割标签分割任务需要的是“像素级”的标签也就是对每一个像素标注它属于哪一类。手工一个像素一个像素画显然不现实正确的做法是画多边形轮廓然后自动填充成掩码Mask。我常用的标注工具是Labelme它是一个开源的图像标注工具安装非常方便界面友好。pip install labelme labelme启动后打开要标注的图片文件夹用“Create Polygons”工具沿着目标边缘点出一圈关键点勾勒出目标轮廓然后给这个多边形命名一个类别标签。一张图里可以标多个目标、多个类别。标完之后保存每张图片会生成一个同名的JSON文件里面记录了多边形的顶点坐标和标签名称。有一点体验上的提醒标注是个体力活非常消耗耐心。如果目标边缘复杂建议适当放大图片再标不然点出来的轮廓锯齿感会很严重训练出来的分割边界也会很难看。4.2 JSON标签转掩码核心转换逻辑Labelme生成的JSON文件不能直接拿来训练得把它解析成模型需要的掩码图。掩码图是一张和原图尺寸相同的灰度图对于二分类问题0表示背景255表示前景对于多分类问题0表示背景1、2、3等整数表示不同类别。转换的核心代码思路是先读取JSON文件用json.load解析出多边形的坐标点然后在掩码图上通过cv2.fillPoly把多边形内部填充成对应的类别值。下面这段代码是我项目中常用的转换逻辑import json import numpy as np import cv2 import os def json_to_mask(json_path, img_width, img_height, label_map): with open(json_path, r, encodingutf-8) as f: data json.load(f) mask np.zeros((img_height, img_width), dtypenp.uint8) for shape in data[shapes]: label shape[label] points np.array(shape[points], dtypenp.int32) if label in label_map: cv2.fillPoly(mask, [points], label_map[label]) return masklabel_map是一个字典例如{background: 0, defect: 255}它把标注时写的类别名称映射成掩码像素值。注意多边形顶点坐标是浮点数转成int32类型因为fillPoly要求整数坐标。生成掩码之后我习惯把原图和掩码都保存成.png格式放到统一的目录结构里这样后续写数据加载类会更方便dataset/ ├── images/ │ ├── 001.png │ ├── 002.png │ └── ... └── masks/ ├── 001.png ├── 002.png └── ...4.3 Dataset类实现PyTorch数据加载的正确姿势有了原图和掩码接下来需要把它们封装成PyTorch的Dataset。PyTorch的torch.utils.data.Dataset类要求实现两个方法__len__返回数据总数__getitem__根据索引返回一对样本和标签。我的实现里还会顺便做数据增强和归一化from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as T class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, img_size256, is_trainTrue): self.image_dir image_dir self.mask_dir mask_dir self.img_size img_size self.is_train is_train self.images sorted(os.listdir(image_dir)) def __len__(self): return len(self.images) def __getitem__(self, idx): img_name self.images[idx] img_path os.path.join(self.image_dir, img_name) mask_name os.path.splitext(img_name)[0] .png mask_path os.path.join(self.mask_dir, mask_name) image Image.open(img_path).convert(RGB) mask Image.open(mask_path).convert(L) # 统一尺寸 image image.resize((self.img_size, self.img_size), Image.BILINEAR) mask mask.resize((self.img_size, self.img_size), Image.NEAREST) # 数据增强 if self.is_train: # 这里可以加随机翻转、旋转等 pass # 转tensor并归一化 image T.ToTensor()(image) mask T.ToTensor()(mask) return image, mask写这一段的时候有几个细节要特别留意。第一掩码resize时一定要用Image.NEAREST最近邻插值不要用双线性或双三次插值。因为掩码是类别标签双线性插值会在类别边界产生中间值比如本来是0和255插值后可能出现128这会让模型非常困惑——你让它预测0到1之间的概率但标签却是0.5语义上完全不可解释。第二T.ToTensor()会把PIL图像的像素值从0到255缩放到0到1这一点模型训练时需要注意。如果后面要可视化记得把张量乘以255再转回numpy。第三__getitem__里最好不要每次都做复杂的文件读取和resize尤其是数据量大时这会拖慢训练速度。一个优化思路是在初始化时就把所有图片提前读入内存前提是显存和内存够用或者使用torch.utils.data.DataLoader的num_workers参数开启多进程加载让数据读取和模型训练并行。4.4 数据增强策略让小数据集发挥大价值标注数据来之不易数量往往有限而深度学习模型非常“贪吃”数据少就容易过拟合。数据增强就是在不改变语义标签的前提下对图片做各种变换让模型看到更多样化的输入。对分割任务而言最常用也最安全的增强方式有水平翻转、垂直翻转、90度旋转。这些操作不会破坏图像和掩码的对应关系因为原图和掩码做完全一样的变换就行。我做增强时的做法是先在__getitem__里用一个随机种子同时作用于图片和掩码保证它们做相同的变换if self.is_train: # 随机水平翻转 if torch.rand(1) 0.5: image T.functional.hflip(image) mask T.functional.hflip(mask) # 随机垂直翻转 if torch.rand(1) 0.5: image T.functional.vflip(image) mask T.functional.vflip(mask)写这段代码有几个实现细节值得提一下。第一种是在__getitem__里逐样本做增强这也是目前大多数人的做法优点是灵活每个epoch模型看到的样本略有不同增强效果随机。第二种是用torchvision.transforms的Compose组合一组增强操作代码更简洁但缺点是不好做到图像和掩码同步变换需要额外处理。第三种是更进阶的做法——在硬盘上离线扩充数据把增强后的图直接保存成新样本这样做的好处是训练时省去了在线增强的计算开销但缺点是磁盘空间消耗大而且每个epoch看到的样本是固定的增强效果打了折扣。我个人的习惯是小数据集用在线增强也就是__getitem__里做变换这样最灵活代码也不复杂。5. 训练流程从损失函数到模型保存数据集就绪模型也搭建完成接下来就是训练环节。训练代码看着不长但里面有不少决定成败的细节比如损失函数怎么配、学习率怎么设、模型怎么保存、怎么判断模型收敛。5.1 损失函数选择BCE Dice混合损失分割任务的损失函数选择很关键我一开始说过了普通交叉熵正负样本不平衡时会出问题。实际项目里我用的比较多的是BCE Loss和Dice Loss的组合。BCE Loss就是二值交叉熵公式上就是每个像素算交叉熵然后求平均。Dice Loss的计算则基于预测和真实标签的重叠程度Dice系数越大表示重叠程度越高而Dice Loss 1 - Dice系数所以数值越小越好。Dice Loss对正负样本不平衡不敏感即使前景只占1%的像素它也能有效驱动模型去学习前景区域。但Dice Loss的梯度在极端情况下可能不太稳定和BCE Loss混合刚好互补。import torch.nn.functional as F def bce_dice_loss(pred, target): bce F.binary_cross_entropy_with_logits(pred, target) smooth 1e-5 pred torch.sigmoid(pred) intersection (pred * target).sum() dice 1 - (2.0 * intersection smooth) / (pred.sum() target.sum() smooth) return bce dice这里pred是模型输出的logits没有经过Sigmoid激活所以用binary_cross_entropy_with_logits而不是binary_cross_entropy这个函数内部帮我们做了Sigmoid数值上更稳定。5.2 优化器与学习率Adam还是SGD优化器方面我的习惯是前期快速收敛用Adam尤其是刚开始调试代码阶段Adam对学习率不那么敏感不太容易出现“loss飞了”的情况。把流程跑通之后再换SGDMomentum做精细调优通常能获得更好的最终精度。学习率的选择Adam我一般初始设置成1e-4或3e-4SGD则用0.01配0.9的Momentum。另外我强烈建议使用学习率衰减策略比如ReduceLROnPlateau它会监控验证集指标连续多少个epoch没有提升时就自动把学习率乘以一个系数。这个策略特别适合训练分割模型因为训练后期loss会陷入平台期学习率不降下来很难继续精调。optimizer torch.optim.Adam(model.parameters(), lr1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemax, factor0.5, patience5, verboseTrue )5.3 训练循环主体一个完整的训练和验证函数接下来说训练循环。一个完整的epoch包含两个阶段训练阶段和验证阶段。训练阶段模型开启model.train()模式会计算梯度并更新参数验证阶段模型开启model.eval()模式不计算梯度不更新参数只做前向推理评估效果。def train_one_epoch(model, dataloader, optimizer, criterion): model.train() running_loss 0.0 for images, masks in dataloader: images images.to(device) masks masks.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, masks) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) return running_loss / len(dataloader.dataset)验证函数和训练函数类似区别在于要加torch.no_grad()上下文管理器提醒PyTorch不要追踪梯度这样能节省大量显存和计算。训练中有些细节值得注意。混合精度训练在较新版本的PyTorch里可以用torch.cuda.amp实现能显著提升训练速度同时显存占用也少很多。如果你用的是不支持自动混合精度的老版本也可以手动做梯度累积或减小batch size来解决显存不足的问题。5.4 模型保存与断点续训模型训练动辄几十个epoch中途断电或者系统重启都是有可能的。我强烈建议不只是保存最后的模型而是定期保存检查点包含模型权重、优化器状态、当前epoch数、最佳指标等所有信息。checkpoint { epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_iou: best_iou, } torch.save(checkpoint, fcheckpoints/unet_epoch_{epoch}.pth)等到要恢复训练时checkpoint torch.load(checkpoints/unet_epoch_20.pth) model.load_state_dict(checkpoint[model_state_dict]) optimizer.load_state_dict(checkpoint[optimizer_state_dict]) epoch checkpoint[epoch]只保存模型权重做推理的话用torch.save(model.state_dict(), model_final.pth)就够了文件体积小很多。我的习惯是训练过程中每隔一定的epoch保存一次完整检查点训练结束后单独再保存一个只含权重的最佳模型文件两种用途分开。6. 训练结果评估与可视化不只是看loss数字训练跑起来了loss也降下来了但你很难从loss数值直接感受到模型分割效果好不好。所以我建议在验证集上做定量评估和可视化这也是判断模型能否真正投入使用的关键环节。6.1 评估指标计算IoU和Dice系数的代码实现IoU的计算方法是用预测正确的正样本数除以预测为正样本或实际为正样本的总数。放在分割任务里就是预测区域和真实区域的交集面积除以并集面积。def compute_iou(pred_mask, true_mask, threshold0.5): pred_mask (pred_mask threshold).astype(int) true_mask (true_mask threshold).astype(int) intersection (pred_mask true_mask).sum() union (pred_mask | true_mask).sum() iou intersection / (union 1e-6) return iouDice系数的计算类似只是它会多算一次交集乘以2作为分子。如果做一个医学影像分割项目Dice系数和IoU都是论文和报告中必须要汇报的核心指标建议两个都算出来。计算指标时需要注意判断预测和标签的数据类型。如果掩码值是0和255需要先转成0和1的布尔数组再进行位运算否则数值对不上算出来的指标会很奇怪。6.2 预测结果可视化原图、掩码、预测图并排看定量指标只能告诉你“效果还行”还是“效果很差”但具体差在哪里哪个区域分割得不好必须靠可视化来直观判断。我的可视化习惯是随机从验证集抽几张图把原图、真实掩码和模型预测结果并排展示。这样一眼就能看出模型是边缘磨糊了、小目标漏检了还是产生了大片的假阳性。下面这段代码展示了如何做预测并保存可视化结果import matplotlib.pyplot as plt def visualize_prediction(model, dataloader, device, save_pathresult.png): model.eval() images, masks next(iter(dataloader)) images, masks images.to(device), masks.to(device) with torch.no_grad(): outputs model(images) preds torch.sigmoid(outputs).cpu().numpy() masks_np masks.cpu().numpy() images_np images.cpu().numpy() fig, axes plt.subplots(3, 3, figsize(12, 12)) for i in range(3): axes[i][0].imshow(images_np[i].transpose(1, 2, 0)) axes[i][0].set_title(Original) axes[i][1].imshow(masks_np[i][0], cmapgray) axes[i][1].set_title(Ground Truth) axes[i][2].imshow(preds[i][0] 0.5, cmapgray) axes[i][2].set_title(Prediction) plt.savefig(save_path)一个实战经验是如果预测结果的边界特别“毛糙”说明模型对边缘的建模能力不足可以考虑增强边界相关的数据或者在损失函数里加入边界损失。如果预测结果整体偏保守目标区域比真实标签小一圈那很可能是Dice Loss的权重过大导致模型不愿意“冒险”输出大范围的前景。7. 常见问题与排查技巧实录训练踩坑是必然的我把自己这几年在U-Net训练过程中遇到的典型问题整理成了一份速查表希望你们遇到问题时不用从头开始踩。现象可能原因解决办法损失值不降学习率过大/过小调整学习率建议先用1e-4试损失值变成NaN学习率过大、数据有异常值调低学习率检查数据归一化预测结果全是黑色Sigmoid后阈值选得不对、数据标签错误检查标签是否全为0调整预测阈值边缘预测模糊输入分辨率太低、模型没有充分训练提高输入图片分辨率增加训练轮数显存不足OOMBatch size太大、图片分辨率过高调小Batch size降分辨率用AMP模型不收敛但loss正常下降过拟合验证集指标低加数据增强、加Dropout、减小模型训练速度极慢没用到GPU、GPU利用率低用nvidia-smi确认GPU在运行调大num_workers训练集准确率高但验证集低过拟合加正则化、早停、增强数据掩码resize后出现中间值用了错误的插值方式改用Image.NEAREST近邻插值7.1 显存不足问题小批量加梯度累积显存不足是新手最常遇到的问题尤其是医学图像分辨率高、模型通道数多的时候更明显。最简单的解决办法是调小batch size但batch size太小会导致BatchNorm统计量不稳定模型训练收敛变慢。一个折中方案是梯度累积也就是每跑几个小batch才更新一次参数模拟出大batch size的效果。比如你实际能跑的最大batch size是4期望的batch size是16那就在4个batch上累积梯度后统一更新accumulation_steps 4 optimizer.zero_grad() for i, (images, masks) in enumerate(dataloader): outputs model(images) loss criterion(outputs, masks) loss loss / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()注意这里loss除以了累积步数这样做的好处是梯度大小不会因为累积而膨胀等效于直接用大batch size计算的平均梯度。7.2 过拟合问题早停和监控指标分割模型参数量大如果数据量只有几百张图片过拟合几乎是必然的。判断过拟合有一个直观的方法训练集loss持续下降但验证集IoU开始下跌或停滞不动。我常用的对策有三个。第一是早停当验证集指标连续多个epoch没有提升时就终止训练避免继续在训练集上“死记硬背”。第二是数据增强这个前面说过了尤其对分割任务效果明显。第三是正则化可以在卷积层后面加Dropout或者DropBlock。在训练时监控指标非常关键我一般会在每个epoch结束打印这样的信息Epoch [20/50] Train Loss: 0.3124 Val Loss: 0.3301 Val IoU: 0.8425 Val Dice: 0.9112观察这些指标的变化趋势比只看训练集loss可靠得多。如果发现验证集IoU连续5个epoch没有增长就果断早停保存最佳模型。7.3 类别不平衡问题加权损失和Focal Loss如果你做的是缺陷检测比如在钢板上找芝麻大小的划痕那正样本像素可能连1%都不到。这种情况下模型很可能快速收敛到把所有像素都预测为背景因为这样loss已经很低了。这时候除了前面说的Dice Loss还可以用Focal Loss。Focal Loss的核心思路是让模型更关注那些难以分类的样本也就是预测概率不高的像素。它的公式我在下面给出实现def focal_loss(pred, target, alpha0.25, gamma2): pred torch.sigmoid(pred) ce_loss F.binary_cross_entropy(pred, target, reductionnone) p_t pred * target (1 - pred) * (1 - target) focal_weight (1 - p_t) ** gamma loss focal_weight * ce_loss return loss.mean()gamma参数控制了难易样本的权重调节力度gamma越大模型越关注难样本。alpha参数用于调节正负样本本身的权重如果前景很少可以把alpha调低负样本权重低让模型更关注正样本。不过要提醒一下Focal Loss的参数调节比较考验经验我用的时候一般会先在验证集上盯几轮效果再决定要不要调整。如果只做二分类且样本不平衡不是特别夸张BCEDice的组合已经够用了不一定非要上Focal Loss。8. 推理部署与后续扩展模型训完才刚走完一半训练出满意的模型这只是整个项目的前半段。在实际应用场景里你得把模型用起来让它接受新图片并给出分割结果。这一步如果事先设计好后面会少很多麻烦。8.1 单张图片推理从加载模型到输出分割掩码推理时不需要重新构建整个训练流程只需要加载模型权重把输入图片处理成模型要求的尺寸前向传播一次再把输出结果还原成可视化掩码即可。这是我的常用推理代码def predict_single_image(model, img_path, device, img_size256, threshold0.5): model.eval() image Image.open(img_path).convert(RGB) original_size image.size image_resized image.resize((img_size, img_size), Image.BILINEAR) input_tensor T.ToTensor()(image_resized).unsqueeze(0).to(device) with torch.no_grad(): output model(input_tensor) pred torch.sigmoid(output) pred (pred threshold).float() pred_np pred.squeeze().cpu().numpy() pred_img Image.fromarray((pred_np * 255).astype(np.uint8), modeL) pred_img pred_img.resize(original_size, Image.NEAREST) return pred_img这段代码里有一个很容易被忽略的细节预测结果resize回原始尺寸时也要用Image.NEAREST原因和之前标签掩码resize相同——最近邻插值不会在类别边界产生伪造的中间像素值。8.2 从U-Net到更多架构后续可以怎么扩展复现完U-Net并跑通自己的数据之后你已经掌握了一套通用的分割项目流程。想在这个基础上前进一步的话方向还挺多的。一个方向是改进模型结构。可以在U-Net里加入注意力机制比如Attention U-Net它在跳跃连接之前加了一个注意力门控能自动抑制无关区域的特征响应对目标区域很小的任务提升明显。另一个实用改进是把普通卷积替换成残差结构的卷积块Residual Block加深网络的同时避免梯度消失。第二个方向是换更强的骨干网络。U-Net的编码器部分可以换成在ImageNet上预训练好的ResNet或EfficientNet利用预训练权重做迁移学习通常能在数据量不大时明显提升精度。第三个方向是数据层面的扩展。除了普通图片你还可以尝试用视频帧序列做时序分割或者把RGB图和深度图组合成多模态输入这些在工业质检和自动驾驶场景都是热门方向。训练这一通下来我对这个流程的感受是U-Net本身不复杂但真正把项目做完整——从标注到训练到评估——有很多细节只有实际操作才能体会到。尤其是数据集制作环节我建议第一次动手时不要追求大数据量找几十张图把整个链路跑通理解每一步发生了什么再考虑扩数据和优化模型。这样你的技术深度和对任务的理解都会扎实很多。
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻