ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

UNet训练自定义数据集完整指南:从数据预处理到源码实现

UNet训练自定义数据集完整指南:从数据预处理到源码实现 简介一套面向深度学习初学者和医学图像分析开发者的UNet语义分割完整源码包聚焦皮肤病分割场景系统覆盖数据标注、数据处理、数据划分以及模型训练与推理全流程解决训练自定义数据集时流程不完整、代码不闭环的痛点。资源共2000个文件以1719张PNG和256张JPG图像为载体包含18个Python脚本、5个XML标注文件以及说明文档压缩包约360.85MB其中图像文件用于训练验证Python脚本实现网络搭建和训练流程XML与文本文件提供标注与使用参考。资源内置皮肤病分割数据集和对应训练权重可直接复现皮肤病灶分割效果也支持替换自定义数据集进行迁移学习。随包附有数据标注工具和训练教程帮助用户快速掌握从原始数据到模型部署的完整链路。目前已有1903人学习下载适合需要快速搭建UNet分割项目的开发者和研究者。从零跑通UNet训练自己的数据集一份能直接抄作业的完整源码笔记我见过太多人卡在同一个地方UNet源码跑Carvana或者CamVid数据集的时候顺顺利利一换到自己的数据集报错、训练不收敛、分割结果一塌糊涂。这个问题的根源几乎不在模型本身而是从数据预处理到训练循环的每一个环节都有隐性前提没有被说透。这篇文章就是我把自己踩过的坑、验证过的方案整合出来的一份完整说明适用于把UNet应用到自己采集的图像分割数据上无论是做医学影像、遥感分割还是工业质检核心流程是通用的。我会沿一条主线走从数据准备、模型结构原理到完整训练源码、训练调试再到推理和模型改进全程给出可以直接复现的代码和思路。在动笔之前需要说明一点这里提供的源码不是一个只能跑通demo的玩具它包含了我实际项目里用到的容错逻辑、指标计算和显存优化是一个可以直接改装进自己项目的骨架。1. 数据准备这个环节决定了训练成败的80%训练自己的数据集最容易翻车的地方不在训练代码而是数据。我调试过不少读者的私信几乎一半的问题根源是图像和对应的mask没有对齐、标签通道数不对、数据增强把图像和标签弄错位。所以第一步要考虑的不是怎么写UNet而是怎么把你的数据集整理成UNet可以直接吃的样子。1.1 图像与标签的配对逻辑先约定一个最简单的文件组织方式这符合大多数自行采集数据的习惯dataset/ ├── images/ │ ├── img_001.jpg │ ├── img_002.jpg │ └── ... └── masks/ ├── img_001.png ├── img_002.png └── ...注意三个细节第一条是图像和标签的前缀必须完全一致这样代码里才能通过文件名建立配对。第二条原始图像推荐用jpg或png都行但标签mask一定保存成png格式因为png是无损压缩能保证类别像素值的精确性jpg是有损压缩会改变边缘像素值如果标签里恰好有需要精确读取的类别编号会带来直接的错误。第三条如果你自己用LabelMe这类工具标注导出的时候会同时包含json和png只需要把png提取出来放到masks文件夹就行。注意如果你的掩码是PNG且是调色板模式P模式别忘了在加载时转换成RGB或灰度索引否则某些框架会直接报通道数不匹配的错误。1.2 灰度图与三通道掩码的取舍很多标注工具导出的mask是三通道的彩色图但这里有一个可以直接提升训练效率的做法在数据加载时把掩码转换为单通道灰度图。语义分割的标签一般有两种存在形式——单通道索引图每个像素值是这个像素所属类别的编号和三通道彩色展示图每个类用不同颜色标出。UNet训练时如果你做的是二分类任务标签可以存成值是0和1的单通道图如果做多分类比如背景、建筑、道路三类标签就应该存成0、1、2这样的单通道索引图配合CrossEntropyLoss一起使用。我这里给出的完整源码将以二分类目标/背景作为默认场景因为这是最普遍的需求也是把UNet跑通自己数据最快的一条路径。1.3 数据增强的具体取舍数据增强是深度学习里特别容易“好心办坏事”的环节。对图像做随机旋转、水平翻转、颜色抖动这些操作时只要忘记对掩码做同步变换损失就会变成NaN训练彻底废掉。我实际使用的增强策略是这样的水平/垂直翻转5成概率图像和mask同步翻转。随机旋转90度适合大多数类别分布相对均匀的场景。随机裁剪到固定尺寸例如512x512缓解显存压力。不推荐随意加入亮度对比度抖动因为很多分割任务对颜色本身是敏感的加了反而掉点。实现同步翻转很容易用numpy的np.flip同时处理图像和掩码即可或者直接沿用下面第4节数据加载器里的写法。如果你要求增强更丰富可以顺手引入albumentations它的核心特性就是同一套变换同时作用在image和mask上用起来最不容易出错。2. UNet结构原理与适配自己数据的修改点UNet是2015年提出的语义分割网络直到今天依旧是绝大多数分割项目的第一选择。它名字里的U来自轮廓左侧是编码器下采样提取特征右侧是解码器上采样恢复分辨率中间通过跳跃连接skip connection把编码器的细节特征拼接到解码器对应层。为什么要这样做因为下采样会丢失空间信息而分割任务恰恰需要像素级的定位精度跳跃连接相当于把宏观语义和微观纹理拼接在一起让模型既能看懂是什么又能定位在哪里。2.1 标准UNet的关键参数与计算量标准的UNet通常做4次下采样和4次上采样。输入尺寸如果是512x512特征图尺寸依次是512、256、128、64、32。这里务必注意由于下采样过程中出现了2的4次方16倍缩放输入图像的长宽必须能被16整除否则转置卷积恢复尺寸时容易遇到形状不匹配的报错。UNet在每层使用两个卷积通常3x3中间夹ReLU激活函数下采样用stride2的最大池化或者stride2的卷积实现上采样采用转置卷积。这部分的源码我一般直接写成可配置的方便在不同数据集上调整初始通道数。2.2 输入输出通道的修改方法适配自己的数据集第一件要改的事就是通道数。输入通道in_channels取决于你的图像。常规RGB图为3如果是灰度医学图像则改为1。输出通道out_channels取决于你的分割目标数量。二分类时输出1个通道配合Sigmoid BCEWithLogitsLoss多分类时输出类别数N类配合Softmax CrossEntropyLoss。我在这篇博文给出的源码中用二分类场景out_channels1作为默认配置并预留了多分类的开关。如果你想做3分类只需要把out_channels改成3并把损失函数换成CrossEntropyLoss即可其余结构不用动。3. 训练时的关键超参从理论到实测的推荐值训练深度学习模型完全照搬别人论文里的学习率十有八九会翻车。我在这里分享一组在自定义数据集上实测表现稳定的配置同时解释每个参数背后的考虑让你遇到问题时有方向可调而不是盲目试。超参数推荐值说明输入尺寸512x512显存不够可降到256但精度会掉Batch Size4~8根据显存动态调整初始学习率1e-4Adam优化器下的常用起点学习率衰减每10轮乘以0.9让损失稳定下降Epoch数50~100看验证集IoU是否继续上涨优化器Adam自适应学习率收敛快损失函数BCE Dice解决正负样本不平衡关于损失函数多说两句如果你的图片里目标占的面积很小比如病变区域只占整张图的几个百分点纯BCE会让模型变成全预测背景的懒惰模式因为那样损失就很小。把Dice Loss加上去之后模型不得不在预测目标区域时做出实质性努力。Dice Loss的公式可以理解为交集乘2除以并集数值越大越好转换成损失就是1减去这个值。我的经验是BCE和Dice Loss按1:1加权对大多数任务都有稳健提升。4. 完整可运行的UNet训练源码这一部分给出核心源码如果你有自己的工程习惯可以直接把数据加载和训练循环抽出去复用。为了保证复现性我会先把核心组件的逻辑讲清楚再给出拼接方式。4.1 UNet模型定义下面是一个紧凑版UNet实现保留了标准结构同时把通道配置做成参数方便扩展。import torch import torch.nn as nn class DoubleConv(nn.Module): 两次卷积 BN ReLU def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels3, out_channels1, features[64, 128, 256, 512, 1024]): super().__init__() self.pool nn.MaxPool2d(2) self.encoder1 DoubleConv(in_channels, features[0]) self.encoder2 DoubleConv(features[0], features[1]) self.encoder3 DoubleConv(features[1], features[2]) self.encoder4 DoubleConv(features[2], features[3]) self.bottleneck DoubleConv(features[3], features[4]) self.up4 nn.ConvTranspose2d(features[4], features[3], 2, 2) self.decoder4 DoubleConv(features[3] features[3], features[3]) self.up3 nn.ConvTranspose2d(features[3], features[2], 2, 2) self.decoder3 DoubleConv(features[2] features[2], features[2]) self.up2 nn.ConvTranspose2d(features[2], features[1], 2, 2) self.decoder2 DoubleConv(features[1] features[1], features[1]) self.up1 nn.ConvTranspose2d(features[1], features[0], 2, 2) self.decoder1 DoubleConv(features[0] features[0], features[0]) self.final nn.Conv2d(features[0], out_channels, 1) def forward(self, x): e1 self.encoder1(x) e2 self.encoder2(self.pool(e1)) e3 self.encoder3(self.pool(e2)) e4 self.encoder4(self.pool(e3)) b self.bottleneck(self.pool(e4)) d4 self.decoder4(torch.cat([self.up4(b), e4], dim1)) d3 self.decoder3(torch.cat([self.up3(d4), e3], dim1)) d2 self.decoder2(torch.cat([self.up2(d3), e2], dim1)) d1 self.decoder1(torch.cat([self.up1(d2), e1], dim1)) return self.final(d1)这段代码里有一个容易被忽略的细节跳跃连接拼接时用的torch.cat是在通道维度dim1上执行的因为特征图张量的形状是[B, C, H, W]。如果你自己手写UNet很容易在这里把维度搞错导致RuntimeError。4.2 数据集加载与预处理数据加载器是整个工程的水电煤我做了一个简化版不依赖第三方加强库但逻辑完整可靠import os import cv2 import numpy as np import torch from torch.utils.data import Dataset class SegmentationDataset(Dataset): def __init__(self, img_dir, mask_dir, image_size(512, 512), trainTrue): self.img_dir img_dir self.mask_dir mask_dir self.image_size image_size self.train train self.ids [f.split(.)[0] for f in os.listdir(img_dir)] def __len__(self): return len(self.ids) def __getitem__(self, idx): img_id self.ids[idx] img_path os.path.join(self.img_dir, img_id .jpg) mask_path os.path.join(self.mask_dir, img_id .png) image cv2.imread(img_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) # 统一尺寸 image cv2.resize(image, self.image_size, interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, self.image_size, interpolationcv2.INTER_NEAREST) # 归一化到[0,1]区间 mask (mask 127).astype(np.float32) # 简单数据增强同步水平翻转 if self.train and np.random.rand() 0.5: image cv2.flip(image, 1) mask cv2.flip(mask, 1) image torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0 mask torch.from_numpy(mask).unsqueeze(0).float() return image, mask需要特别关注一个隐藏很深的坑cv2.resize处理mask时必须把插值方式设为INTER_NEAREST也就是最近邻插值。如果用默认的线性插值原本值是0和1的标签会被插值成0.4、0.7这类中间值导致训练时损失计算混乱。这个问题我几乎每隔一段时间就会在读者报错里看到一次。4.3 训练循环与指标监控训练循环本身不复杂但我想强调三个容易忽略的细节模型模式切换、梯度清零、验证集评估。import torch.nn as nn from torch.utils.data import DataLoader from tqdm import tqdm def dice_loss(pred, target, smooth1e-6): pred torch.sigmoid(pred) intersection (pred * target).sum() return 1 - (2.0 * intersection smooth) / (pred.sum() target.sum() smooth) def bce_dice_loss(pred, target): bce nn.BCEWithLogitsLoss()(pred, target) dice dice_loss(pred, target) return bce dice def iou_score(pred, target, threshold0.5): pred (torch.sigmoid(pred) threshold).float() intersection (pred * target).sum() union pred.sum() target.sum() - intersection return (intersection 1e-6) / (union 1e-6) def train_one_epoch(model, loader, optimizer, device): model.train() total_loss 0.0 for images, masks in tqdm(loader): images, masks images.to(device), masks.to(device) optimizer.zero_grad() outputs model(images) loss bce_dice_loss(outputs, masks) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader) torch.no_grad() def validate(model, loader, device): model.eval() total_iou 0.0 for images, masks in loader: images, masks images.to(device), masks.to(device) outputs model(images) total_iou iou_score(outputs, masks).item() return total_iou / len(loader)在主训练脚本里把模型、优化器、数据加载器串起来同时加入学习率衰减和模型保存逻辑。模型保存的时机很有讲究我一般只看验证集的IoU当它比之前最好的成绩更高时才保存一次权重。单纯靠训练集损失下降来保存很容易存下一个过拟合模型。5. 训练过程中的三大拦路虎显存、过拟合、不收敛训练的坑基本集中在显存溢出、过拟合、损失不下降这三类问题上下面逐个说清。5.1 显存不足菜鸟容易忽视的层数和通道数耦合默认的UNet结构在512x512输入下一个batch假设batch size为4大约会吃掉8~12GB显存。如果你的显卡只有6GB第一步把batch size降为2输入尺寸降到256x256或者把features通道数降为[32, 64, 128, 256, 512]显存占用会大幅下降。还有一个技巧是开启混合精度训练AMP在PyTorch中只需要在训练循环中加入torch.cuda.amp.autocast()和GradScaler()两行逻辑显存占用和训练时间都能得到明显改善精度损失几乎可以忽略。对于显存小于8GB的用户这是性价比最高的一招。5.2 过拟合医学小数据集最典型的症状当你的训练集只有一两百张图时过拟合几乎是必然趋势典型症状是训练损失不断下降但验证IoU停滞甚至倒退。我的应对策略有三个按优先级排序第一提高数据增强强度包括随机旋转、缩放、弹性形变等。第二把Dropout加在瓶颈层的后面注意UNet原始结构是没有Dropout的这部分需要自己加。第三采用早停机制连续10轮验证集IoU不上升就停止训练并回滚到最佳权重。5.3 损失不下降静的代码和活的调试损失完全不下降或者直接变成NaN大概率是数据出了问题。先用可视化几对图像和mask确定mask内容不是全黑。接着检查归一化方式图像归一化到0~1还是0~255会影响初始损失的大小但并不影响是否正确收敛。最容易被忽视的是学习率如果初始学习率过高交叉熵和Dice组合的损失可能直接炸掉。根据我的经验Adam优化器搭配1e-4的学习率在UNet上基本不会出大问题如果换成了SGD就需要把学习率调到0.01再配动量和学习率衰减。6. 推理与效果评估训练好之后如何判断模型到底行不行训练完成之后还要做一套完整的推理和可视化流程才能算出有说服力的效果指标并把分割结果直接导出应用到实际场景里。6.1 推理单张图像的完整流程def predict_image(model, image_path, device, size(512, 512), threshold0.5): model.eval() image cv2.imread(image_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) original_size image.shape[:2] image cv2.resize(image, size, interpolationcv2.INTER_LINEAR) image_tensor torch.from_numpy(image.transpose(2, 0, 1)).float().unsqueeze(0) / 255.0 image_tensor image_tensor.to(device) with torch.no_grad(): output model(image_tensor) prob torch.sigmoid(output).cpu().numpy().squeeze(0).squeeze(0) mask (prob threshold).astype(np.uint8) * 255 mask cv2.resize(mask, (original_size[1], original_size[0]), interpolationcv2.INTER_NEAREST) return mask注意这里的最后一个cv2.resize我特意把分割结果恢复到了原图尺寸这样方便叠加显示和后续计算面积。阈值0.5是一般默认值如果你发现预测结果中目标区域偏大或偏小可以把阈值在0.3~0.7之间调整这是一个成本几乎为零的调优手段。6.2 用IoU、Dice系数和像素精度客观评价效果语义分割最常用的三个评价指标是IoU、Dice系数和像素精度Pixel Accuracy。IoU是预测区域和真实区域的交集除以并集Dice系数与IoU本质相通但有轻微权重偏向数值更接近且对小目标更敏感像素精度是预测正确的像素占总像素的比例。写项目报告时建议三个指标都算出来因为单一指标可能会掩盖问题。我在代码里已经实现了IoU的计算Dice系数可以直接复用上面的dice_loss函数只需要用1 - dice_loss(...)即可得到。7. UNet的主观局限与4条进阶改进方向如果你的数据集难度比较大直接使用标准UNet的效果可能不够理想。结合我自己的实际使用感受给出以下四条实证有效的改进方向。7.1 注意力机制让网络自动关注关键区域在跳跃连接处或解码器部分加入Attention Gate可以让模型自动抑制背景区域的响应、突出目标区域的响应。最常见的是在每次上采样之后、和编码器特征拼接之前加一个注意力系数计算。代码量不大但对小目标分割的提升肉眼可见。7.2 残差连接与更深的编码器缓解深层网络退化问题把编码器部分替换成ResNet34或者ResNet50的预训练权重属于时下常用的预训练编码器UNet解码器方案适合数据量中等的场景因为预训练权重已经学到了大量通用视觉特征。替换之后解码器部分不需要大改只要把跳跃连接的通道数对齐即可。7.3 多尺度输入让网络兼顾全局语义和局部细节使用类似DeepLab的ASPP空洞空间金字塔池化模块替换瓶颈层可以在不降低分辨率的情况下获得多个感受野的信息。对场景分割这类需要同时辨识大物体和小物体的任务这个改动效果稳定。具体实现中用不同空洞率的3x3卷积并行提取特征最后拼接融合成瓶颈输出。7.4 把UNet换成分割一切类模型什么时候值得考虑这件事近几年基础分割模型在公开场景上表现很强但这类模型的部署门槛和显存要求都比较高而且针对小众数据仍然需要微调。以我的经验先跑通UNet得到一套基准指标再根据指标短板决定是否切换到更强的模型这是一条性价比最高的技术路线。如果你连UNet的基准都没跑出来直接换更重的模型只会让问题更难定位。8. 一个容易让新手崩溃却极少被提及的问题尺寸对齐尺寸对齐是UNet训练中极其隐蔽但危害巨大的坑。前面提到输入尺寸需要能被16整除是因为4次下采样后的特征图长宽是输入的1/16。如果输入是512最后一层特征图大小是32如果输入是513特征图大小是32.0625这在卷积计算时会发生维度取整的错位。而转置卷积上采样时拼接操作又要求解码器当前特征图和跳跃连接传过来的特征图长宽完全一致一旦不一致就会直接报错。为了彻底躲避这类问题我在数据加载器里做了一次保险处理先把图像缩放到短边为512或256然后再中心裁剪到固定尺寸512x512。这样配合自定义的Dataset类无论你喂进来的原始图是多大最终进入模型的数据尺寸都是一致的训练稳定性会高很多。其实写到这里UNet训练自己数据集的完整链路已经讲完了。最后分享一个我自己的习惯每次启动一组新数据的训练我会先用5张图跑一个冒烟测试只训练1个epoch确认损失能下降、验证脚本能跑通再开启完整的50~100轮训练。这个习惯曾经帮我节省过大量定位bug的时间建议你也试试。如果你在复现过程中遇到数据集或代码层面的报错把完整的错误信息、数据目录结构、输入图像尺寸发出来我可以帮你定位是哪个环节的问题。这套源码我后续也会持续更新加入更多稳健性处理和高级改进的实现。本文还有配套的精品资源点击获取
返回列表