ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch UNet语义分割实战:从数据准备到模型训练的核心技巧

PyTorch UNet语义分割实战:从数据准备到模型训练的核心技巧 简介基于PyTorch搭建U-Net网络并训练自定义数据集的完整工程包面向图像分割方向的研究人员与开发者尤其适合医学影像分割、卫星图像分割等场景的入门与迁移学习。工程覆盖了U-Net核心组件——编码器、解码器、跳跃连接以及数据预处理、数据增强、损失函数与优化器配置并提供了可运行的训练、测试与评估脚本方便直接参考或改造。压缩包共27个文件以7个Python源码为主包括网络定义、训练、验证、推理等模块另附有5份markdown说明文档、5张示例分割结果图以及适合PyCharm的工程配置文件整包仅602KB体积小巧便于快速部署。当前已有773人学习下载。借助该资源可清晰理解PyTorch实现U-Net的完整流程并在自己的标注数据集上完成训练与预测是搭建图像分割模型、撰写相关博文或课程设计的高性价比参考。1. 从拿到 train.py 到自己训练的完整链路UNet 其实比你想的更务实在 PyTorch 里搭建自己的 UNet 网络、训练自己的数据集这事听起来像是“解压一个 zip 跑通 demo”就能交差真正动手后你才会发现卡点全在数据准备和训练细节上。UNet 是语义分割里结构最透明的网络之一编码器往下压分辨率、解码器往上升分辨率、跳跃连接把编码器细节直接传给解码器。它的价值在于数据量不大时也能训得像样几百张标注图就能出可用的轮廓这一点比 Transformer 系分割模型动辄要上万张数据省心得多。下面不是复刻某个仓库的 readme而是把“从零到能交差”的路径拆给你数据怎么组织、网络怎么写、哪些参数值得调、哪些坑我替你踩过。适合正在用手头图片训练分割模型的人读完能少走两晚上弯路。2. 数据准备把任意图片变成 UNet 能直接吃的数据集2.1 数据集目录与掩码转换彩色标注必须转成单通道索引图UNet 训练的输入输出本质上是一对图输入是原始图像输出是掩码mask掩码里每个像素的值代表类别编号。常见做法是准备两个目录比如 images 和 masks文件名一一对应。你可以用任何标注工具Labelme、Label Studio、图片编辑软件都行但导出的标注图大多是 RGB 彩色图例如背景是黑色、目标物体是白色。RGB 图不能直接送进 UNet 当监督信号二分类任务里模型输出的是一个单通道概率图你的监督标签也必须是单通道像素值 0 表示背景、1 表示前景。所以拿到标注后的第一步是把 RGB 掩码转成单通道索引图。下面这个脚本按颜色映射转成 0/1 标签多类别场景只要在 color2id 里多写几组颜色映射即可import numpy as np from PIL import Image import glob, os # 颜色 - 类别id按你标注时的实际颜色写 color2id { (0, 0, 0): 0, # 背景 (255, 255, 255): 1, # 前景 } os.makedirs(masks, exist_okTrue) for png in sorted(glob.glob(annotations/*.png)): arr np.array(Image.open(png).convert(RGB)) idx np.zeros(arr.shape[:2], dtypenp.uint8) for color, cid in color2id.items(): idx[np.all(arr color, axis-1)] cid # 按文件名写入masks目录保持和images目录对应 name os.path.basename(png).replace(.png, _mask.png) Image.fromarray(idx).save(os.path.join(masks, name))注意两个细节。第一np.all 是在最后一维做颜色匹配三个通道值完全相等才赋值如果标注图边缘有抗锯齿颜色会介于两者之间这些像素会被留在背景里建议标注时关闭羽化或平滑。第二转换完打开一张 mask 看一眼确认它显示出来是“纯黑 纯白”而不是灰色渐变灰色说明你存成了 0-255 的灰度语义而不是 0/1。这一步错了后面训练出来全是黑的你都不知道是哪出的问题。2.2 自定义 Dataset裁剪、归一化与同步增强数据集做好之后要写一个继承torch.utils.data.Dataset的类。核心工作是三件事读取图像和掩码、做随机裁剪、做图像归一化。裁剪非常重要常见做法是把训练图统一裁成 256x256 或 512x512一方面控制显存另一方面 UNet 对输入尺寸没有硬性要求随机裁剪本身就是一种数据增强。import torch import numpy as np from PIL import Image import glob, os class SegDataset(torch.utils.data.Dataset): def __init__(self, img_dir, mask_dir, crop256, trainTrue): self.imgs sorted(glob.glob(os.path.join(img_dir, *.png))) self.masks [os.path.join(mask_dir, os.path.basename(p).replace(.png, _mask.png)) for p in self.imgs] self.crop crop self.train train def __len__(self): return len(self.imgs) def __getitem__(self, i): img np.array(Image.open(self.imgs[i]).convert(L)) # 灰度图单通道 mask np.array(Image.open(self.masks[i])).astype(np.uint8) h, w mask.shape # 图像和掩码必须用同一个偏移量裁剪否则目标错位 if self.train: y np.random.randint(0, h - self.crop) if h self.crop else 0 x np.random.randint(0, w - self.crop) if w self.crop else 0 else: y max(0, (h - self.crop) // 2) x max(0, (w - self.crop) // 2) img img[y:y self.crop, x:x self.crop] mask mask[y:y self.crop, x:x self.crop] # 归一化到 0~1给通道维mask 保持 0/1不要除以 255 img torch.from_numpy(img / 255.0).float().unsqueeze(0) mask torch.from_numpy(mask).float().unsqueeze(0) return img, mask这段代码里有三个极其容易翻车的点。第一mask 读取时不要加.convert(RGB)灰度标注转成 RGB 会变成三通道完全相同的数组送到损失函数时和模型输出的单通道对不上还会白白多吃三倍显存。第二mask 不能除以 255二分类标签保持 0/1如果归一化成 0-255 的浮点BCE 损失函数会认为前景像素值是 255模型永远无法收敛这是新手常踩的暗坑。第三裁剪偏移量 y、x 必须是同一个变量图像和掩码错位是训练 loss 降不下去最常见的隐性原因。随机翻转也是增强的一部分自己写的话注意要用同一个随机决定翻转与否。我一般这样处理在__getitem__里生成一个if np.random.rand() 0.5的判断然后 img 和 mask 同时做np.fliplr而不是用两个独立的随机种子。3. 用 PyTorch 搭建 UNet 网络编码器、解码器与跳跃连接3.1 主干结构四层下采样与四层上采样怎么拼UNet 的结构不复杂编码器由四个卷积块组成每一块是两个 3x3 卷积加 ReLU块之间用 2x2 最大池化降采样通道数逐层翻倍。解码器对称地做四层上采样每层先把分辨率翻倍再和编码器对应层的输出拼在一起之后过两个卷积。先看完整的模型代码我习惯把它写在 unet_model.py 里import torch import torch.nn as nn def conv_block(in_ch, out_ch): # 每个block两次 convBNReLU尺寸不变 return 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) ) class UNet(nn.Module): def __init__(self, in_ch1, out_ch1, base64): super().__init__() # 编码器 4 层 self.e1 conv_block(in_ch, base) self.e2 conv_block(base, base * 2) self.e3 conv_block(base * 2, base * 4) self.e4 conv_block(base * 4, base * 8) # 瓶颈层 self.bridge conv_block(base * 8, base * 16) # 解码器 4 层 self.d4 conv_block(base * 16, base * 8) self.d3 conv_block(base * 8, base * 4) self.d2 conv_block(base * 4, base * 2) self.d1 conv_block(base * 2, base) # 输出层1x1 卷积降通道 self.out nn.Conv2d(base, out_ch, 1) self.pool nn.MaxPool2d(2) # 上采样用双线性插值不参与梯度回传 self.up nn.Upsample(scale_factor2, modebilinear, align_cornersTrue) def forward(self, x): # 编码器保存跳跃连接要用的中间特征 e1 self.e1(x) e2 self.e2(self.pool(e1)) e3 self.e3(self.pool(e2)) e4 self.e4(self.pool(e3)) # 瓶颈 b self.bridge(self.pool(e4)) # 解码器先上采样再和编码器特征拼接 d4 self.d4(torch.cat([self.up(b), e4], dim1)) d3 self.d3(torch.cat([self.up(d4), e3], dim1)) d2 self.d2(torch.cat([self.up(d3), e2], dim1)) d1 self.d1(torch.cat([self.up(d2), e1], dim1)) return self.out(d1)forward 里的 torch.cat 是整张网络的灵魂。以 d4 这行为例bridge 输出是 H/16 x W/16 的 1024 通道特征self.up 把它放大到 H/8 x W/8此时 e4 的尺寸正好也是 H/8 x W/8两者在通道维拼接后变成 10245121536 通道再进卷积块降回 512。每一层解码器都在“用深层语义 浅层细节”共同决定分割结果。这种设计让 UNet 在目标边缘处有很强的定位能力也是它在中小数据集上比纯卷积分类器改出来的分割网络更稳的原因。3.2 关键参数输入通道、padding 与上采样方式搭建时最常改的参数是 in_ch 和 base。灰度图 in_ch 设 1RGB 图设 3如果你的 mask 是多类别而不是二分类out_ch 改成类别数并且训练时损失函数要换成 CrossEntropyLoss不能继续用 BCE。base 默认 64显存不够时可以降到 32代价是分割细节变差一些。三个必须注意的细节。第一所有卷积层都必须带 padding1否则特征图每过一层缩小 2 像素解码器拼接时两侧尺寸不匹配直接报 RuntimeError。第二上采样方式我倾向于用 bilinear 而不是转置卷积转置卷积参数量大、容易在训练初期产生棋盘格噪声而 bilinear 没有可学习参数解码器照样能恢复细节对显存和收敛速度都更友好。第三如果你在拿到手的pytorch-UNet.zip里看到的是两层卷积后没有 BatchNorm 的版本建议加上 BNBN 能让训练损失曲线平稳很多。下面这张表是 1 通道灰度图输入、512x512 尺寸下的中间特征图情况方便你排查显存和尺寸问题层输入尺寸输出通道输出尺寸e1512x51264512x512pool512x51264256x256e2256x256128256x256e3128x128256128x128e464x6451264x64bridge32x32102432x32d432x3251264x64d364x64256128x128d2128x128128256x256d1256x25664512x512如果你输入尺寸不是 2 的整数次幂比如 300x300也完全能跑因为下采样和上采样各层之间尺寸能对齐。真正会出问题的是输入长宽为奇数时MaxPool2d 默认向下取整导致解码器拼接时差 1 像素。解决办法是在 forward 里对 encoder 特征做一次中心裁剪裁剪量和差值对应。实际项目里我一般直接img img[:, :, :h - h % 16, :w - w % 16]在数据集里先把尺寸规整到 16 的倍数一劳永逸。4. 训练自己的数据损失函数、优化器与训练循环4.1 损失函数为什么二分类分割用 BCE Dice Loss 而不是 MSE分割任务的损失函数选择直接决定模型能不能学出来。新手常见的误用是拿 MSE 做二分类分割结果轮廓模糊一片。原因很简单MSE 把每个像素当成独立回归问题对分类边界附近的置信度不敏感。更合理的做法是BCEWithLogitsLoss加DiceLoss的组合。BCEWithLogitsLoss 内部把 sigmoid 和交叉熵合并了比手动F.sigmoid后再算 BCE 数值上更稳定因为它在 logits 空间直接计算。Dice Loss 衡量的是预测区域和真实区域的重叠度天然对类别不平衡不敏感——当你的目标只占整张图的 2% 时BCE 会让模型倾向于全预测背景因为那样损失也很低但 Dice Loss 会强制它把目标区域找出来。我一般按 1:1 加权代码如下import torch import torch.nn.functional as F def dice_loss(logits, target, smooth1.0): # 先过sigmoid得到概率 prob torch.sigmoid(logits) prob prob.contiguous().view(-1) target target.contiguous().view(-1) intersection (prob * target).sum() # smooth防止分母为0也避免loss在极小区域上剧烈波动 dice (2.0 * intersection smooth) / (prob.sum() target.sum() smooth) return 1.0 - dice # 训练时组合使用 bce nn.BCEWithLogitsLoss() loss bce(logits, mask) dice_loss(logits, mask)注意 dice_loss 里一定要用 sigmoid 而不是 softmax二分类分割的模型输出只有一个通道softmax 在两个类别上做归一化会回传错误的梯度。多分类分割时才用 CrossEntropyLoss此时模型输出通道数等于类别数标签是 long 类型的索引图和这里二分类的 float 标签处理方式完全不同别混用。4.2 训练循环优化器、学习率与模型保存策略优化器我推荐 Adam学习率设 1e-4在 UNet 这种编码器-解码器结构上比 SGD 收敛快很多。权重衰减我一般设 1e-5 或直接不开分割任务的数据量通常不大L2 正则加多了容易欠拟合。学习率调度用ReduceLROnPlateaupatience 设 5-8 个 epoch监控验证集的 loss 或者 dice。下面是训练循环的核心代码包含 AMP 混合精度import torch from torch.cuda.amp import autocast, GradScaler model UNet(in_ch1, out_ch1).cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemax, factor0.5, patience5 ) scaler GradScaler() # AMP混合精度显存不够时的第一选择 best_dice 0.0 for epoch in range(100): model.train() for images, masks in train_loader: images images.cuda() masks masks.cuda() optimizer.zero_grad() with autocast(): logits model(images) loss bce(logits, masks) dice_loss(logits, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() # 验证集上算dice用它决定是否保存模型 val_dice evaluate(model, val_loader) scheduler.step(val_dice) if val_dice best_dice: best_dice val_dice torch.save(model.state_dict(), best_unet.pth)AMP 是显存不足时性价比最高的优化手段autocast让卷积和矩阵运算自动用半精度GradScaler 负责把梯度放大再回传避免半精度下小梯度变成 0。它对最终精度几乎无影响但显存占用能降 40% 左右。如果你的输入图像是 512x512、批量大小 8显存不够时先把 batch 降到 4再开 AMP基本够用。模型保存只保存state_dict()而不保存整个 model 对象这是行业惯例因为 PyTorch 版本变了整个模型文件容易加载失败而 state_dict 是纯张量字典兼容性好很多。恢复训练时你需要额外保存 optimizer 的 state_dict 和当前 epoch我在项目长训时会用文件开头注释写下当时的 lr 和 batch_size因为模型文件多了以后你根本记不住哪个模型是用什么配置训出来的这算我的一个小习惯。5. UNet 训练避坑清单显存、nan 与类别失衡的 5 个现场5.1 显存溢出改输入尺寸比换显卡更实际现象训练跑了几个 batch 后直接报CUDA out of memory不是第一个 batch 就炸而是随着显存碎片累积越来越慢最终爆掉。原因第一个是 batch_size 太乐观512x512 输入下 UNet 中间层 feature map 峰值接近 1GBbatch_size 8 叠加很容易爆第二个隐蔽原因是数据集里灰度图被转成了 3 通道显存直接翻三倍第三个是没有用 AMP。解决优先开 AMP再把 batch_size 降到 4 或 2如果还爆把随机裁剪尺寸从 512 降到 256显存占用会降到原来的四分之一这是最直接的手段。灰度图就保持 1 通道输入没必要为了“对齐 ImageNet 预训练”强行转 3 通道因为 UNet 几乎没有预训练权重可用你从头训1 通道完全没问题。5.2 loss 不降与 nan先从数据集找原因再怀疑优化器现象loss 曲线一直是 0.69 左右纹丝不动或者训着训着突然变成 nan。原因0.69 是log(2)说明模型一直在输出“所有像素都是背景”你的 mask 里正样本占比过低BCE 占主导压过了 Dice Lossnan 一般是学习率过大导致梯度爆炸但还有一个经常被忽略的原因——mask 里存在超出 0/1 范围的值比如用 PIL 读取时不小心做了除法或者保存成了 uint8 的 255。解决先用np.unique(mask)检查标签值域确认只有 0 和 1。如果类别失衡严重把 dice_loss 的权重调到 1.5 或者 2同时给 BCE 设置pos_weight。nan 的话先把学习率降到 1e-5 跑几十个 step能恢复说明是 lr 问题不能恢复就检查数据里有没有无穷值。5.3 预测全黑或全白推理阶段没做 sigmoid 与阈值化现象训练时 loss 很低验证 dice 也不差但保存出来的预测图全黑或者全白。原因模型输出的是 logits不是概率。很多人直接torch.save(pred, result.png)把负值 logits 直接当成灰度存了。另一个原因是归一化方向搞反了把 0-1 的概率乘成了 255 的灰度值还嫌它太暗。解决推理时必须先torch.sigmoid(logits)再和 0.5 比较得到二值图最后乘 255 转 uint8 再保存model.eval() with torch.no_grad(): logits model(img.unsqueeze(0).cuda()) # 输出是logits prob torch.sigmoid(logits) # 转成0~1概率 pred (prob 0.5).float() # 0/1二值化 pred_img (pred.squeeze().cpu().numpy() * 255).astype(np.uint8)这段代码写进了我所有分割推理脚本里。判断预测是否正常的快速方法看概率图的直方图是不是两个峰如果所有值都堆在 0.1 以下说明训练时正样本根本没学会如果堆在 0.9 以上说明模型过自信了检查数据集是不是只有一两张图。5.4 图像与掩码错位增强时的随机种子必须同步现象loss 在 0.5 上下死活降不下去验证集 dice 只有 0.3 左右可视化训练样本发现目标轮廓和原图对不上。原因数据增强时用了两次独立随机操作。比如用 albumentations 时只对 image 调用了 transform对 mask 忘了调用或者自写水平翻转时 image 和 mask 用了两个np.random.rand()判断导致两者翻转向不同。解决增强的黄金法则是“图像和标签共用同一个随机状态”。自写 Dataset 时就按第 2 章代码那样裁剪偏移量用同一个变量翻转用一个 if 判断同时作用于两者。用第三方库时优先选支持同时传入 image 和 mask 的 API再确认两个参数用了同一次 transform 调用。5.5 边界锯齿与空洞上采样方式和后处理的取舍现象dice 分数已经到 0.9 了但预测图的边界像锯齿或者目标内部有零星空洞。原因用 nearest 上采样会丢失亚像素位置信息边界自然毛糙空洞则是 mask 标注本身有小孔或者模型对细小结构把握不足这不是玄学是编码器下采样到 H/16 时小结构信息已经被丢了。解决把网络里的上采样全部换成 bilinear第 3 章的代码就是这么写的边界会平滑一档空洞可以做一次形态学闭运算填补但这是治标如果你发现空洞很大更有效的做法是减小下采样层数或者把 base 通道数加大。一个反直觉的经验不要完全迷信后处理模型本身产出的质量才是上限后处理只能在你“上线前最后一公里”救急。6. 验证你的模型从预测图到 IoU 和 ONNX 导出训练的终点不是 loss 降到最低而是模型能在验证集上稳定复现目标轮廓。我每次训完只做三件事跑一遍验证集可视化、算一次 IoU、导出一个 ONNX。可视化和 IoU 要放在一起做因为单独看数字经常被骗。比如 dice 0.9 但图上一团糊这种模型实际用起来会非常不稳。看 IoU 时不要用 scikit-learn 封装好的函数自己十行代码反而更直观def compute_iou(pred, mask): pred (pred 0.5).float() mask (mask 0.5).float() inter (pred * mask).sum() union (pred mask).clamp(max1).sum() return (inter / (union 1e-6)).item() def predict_and_save(model, img_path, mask_path, save_path): img preprocess(img_path) # 和训练时相同的预处理 mask torch.from_numpy(np.load(mask_path)) model.eval() with torch.no_grad(): prob torch.sigmoid(model(img.unsqueeze(0).cuda())) pred (prob 0.5).float().squeeze().cpu() iou compute_iou(pred, mask) # 原图和预测拼在一起保存方便肉眼检查 vis np.hstack([img.squeeze().numpy(), pred.numpy()]) Image.fromarray((vis * 255).astype(np.uint8)).save(save_path) return iouIoU 比 accuracy 更适合分割因为分割的类别比例经常严重失衡accuracy 会把大量背景像素的“正确”算进分里而 IoU 只看你真正关心的目标区域重叠度。一个合格的单类分割模型验证集 IoU 至少大于 0.6如果只有 0.3 左右大概率是数据量不够或者标注质量不高先不要急着调模型结构。ONNX 导出是我最想强调的一个收尾动作。PyTorch 训练好的模型部署到服务端或者手机端几乎都要经过 ONNX 这一步dummy torch.randn(1, 1, 256, 256).cuda() torch.onnx.export( model, dummy, unet.onnx, opset_version11, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}} )注意导出时设置 dynamic_axes让 batch 维度可变这样实际推实时可以单张图也可以一次多张。导出后用 onnxruntime 加载并跑一次推理和 PyTorch 输出对比误差在 1e-4 量级就说明导出没问题。这个验证不能省模型一旦转成 ONNX推理效率比 PyTorch 原生快不少同时也能避免后端环境没有 PyTorch 时直接干瞪眼的局面。我自己每次训 UNet 的固定习惯是训练到一半随机抽五张验证集的图把原图、真实掩码、预测掩码拼成一张三联图肉眼扫一遍。数字会骗人图不会。边界糊了就回去调数据增强空洞多了就回去看分类权重。这个习惯帮我避开了无数次“测出来很高、用起来翻车”的情况。希望帮到你。本文还有配套的精品资源点击获取
返回列表