ARTICLE DETAIL

资讯详情

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

PyTorch从零搭建UNet:掌握编码器、跳跃连接与图像分割实战

PyTorch从零搭建UNet:掌握编码器、跳跃连接与图像分割实战 简介面向深度学习初学者与图像分割研究人员这是一份基于PyTorch搭建U-Net网络并训练自定义数据集的完整工程资源可应用于医学图像分割、卫星影像分析、视频目标分割等场景。资源系统梳理了U-Net的核心结构编码器通过卷积与池化逐步下采样提取语义特征解码器借助转置卷积恢复空间分辨率并通过跳跃连接融合浅层的边缘纹理与深层的语义信息同时给出图像预处理、标注数据准备与数据增强的具体思路帮助读者将自有图片整理成可训练样本。压缩包共27个文件以7个Python脚本和5个Markdown文档为主体覆盖数据加载、网络定义、训练循环、测试评估与结果可视化另有5个XML工程配置、图像样例与许可证文件整体大小约602KB目录结构清晰便于按功能定位代码模块。已有773人学习本资源对于希望从零实现U-Net并迁移到自身项目的开发者和研究人员可直接参考训练与测试流程有效缩短模型搭建和调参周期从数据准备到模型评估形成完整闭环是一份适合上手和实践的代码型参考包。1. pytorch搭建自己的unet网络为什么我劝你别再用别人的现成权重pytorch搭建自己的unet网络这件事看起来只是把公开代码跑一遍实际是图像分割入门里最值回票价的一条路。很多人直接下载一个训练好的权重做推理遇到自己的数据集就懵了类别不一样、图像尺寸不对、分割效果稀烂还不知道模型内部到底改哪里。自己搭建unet网络的价值在于你能掌握编码器、解码器、跳跃连接这些部件的实际张量变化知道自己改了什么、为什么改。这个方向适合做医学影像分割、遥感地物提取、工业表面缺陷检测的人也适合想弄懂分割网络原理的算法工程师。下面我从网络结构、数据集组织、训练循环到避坑排查把一条完整路径讲清楚。2. 从网络结构图到PyTorch代码编码器、解码器与跳跃连接怎么落地2.1 先看懂UNet的四次下采样为什么图像尺寸要设计成能被16整除UNet最初是医学图像分割论文提出来的结构上分成左边的收缩路径编码器和右边的扩张路径解码器。编码器做四次下采样每一步用MaxPool把特征图尺寸减半通道数翻倍解码器做四次上采样把特征图尺寸逐步恢复回输入大小。四次下采样意味着一个512x512的输入经过池化后会变成256、128、64、32最后瓶颈层是32x32。这里引出一个最常被忽略的问题输入图像的高度和宽度必须能被16整除否则最后一次池化后的特征图尺寸不是整数网络forward直接报错或者尺寸对不齐。我一般要求数据集预处理时把图像统一resize到64的倍数512x512最省心256x256也行因为这样不管怎么实验都不会触到尺寸边界。2.2 双卷积块与跳跃连接先把骨干写出来UNet里最基本的组件是双卷积块也就是两次“3x3卷积 批归一化 ReLU”。这里用3x3卷积加padding1好处是特征图尺寸经过卷积后不变化尺寸只由池化和上采样控制这让张量形状的推演变得非常简单。跳跃连接则是在每个下采样层级把编码器输出的特征图保存下来等解码器上采样之后在通道维度上拼接起来。这样做的直觉是解码器不仅能拿到上采样后的抽象语义特征还能直接看到编码器阶段的高分辨率细节对小目标分割特别重要。2.3 最小可运行的UNet代码骨架常见做法是把UNet拆成几个模块来写便于后续加注意力机制或者替换上采样方式。下面这个版本是经典结构输入通道数、输出通道数和特征通道都做成了参数import torch import torch.nn as nn def conv_block(in_ch, out_ch): # 双卷积块两次 3x3 卷积 BN ReLU 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_ch3, out_ch1, features(64, 128, 256, 512)): super().__init__() # 编码器四个阶段 self.enc1 conv_block(in_ch, features[0]) self.pool1 nn.MaxPool2d(2) self.enc2 conv_block(features[0], features[1]) self.pool2 nn.MaxPool2d(2) self.enc3 conv_block(features[1], features[2]) self.pool3 nn.MaxPool2d(2) self.enc4 conv_block(features[2], features[3]) self.pool4 nn.MaxPool2d(2) # 瓶颈层 self.bridge conv_block(features[3], features[3] * 2) # 解码器上采样 跳跃连接拼接 self.up4 nn.ConvTranspose2d(features[3] * 2, features[3], kernel_size2, stride2) self.dec4 conv_block(features[3] * 2, features[3]) self.up3 nn.ConvTranspose2d(features[3], features[2], kernel_size2, stride2) self.dec3 conv_block(features[2] * 2, features[2]) self.up2 nn.ConvTranspose2d(features[2], features[1], kernel_size2, stride2) self.dec2 conv_block(features[1] * 2, features[1]) self.up1 nn.ConvTranspose2d(features[1], features[0], kernel_size2, stride2) self.dec1 conv_block(features[0] * 2, features[0]) # 输出层1x1 卷积把通道数压到类别数 self.out nn.Conv2d(features[0], out_ch, 1) def forward(self, x): # 编码阶段 e1 self.enc1(x) e2 self.enc2(self.pool1(e1)) e3 self.enc3(self.pool2(e2)) e4 self.enc4(self.pool3(e3)) # 瓶颈 b self.bridge(self.pool4(e4)) # 解码阶段注意 cat 是在通道维度 d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out(d1)这个实现里的关键点有两个一是跳跃连接拼接时用的dim1因为PyTorch里通道维度是第二个维度二是ConvTranspose2d做上采样时kernel_size2, stride2正好让特征图尺寸翻倍不做任何插值。如果你把features参数改成(32, 64, 128, 256)网络参数量会小很多适合显存吃紧或者数据量小的场景。经典UNet第一层卷积核数量是64一般来说不要低于32否则编码器提特征能力不够。2.4 前向传播的形状验证这一步能拦住绝大多数报错写完网络别急着训练先用一个小张量跑一次forward确认输出尺寸和输入尺寸一致。这个习惯能帮你把“网络结构不对”和“数据预处理不对”两类问题分开。model UNet(in_ch3, out_ch1) x torch.randn(1, 3, 512, 512) y model(x) print(y.shape) # 期望输出 torch.Size([1, 1, 512, 512])如果输出尺寸不对先检查输入尺寸是不是16的倍数如果报错说尺寸不匹配打印每一层的形状定位到具体是哪个拼接出了问题。我见过不少人直接把VOC数据集的图片不resize就塞进网络结果有的图是375x500直接就炸了。统一用Dataset预处理做好resize别靠网络自适应。3. 把自己的数据集变成UNet的输入Dataset类与掩码预处理3.1 两种常见的数据集组织方式UNet训练需要成对的“原图 掩码图”。实际项目里有两种主流组织方式。第一种是VOC风格的一个图片文件配一个同名的PNG掩码文件掩码里每个像素值代表类别编号背景是0目标区域是1、2、3……第二种是单通道二值掩码背景黑、前景白适合只分前景背景的任务比如缺陷检测、道路提取。如果你手里的数据集是COCO格式的JSON标注或者像CCPD那样的车牌检测数据集需要先写一个脚本把多边形标注渲染成掩码图再做训练。这一步没有现成轮子就是把polygon画到黑色画布上cv2.fillPoly就能搞定。3.2 自定义Dataset类的标准写法我这里给出一个不带第三方增强库的干净版本逻辑清楚适合先跑通import os import cv2 import torch from torch.utils.data import Dataset class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, size(512, 512)): self.img_dir img_dir self.mask_dir mask_dir self.size size self.images sorted(os.listdir(img_dir)) # 要求图片和掩码同名 def __len__(self): return len(self.images) def __getitem__(self, idx): name self.images[idx] img_path os.path.join(self.img_dir, name) mask_path os.path.join(self.mask_dir, name) # cv2 读出来是 HWC 顺序这里转成 RGB image cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) # 统一尺寸掩码用最近邻插值防止类别标签被平滑 image cv2.resize(image, self.size) mask cv2.resize(mask, self.size, interpolationcv2.INTER_NEAREST) # 归一化到 0~1 image image.astype(float32) / 255.0 mask mask.astype(float32) / 255.0 # 转成 CHW image torch.from_numpy(image).permute(2, 0, 1).float() mask torch.from_numpy(mask).unsqueeze(0).float() # [1, H, W] return image, mask这段代码里有三个容易踩的细节。第一掩码resize必须用INTER_NEAREST用线性插值会把0和255之间的边界变成灰色过渡后续算loss的时候平白多出大量非0非1的预测目标。第二如果你的掩码图是调色板PNG常见于标注工具导出的伪彩色图必须用IMREAD_GRAYSCALE读取否则读出来是三通道的RGB直接拼接维度就错了。第三掩码除以255是有前提的掩码像素值只有0和255。如果标注工具导出的是0和1再除255就把前景值变成0.0039训练出来的模型预测图会全黑。这个我在第5章还会专门讲。3.3 数据增强同步变换是最大坑点有经验的读者到这里会问为什么不用torchvision.transforms说实话torchvision的很多几何变换是单图操作不能保证原图和掩码做完全相同的随机变换即使勉强用RandomResizedCrop也得手动固定随机种子非常容易翻车。我一般直接用albumentations它对“同时处理image和mask”这是原生支持import albumentations as A train_transform A.Compose([ A.RandomResizedCrop(512, 512, scale(0.75, 1.0)), A.HorizontalFlip(p0.5), A.VerticalFlip(p0.2), A.RandomBrightnessContrast(p0.3), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ])注意Normalize放到最后因为随机亮度对比度操作要在原始像素范围上做才有意义。使用的时候调用方式统一传关键字参数augmented train_transform(imageimage, maskmask)返回的augmented[image]是(H, W, C)的numpy数组需要再转成CHW并转成tensor。这里有个细节如果mask是多类别分割的albumentations的随机增强不会改变像素值所以在变换之后不需要重新归一化类别如果是二值图也不需要再做二值化。原图做了Normalize掩码千万不要做Normalize。3.4 快速验证数据管线数据管线写完后用下面这段脚本做一次体检比直接开训练省时间from torch.utils.data import DataLoader dataset SegDataset(train_images, train_masks) loader DataLoader(dataset, batch_size4, shuffleTrue) for images, masks in loader: print(images.shape, masks.shape) # [4,3,512,512] [4,1,512,512] print(pixel range:, images.min().item(), images.max().item()) print(mask unique:, torch.unique(masks).tolist()) # 期望 [0.0, 1.0] break这一步能同时确认尺寸、归一化范围、掩码取值三个关键信息。我自己的习惯是每次换数据集都先跑这个脚本看到mask unique是[0.0, 1.0]才放心进训练。如果你在这里就发现mask里有0.0039这样的值说明掩码本身是0和1却被除以255回头改Dataset就好。4. 训练自己的数据集损失函数、优化器与训练循环的取舍4.1 损失函数选择二分类用BCEWithLogitsLoss多分类用CrossEntropyLossUNet的输出层是1x1卷积输出通道数等于类别数。二分类任务out_ch1输出张量形状是[B, 1, H, W]推荐用BCEWithLogitsLoss。这个损失函数内部先做了Sigmoid再做交叉熵数值上比“手动Sigmoid BCELoss”稳定。多分类任务out_ch类别数配合CrossEntropyLoss交叉熵内部的Softmax会把通道维当作类别维。这里有个常见的张量形状误区BCEWithLogitsLoss要求预测和标签都是[B, 1, H, W]标签必须是floatCrossEntropyLoss要求预测是[B, C, H, W]标签是[B, H, W]的长整型索引不是one-hot。搞反这两种格式训练时要么报错要么loss值诡异地大。4.2 一个epoch的完整训练循环训练循环本身不复杂难点在处理好设备、梯度清零、验证时机device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_ch3, out_ch1).to(device) criterion torch.nn.BCEWithLogitsLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, factor0.5, patience5) def train_one_epoch(model, loader, criterion, optimizer): model.train() total_loss 0.0 for images, masks in loader: images, masks images.to(device), masks.to(device) preds model(images) # 前向 loss criterion(preds, masks) # 计算损失 optimizer.zero_grad() # 梯度清零 loss.backward() # 反向传播 optimizer.step() # 更新参数 total_loss loss.item() * images.size(0) return total_loss / len(loader.dataset)参数上的讲究主要在初始化学习率。UNet用Adam优化器初始学习率我一般设置在1e-4到3e-4之间数据量小、类别不平衡严重时用1e-4更稳因为大学习率容易让模型一开始就进入全背景的局部最优。ReduceLROnPlateau按验证集loss降低学习率比固定步长衰减省心patience设为5到8个epoch即可。num_workers建议设为4以上pin_memoryTrue能减少数据从CPU到GPU的拷贝时间代价是占用更多内存。4.3 验证集指标别只盯着loss算一下mIoU训练过程中每1个或2个epoch在验证集上算一次mIoU比只看loss更能反映分割质量。loss是平滑的但二值分割里可能出现“loss一直在降、预测图却全黑”的假象。mIoU的公式是交集除以并集下面是一个适合二分类的极简实现def compute_iou(preds, masks, eps1e-6): preds (torch.sigmoid(preds) 0.5).int() masks masks.int() intersection (preds masks).sum(dim(1, 2, 3)).float() union (preds | masks).sum(dim(1, 2, 3)).float() iou (intersection eps) / (union eps) return iou.mean().item() def evaluate(model, loader): model.eval() ious [] with torch.no_grad(): for images, masks in loader: images, masks images.to(device), masks.to(device) preds model(images) ious.append(compute_iou(preds, masks)) return sum(ious) / len(ious)注意sigmoid(preds) 0.5这个阈值不是固定的。类别不平衡时0.5可能不是最优阈值可以画一次PR曲线选最佳阈值。但初版模型先用0.5没问题。eps是为了防止某张图预测和标签全为空导致除零。验证时记得加model.eval()和no_grad()否则BatchNorm的统计量会被验证数据污染而且推理阶段还计算梯度非常浪费显存。4.4 保存模型只保存权重还是保存整个模型我建议训练过程中每个epoch都记录最佳验证mIoU并保存对应的权重best_iou 0.0 for epoch in range(epochs): train_loss train_one_epoch(model, train_loader, criterion, optimizer) val_iou evaluate(model, val_loader) scheduler.step(val_iou) if val_iou best_iou: best_iou val_iou torch.save(model.state_dict(), best_unet.pth) print(fepoch {epoch}: saved, iou{val_iou:.4f})只保存state_dict不带模型结构好处是结构改动后老权重不至于完全没法用坏处是加载时要手动重新定义模型。如果你想省事torch.save(model, best_full.pth)也是合法的但后续要部署到别处时容易被PyTorch版本迁移问题卡住。我的习惯是两者都存权重文件用来做推理完整模型文件用来调试。5. UNet训练避坑指南5个反复出现的问题与排查顺序5.1 现象loss在下降但预测图全黑或全白这个坑我踩过不止一次。训练过程看起来完全正常loss从0.7一路降到0.1但推理时输出整张图都是黑色或者偶尔全白。原因几乎都出在掩码预处理灰度掩码图里前景是255除以255成为1.0这是对的但如果标注工具导出的是0和1的掩码或者导出的是单通道二值图但像素值只有2和3这种索引值除以255后前景变成0.0039或0.0118模型学到的是“全部输出0就是最优解”。排查方法是在训练前打印torch.unique(mask)确认标签值只有0和1。解决方式在Dataset的__getitem__里先判断像素最大值如果最大像素值大于1再做除法否则直接用原始值。这个判断写进代码里能一劳永逸。5.2 现象训练到一半显存溢出CUDA out of memory特征图占用显存最多的环节是跳跃连接拼接后的解码器层因为那里同时保存了编码器的高分辨率特征图。512x512输入、batch_size8、经典64通道配置大概需要11GB左右显存。如果你的显卡只有8GB优先把batch_size降到4或2而不是砍模型通道数。如果降batch_size还不行可以开启梯度累积每两个小batch累积一次梯度再更新参数效果近似大batch。还有一个很多人忽略的点torch.backends.cudnn.benchmark True在输入尺寸固定的情况下能加速卷积但每次输入尺寸变化时它反而会拖慢——所以统一resize到固定尺寸不仅是正确性问题也是性能问题。5.3 现象输入图片不是16的倍数forward报size mismatch一旦输入尺寸不满足16的倍数报错位置通常在torch.cat这一步因为上采样后的特征图尺寸和跳跃连接传过来的特征图尺寸对不上。报错信息会提示两个张量的size不同但不会告诉你根因是输入尺寸没满足池化要求。排查方法是第一步在forward里打印每一层输出的shape第二步检查输入图像原始尺寸。解决方法是直接在Dataset里resize到固定尺寸比如512x512。另一种做法是在数据加载后用torch.nn.functional.pad把图补到16的倍数但分割任务的标签也要同步pad而且pad区域不应该计入loss处理起来比较脏不如直接resize。5.4 现象数据增强后标签和图像错位用torchvision做随机翻转和裁剪时原图和掩码的随机因子可能不一致导致训练时模型看到的“原图是汽车标签却是道路”。这个问题非常隐蔽因为loss不会明显异常只会导致验证指标上不去。我的解决方式很简单放弃torchvision的几何变换统一用albumentations它对image和mask用的是同一套坐标变换参数。另外还有一个细节不要把Normalize应用到mask上mask经过归一化后就不再是类别标签了0.5的阈值会把它变成全零。5.5 现象验证集mIoU很高但放到真实场景里分割效果很差这类问题本质是训练分布和真实分布不一致。常见原因有三个一是数据泄露训练集和验证集里有重复图片或同一场景的邻近帧导致验证指标虚高二是验证集只包含简单样本比如只挑了大目标、对比度高的图三是真实场景的输入分辨率和训练分辨率不一致模型没见过更大或更小的目标尺度。应对方式按场景或按视频序列切分数据保证验证集独立验证时也统一resize到训练尺寸如果真实场景目标尺度多变考虑把训练数据做多尺度增强而不是只固定512x512。6. 进阶从权重保存到单张推理给模型做一次完整体检6.1 把训练好的权重加载回来做推理权重加载要保证模型结构和训练时一致。我用map_location指定设备避免在无GPU机器上加载时报错model UNet(in_ch3, out_ch1) model.load_state_dict(torch.load(best_unet.pth, map_locationcpu)) model.eval()这里有个容易被新手忽略的点加载权重后必须调用model.eval()否则BatchNorm层还在用训练时的数据统计量推理结果会不稳定。做完这一步拿一张验证集之外的图跑一次前向把输出转换成8位整数保存成可视化结果import numpy as np import cv2 image cv2.imread(test.jpg) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) image cv2.resize(image, (512, 512)) tensor torch.from_numpy(image / 255.0).permute(2, 0, 1).unsqueeze(0).float() with torch.no_grad(): pred torch.sigmoid(model(tensor))[0, 0].cpu().numpy() pred (pred 0.5).astype(uint8) * 255 cv2.imwrite(pred.png, pred)这一小段就是完整的推理管线。真正的分割项目还要做一步后处理去掉面积过小的连通域或者用形态学闭运算填充孔洞但对初版模型来说能稳定地看到分割轮廓就已经说明整条链路是通的。6.2 模型体检看看哪里分割错了单张图可视化只能看出大概效果系统性体检建议做两类分析。第一类是分尺寸统计mIoU把验证集图片按目标尺寸排序看你的模型在哪个尺度范围效果最差这直接决定你要不要做多尺度训练或使用更大输入尺寸。第二类是错误类型统计预测图里是假阳性多还是假阴性多。假阳性多说明模型把背景纹理误判成目标适当提高置信度阈值有效假阴性多说明小目标被漏检优先考虑数据增强中增加小目标出现的概率而不是盲目加深网络。6.3 后续值得做的三个方向训练跑通只是起点把它用起来的价值更大。常见的下一步是换损失函数BCEWithLogitsLoss配合DiceLoss加权组合早期训练能更快突破极小目标的困境。第二个方向是把ConvTranspose2d上采样换成Upsample 3x3卷积会在边缘细节上更平滑但注意这种改法会增加一点参数量。第三个方向是导出成ONNX做部署验证这在服务端推理时会省掉PyTorch环境的依赖排查问题也更直接。我自己做了一次这种“从零搭建到落地”之后最大的习惯改变是每改一个环节先跑最小验证脚本再开完整训练。这个习惯帮我把好几次差点在错误数据集上浪费一整天的事情拦了下来。希望帮到你。本文还有配套的精品资源点击获取
返回列表