首页 / 资讯中心 / 文章详情

PyTorch复现U-Net:自制数据集训练图像分割模型完整流程

PyTorch复现U-Net:自制数据集训练图像分割模型完整流程 ★ FEATURED ARTICLE
做图像分割的同学大概率绕不开U-Net。不管是医学影像里的器官分割、遥感图像里的地物提取还是工业质检中的缺陷定位U-Net 几乎都是入门必跑的第一个深度学习分割网络。这次我完整复现了一遍 U-Net 的 PyTorch 版本然后自己标注了一套数据集从环境搭建到训练、推理把整个链路跑通了一遍。这篇文章会围绕 U-Net 复现、PyTorch 实现、自制数据集和完整训练流程展开适合手头没有现成数据集、想自己造数据来训练分割模型的人参考。1. U-Net 到底在解决什么问题复现之前先把架构动机搞清楚1.1 分割任务和分类任务本质上是两件事图像分类你只要给出一个标签这张图是猫还是狗。分割不一样它要求对图像里的每一个像素给出类别比如 512×512 的图像就有 26 万个像素需要判断。这决定了分割网络不能像分类网络那样最后接一个全连接层输出概率而是必须做到逐像素输出。早期基于 FCN 的思路已经提出了全卷积的概念但 FCN 有一个明显的硬伤连续下采样之后特征图分辨率越来越小等上采样回原图大小时目标的边缘细节已经丢了很多。表现在结果上就是物体大概出现在该出现的位置但边界糊成一片小目标直接消失。1.2 U-Net 的解法编码器-解码器结构加跳跃连接U-Net 这个名字起的很形象网络结构画出来就是一个 U 形。左边是编码器负责不断下采样提取高语义特征解决这是什么的问题右边是解码器负责逐步上采样恢复分辨率解决这东西在哪的问题。关键的创新在跳跃连接。解码器每一层上采样之后不是直接往下走而是把编码器对应层的特征图拿过来拼接。这一步的效果就像你在写完一篇摘要之后又把你之前画的重点标注直接贴回原文里。浅层特征保留了大量的边缘、纹理细节深层特征有更强的语义信息两者一拼接模型既知道目标是什么类别也知道边界在哪里。这也是 U-Net 结构简单但泛化能力强的根本原因。后来的很多分割模型不管是 SegNet、DeepLab 还是 Transformer 系的分割模型或多或少的思路都受到了这种编码器-解码器 跳跃连接结构的影响。1.3 为什么选 PyTorch 不做成 TensorFlow 版本坦率地说PyTorch 的调试体验比 TensorFlow 早期版本舒服太多动态图机制让你可以在任何一个中间层直接 print 张量形状出问题一眼就能定位。做复现这件事调试效率基本就是一切。另外 PyTorch 生态里数据加载、模型定义、分布式训练都有非常成熟的现成组件代码写起来很直观。如果你也正打算第一次在 PyTorch 里跑通一个完整的视觉模型U-Net 是一个非常好的练手对象它的模块化程度高、代码量不大但足以让你把 Dataset、DataLoader、模型定义、训练循环、评估指标这些环节全部过一遍。2. 从零搭建训练环境conda 虚拟环境、CUDA 匹配与依赖清单2.1 先建一个干净的 conda 虚拟环境别再祸祸 base 环境了我一开始学 PyTorch 的时候偷懒直接在 Anaconda 的 base 环境里 pip install结果某个依赖包把 numpy 版本直接顶掉其他项目全部跑不起来。后来老老实实养成了每次开新项目都创建独立虚拟环境的习惯。conda create -n unet python3.9 -y conda activate unetPython 版本我建议选 3.9 或 3.10因为一些视觉库的预编译 wheel 在 3.11、3.12 上偶尔会有兼容性问题没必要在环境上浪费时间。2.2 PyTorch 安装与 CUDA 版本匹配的关键细节PyTorch 的安装命令建议直接从 PyTorch 官网生成不要自己凭记忆写。先在终端执行nvidia-smi看右上角的 CUDA Version这是当前显卡驱动所支持的最高 CUDA 版本。我的显卡驱动支持到 CUDA 12.1就选择了对应的 PyTorch 版本。安装命令大概长这样pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121装完一定要验证 GPU 是否真的可用import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name())如果torch.cuda.is_available()返回 False先别急着怀疑显卡坏了大概率是 PyTorch 版本和 CUDA 版本不匹配或者装成了 CPU 版本。执行pip list | grep torch看一下 torch 后缀是cu121还是cpu能直接判断问题。没有 NVIDIA GPU 的同学怎么办装 CPU 版本也能跑通整个流程pip install torch torchvision只是训练速度会慢很多建议把输入图像缩小到 256×256batch_size 调成 2 甚至 1流程照样能通只是别指望短时间收敛。2.3 其余依赖清单库名用途numpy数组操作读掩码必备opencv-python图像读写、多边形填充pillow图像基础处理通常和 opencv 一起装matplotlib绘制训练曲线、可视化预测结果tqdm训练进度条没它我真的不想盯着终端labelme标注工具制作分割数据集albumentations数据增强图像和掩码同步变换tensorboard训练指标可视化可选一次性安装pip install numpy opencv-python pillow matplotlib tqdm labelme albumentations tensorboard这套组合拳基本覆盖了从标注到训练的所有环节后面不会再遇到缺少某个库的情况。3. 自制数据集影像采集、Labelme 标注、JSON 转掩码的一条龙流程3.1 数据采集阶段最容易犯的错误以及怎么避免很多初学者一开始就盯着公开数据集下载但 U-Net 复现这个场景下我更推荐自制数据因为自制数据能让你彻底理解标注格式、掩码生成和加载逻辑。我这次选的是墙面裂缝分割这个任务数据容易获取、目标边界明确很适合拿来练手。图像数量方面分割任务起步建议至少 200 张。别贪多关键是质量。采集的时候注意几点目标不要太小至少占画面 5% 以上否则标注和训练都很难受光线、角度尽量多样化这样模型泛化能力才够同一张图别把目标挤成一团训练时网络会分不清实例边界分辨率别太低至少 960×720因为后续可能要 resize 到 512×512。3.2 Labelme 标注多边形是自制分割数据集的标配数据标注工具里Labelme 是最常见的选择尤其适合做不规则形状的语义分割标注。启动方式很简单labelme启动之后打开图片目录用Create Polygons沿着目标边缘打点最后给每个多边形起一个标签名比如crack。保存后每张图片旁边会生成一个同名的 JSON 文件里面记录了这个多边形的所有顶点坐标和标签名称。标注环节的实用建议打点别太密轮廓大致贴合就行训练时数据增强里的旋转、缩放会把小误差自然抹平打点太密反而会让标注文件巨大而且模型学不到更本质的形状特征。我标注了大约 240 张图每张图平均耗时 1 到 2 分钟整个过程大半天就完成了。3.3 JSON 转掩码一次到位直接生成 PNG 掩码Labelme 保存的是 JSON 矢量格式训练时必须把它转成和原图尺寸一致的 PNG 掩码图。这个转换环节是自制数据集最容易出错的地方坐标比例、多边形填充方式、类别值设置任何一个环节出问题都会让模型学到错误答案。import json import os import cv2 import numpy as np from tqdm import tqdm # 标注json路径和输出路径 json_dir data/labelme_json mask_dir data/masks os.makedirs(mask_dir, exist_okTrue) categories {crack: 1} # 背景默认0裂缝类别1 for filename in tqdm(os.listdir(json_dir)): if not filename.endswith(.json): continue json_path os.path.join(json_dir, filename) with open(json_path, r, encodingutf-8) as f: data json.load(f) # 获取原始图像尺寸 img_path os.path.join(os.path.dirname(json_path), data[imagePath]) img cv2.imread(img_path) h, w img.shape[:2] mask np.zeros((h, w), dtypenp.uint8) for shape in data[shapes]: label shape[label] points np.array(shape[points], dtypenp.int32) mask cv2.fillPoly(mask, [points], colorcategories[label]) # 输出同名前缀的png掩码 base_name os.path.splitext(filename)[0] cv2.imwrite(os.path.join(mask_dir, base_name .png), mask)这段脚本的核心思路是每个 JSON 文件对应一张图先建一张全零掩码然后逐多边形填充。背景默认是 0裂缝填充为 1。最后保存为 PNG二值掩码图看起来应该是黑底白纹。这中间有个容易被忽略的细节Labelme 自动生成的 JSON 里imagePath保存的是相对路径或者原始文件名有时候因为文件移动位置会读不到图最好直接用你自己的图片路径替换。数据集目录结构最终长这样dataset/ ├── images/ # 原始图像 ├── masks/ # 对应掩码png ├── train.txt # 训练集图片文件名列表 └── val.txt # 验证集图片文件名列表划分训练集和验证集时我用 8:2 比例做随机划分并且固定了随机种子保证每次复现实验结果一致。这一步强烈建议用脚本生成而不是手动复制否则传错一张图都会让实验结果失真。3.4 数据增强必须图像和掩码同步变换否则白训练分割任务的数据增强和分类任务的差异在于图像和掩码必须用完全相同的变换参数。比如随机翻转图像水平翻转之后掩码也必须水平翻转否则像素对应关系就错了。我用的增强策略如下import albumentations as A train_transform A.Compose([ A.Resize(512, 512), A.RandomCrop(448, 448), A.HorizontalFlip(p0.5), A.Rotate(limit15, p0.5), A.RandomBrightnessContrast(p0.3), ])RandomCrop放在Resize后面是为了先统一尺度再随机裁剪既增加了样本多样性又避免让目标变形太夸张。注意不要用A.Normalize放在这里归一化我更习惯在 Dataset 里单独做因为推理阶段只需要对单张图归一化没必要绑定在增强流程里。4. 数据加载器调试自定义 Dataset 与 DataLoader 的关键细节4.1 Dataset 类的三个方法init、len、getitemU-Net 训练的第一步是写一个继承torch.utils.data.Dataset的自定义类。关键点在于__getitem__方法返回的内容必须是图像张量 掩码张量的元组而且这两个张量必须尺寸一致、空间位置对齐。import os import cv2 import torch from torch.utils.data import Dataset class CrackDataset(Dataset): def __init__(self, img_dir, mask_dir, file_list, transformNone): with open(file_list, r) as f: self.names [line.strip() for line in f.readlines()] self.img_dir img_dir self.mask_dir mask_dir self.transform transform def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] img_path os.path.join(self.img_dir, name .jpg) mask_path os.path.join(self.mask_dir, name .png) image cv2.imread(img_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) # 掩码里非0的都置为1保证类别值是0和1 mask[mask 0] 1 if self.transform: augmented self.transform(imageimage, maskmask) image augmented[image] mask augmented[mask] # 转成tensor并调整通道顺序 HWC - CHW image torch.from_numpy(image).permute(2, 0, 1).float() / 255.0 mask torch.from_numpy(mask).long() return image, mask掩码读取时一定要用cv2.IMREAD_GRAYSCALE否则读进来是三个通道和图像的通道数不一致后面算 loss 的时候会直接报错。4.2 图像和掩码的同步变换是 Dataset 里最容易翻车的地方如果不用 albumentations而是自己写 transform那就必须手动保证图像和掩码用同一组随机种子。很多人在这一步偷懒结果训练时发现 loss 怎么都不降一检查才发现掩码和图像对不上了。albumentations 的优势就在这里它接受image和mask两个参数内部会保证随机变换参数完全一致还会自动把掩码的插值方式锁定为最近邻不会因为旋转、缩放产生模糊的类别边界。4.3 DataLoader 参数怎么配才能不爆内存又不卡训练from torch.utils.data import DataLoader train_dataset CrackDataset(dataset/images, dataset/masks, dataset/train.txt, transformtrain_transform) val_dataset CrackDataset(dataset/images, dataset/masks, dataset/val.txt, transformval_transform) train_loader DataLoader(train_dataset, batch_size8, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size8, shuffleFalse, num_workers4, pin_memoryTrue)几个参数的经验值batch_size显存 8G 以下建议 416G 可以 8num_workersLinux 下直接给 4 到 8Windows 下如果报错就设 0这是无数人踩过的坑pin_memory配合 GPU 训练建议开启能减少数据从内存拷贝到显存的时间。写完之后先用一个简单的循环测试 data loader 能不能正常吐数据别直接开始训练for images, masks in train_loader: print(images.shape, masks.shape) break期望输出是torch.Size([8, 3, 448, 448]) torch.Size([8, 448, 448])。如果 mask 多了一个通道或者尺寸和图像对不上趁早回去改代码不要带着 bug 硬训。5. 核心网络实现U-Net 的编码器、解码器与跳跃连接5.1 DoubleConv、Down、Up 三个模块拆解U-Net 的核心模块可以抽象成三种结构。第一个是DoubleConv包含两个 3×3 卷积、BatchNorm 和 ReLU 激活。卷积都设置padding1这样卷积前后特征图尺寸不变方便跳跃连接进行拼接。第二个是Down先做一次 2×2 最大池化下采样然后接一个DoubleConv。每经过一个 Down特征图尺寸减半通道数翻倍。第三个是Up先通过转置卷积把特征图上采样一倍然后把编码器对应层的特征在通道维度上拼接起来最后接一个DoubleConv。5.2 完整可运行的 U-Net 模型代码import torch import torch.nn as nn class DoubleConv(nn.Module): 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 Down(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.pool nn.MaxPool2d(2) self.conv DoubleConv(in_ch, out_ch) def forward(self, x): return self.conv(self.pool(x)) class Up(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.up nn.ConvTranspose2d(in_ch, in_ch // 2, kernel_size2, stride2) 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 nn.functional.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x torch.cat([x2, x1], dim1) return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels3, num_classes1, features[64, 128, 256, 512, 1024]): super().__init__() self.inc DoubleConv(in_channels, features[0]) self.down1 Down(features[0], features[1]) self.down2 Down(features[1], features[2]) self.down3 Down(features[2], features[3]) self.down4 Down(features[3], features[4]) self.up1 Up(features[4], features[3]) self.up2 Up(features[3], features[2]) self.up3 Up(features[2], features[1]) self.up4 Up(features[1], features[0]) self.outc nn.Conv2d(features[0], num_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)注意Up模块里的concat之前我用pad对 x1 做了中心补边处理。理论上只要输入尺寸是 2 的倍数每一层编码器和解码器的特征图尺寸应该恰好匹配但保险起见还是加上这个处理否则输入长宽不是 2 的幂次时直接torch.cat会报错。5.3 输出通道到底填 1 还是填 2二分类分割比如裂缝 vs 背景num_classes可以填 1输出层是单通道后面接BCEWithLogitsLoss直接在模型输出上做 sigmoid。但如果你习惯了多分类框架也可以把num_classes填 2用CrossEntropyLoss两种写法都能跑通。我的建议是二分类场景直接用单通道 BCE少一个通道的显存开销不说后处理也简单。多类别分割场景再用num_classesN 交叉熵。6. 训练流程与损失函数选择让模型真正把前景分出来6.1 损失函数对比为什么 BCE 在裂缝分割上不够好用分类任务里常用的交叉熵损失在分割任务里能直接用但一个典型问题是当背景像素远多于前景像素时模型会倾向于把所有像素都预测成背景因为这样 loss 也低。裂缝这种目标在整幅图里占比有时候只有几个百分点用纯 BCE 训出来的网络很可能出现全黑掩码的情况。损失函数优点缺点适用场景BCEWithLogitsLoss实现简单收敛稳类别不平衡时偏向多数类前景背景比例接近Dice Loss直接优化 IoU 指标天然处理不平衡训练初期可能出现梯度不稳定前景占比小的分割任务BCE Dice 组合兼顾逐像素精度和区域重叠度需要调两个损失的权重推荐默认方案组合损失是我现在最常用的方案代码里可以直接在训练脚本中叠加两个损失import torch.nn.functional as F def dice_loss(pred, target, smooth1.0): pred torch.sigmoid(pred) pred pred.contiguous().view(-1) target target.contiguous().view(-1) intersection (pred * target).sum() return 1 - (2.0 * intersection smooth) / (pred.sum() target.sum() smooth) def combined_loss(pred, target): bce F.binary_cross_entropy_with_logits(pred, target.float()) dice dice_loss(pred, target) return bce dicetarget用float()的原因是 BCE 需要浮点类型的标签而dice_loss里sigmoid(pred)和target做乘法也需要浮点。6.2 训练循环主结构核心几步别搞乱训练循环本身没有太多花活但每一步的先后顺序必须清楚前向传播算 loss反向传播算梯度优化器更新参数。这三步的顺序反了或者漏了某一步模型都会处于假装在训练的状态。model UNet(in_channels3, num_classes1).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-4, weight_decay1e-5) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, patience5, factor0.5) best_iou 0.0 for epoch in range(epochs): model.train() train_loss 0.0 for images, masks in train_loader: images, masks images.to(device), masks.to(device) pred model(images) # 前向传播 loss combined_loss(pred, masks) # 计算损失 optimizer.zero_grad() # 梯度清零 loss.backward() # 反向传播 optimizer.step() # 更新参数 train_loss loss.item() * images.size(0) # 验证阶段 model.eval() val_iou 0.0 with torch.no_grad(): for images, masks in val_loader: images, masks images.to(device), masks.to(device) pred model(images) val_loss combined_loss(pred, masks) val_iou compute_iou(pred, masks) * images.size(0) # 保存验证集上IoU最高的权重 if val_iou best_iou: best_iou val_iou torch.save(model.state_dict(), best_model.pth)写训练脚本时model.train()和model.eval()千万别省。虽然你的网络里有 BatchNorm 层这两个方法会直接影响 BatchNorm 在训练和推理阶段的统计方式忘了切换可能导致验证指标忽高忽低。6.3 评估指标只用 loss 看模型好坏是不够的训练过程中必须计算验证集上的 IoUIntersection over Union。IoU 的含义是两个区域交集除以并集数值越接近 1 越好。计算公式如下def compute_iou(pred, target, threshold0.5): pred torch.sigmoid(pred) pred (pred threshold).float() target target.float() intersection (pred * target).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target.sum(dim(2, 3)) - intersection iou (intersection 1e-6) / (union 1e-6) return iou.mean().item()我在训练时用了ReduceLROnPlateau调度器当验证集 loss 连续 5 个 epoch 不再下降时学习率就减半。这个策略比固定学习率省心得多不用手动反复试。6.4 训练曲线怎么看正常的和异常的信号正常训练中train loss 和 val loss 应该逐步下降val IoU 稳步上升。如果出现这些信号就要警惕train loss 下降但 val loss 不降过拟合增强数据增强强度、加 dropout、调大 weight_decaytrain loss 和 val loss 都纹丝不动学习率太低或者数据加载环节就出问题了train loss 正常下降但 val IoU 一直是 0很可能是掩码标注有误或者数据划分泄漏先可视化几张验证集样本看看。我这次训练 80 个 epoch 后验证集 IoU 到了 0.78 左右单卡 2080Ti 上每个 epoch 大概 40 秒。裂缝分割这种结构比较简单的小目标任务U-Net 在小数据集上就能达到不错的效果。7. 推理与效果评估加载权重输出分割结果7.1 推理脚本预训练权重加载与结果保存训练完成后最终交付的是一个训练好的权重文件best_model.pth。推理阶段要做的事情是加载这个权重对新图片做和训练时完全一致的预处理前向传播得到输出再通过 sigmoid 和阈值得到最终掩码。import cv2 import torch import numpy as np # 加载模型 device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, num_classes1).to(device) model.load_state_dict(torch.load(best_model.pth, map_locationdevice)) model.eval() def predict(image_path, model, device, threshold0.5): image cv2.imread(image_path) image_rgb cv2.cvtColor(image, cv2.COLOR_BGR2RGB) h, w image.shape[:2] # 和训练时保持一致的resize resized cv2.resize(image_rgb, (448, 448), interpolationcv2.INTER_LINEAR) tensor torch.from_numpy(resized).permute(2, 0, 1).float().unsqueeze(0) / 255.0 tensor tensor.to(device) with torch.no_grad(): output model(tensor) prob torch.sigmoid(output).cpu().numpy()[0, 0] binary_mask (prob threshold).astype(np.uint8) # 恢复到原图尺寸 mask cv2.resize(binary_mask, (w, h), interpolationcv2.INTER_NEAREST) return mask # 推理并可视化 mask predict(test.jpg, model, device) overlay cv2.imread(test.jpg) overlay[mask 1] (0, 0, 255) # 在原图上标红 cv2.imwrite(result_overlay.jpg, overlay)推理时的注意事项resize 回原尺寸时插值方式保持INTER_NEAREST和训练时保持一致这样不会因为插值引入伪类别。同时要记得在推理前调用model.eval()否则 BatchNorm 层的行为会和训练时不一致导致输出层 prob 明显异常。7.2 结果可视化时重点看哪些位置训练结束后我习惯随机抽几张验证集图片放在一张大图里对比原图、真实掩码、预测掩码、叠加图。重点观察三类位置目标的边缘是否干净如果边缘锯齿严重说明上采样恢复的细节不够可以尝试增大 U-Net 的通道数小目标裂缝是否被保留如果小裂缝直接没检出来大概率是下采样过深导致细节损失可以考虑减少下采样次数背景是否有大面积误检如果有检查标签里有没有漏标的目标漏标会让模型把同一区域既当正样本又当负样本。8. 复现过程中的高频坑与排查思路8.1 掩码读取通道数不对的连锁反应用cv2.imread(mask_path)读掩码默认读成 3 通道 BGR 图像。虽然灰色图像的三个通道数值相同看起来不报错但当你把 mask 转成long类型 tensor 后和单通道预测结果计算 loss 时会直接报 shape mismatch。正确做法是cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)。8.2 输入图像尺寸必须是 2 的倍数吗U-Net 内部一共有 4 次下采样和 4 次上采样。虽然我在代码里加了 padding 补边逻辑正常情况下任意尺寸都能跑通但尺寸不是 2 的倍数时上采样后的特征图和跳跃连接的特征图会有一两个像素的偏差补边后虽然能跑但会有微小的空间对齐误差。最稳妥的输入尺寸是 512×512、448×448、384×384 这类 2 的幂次相关尺寸。8.3 训练时显存不足怎么破先从最简单的batch_size开始降从 8 降到 4 再降到 2。如果还是不够把输入尺寸从 512 降到 384。还不行就用混合精度训练pip install torch.cuda.ampfrom torch.cuda.amp import autocast, GradScaler scaler GradScaler() with autocast(): pred model(images) loss combined_loss(pred, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度在 30 系以上显卡上几乎无损地减少一半显存占用速度也有提升是我现在的默认配置。8.4 训练 loss 一直不降的排查链路如果你的 loss 在 10 个 epoch 内完全不动不要动模型结构先按这个顺序排查确认数据加载正确随机打印 5 张图同时输出图像和掩码肉眼看掩码和图像是否对齐检查标签值域二分类掩码必须只包含 0 和 1如果误标成 255BCE loss 会在数值上表现异常检查输出层通道数单通道输出配 BCE多通道输出配 CrossEntropy混用必出事检查学习率1e-4 是 U-Net 的常用起点如果设成 1e-2loss 会一直震荡甚至爆掉检查 loss 是否返回的是一个标量而不是 tensor 数组。8.5 自制数据集标注中的几个经验教训标注阶段我踩过的坑主要在两个地方。第一个是多边形坐标Labelme 输出的 points 是 float 类型直接传给cv2.fillPoly会报警告甚至填充错误必须转成np.int32。第二个是类别值多类别分割时类别编号必须从 0 开始连续递增比如 0、1、2跳号会让 softmax 输出通道数量对不上loss 计算直接抛错。最后说一点我自己的体会整个 U-Net 复现下来最大的感受是网络本身并不复杂真正花时间的地方在数据和工程细节上。数据标注占了大约 60% 的时间但是数据质量直接决定了模型上限这个投入非常值得。环境配置、踩坑排查这些工作在第一次做的时候会觉得繁琐但把这些经验沉淀下来后面再跑其他分割模型就会顺手很多。如果你也准备复现 U-Net我的建议是先别急着改网络结构老老实实把自制数据的链路跑通再在这个基础上逐步优化效果会比你想象的来得快。
阅读完成 · 觉得有帮助?
咨询建站