ARTICLE DETAIL

资讯详情

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

PyTorch实现Unet图像分割:网络结构详解与torchsummary可视化

PyTorch实现Unet图像分割:网络结构详解与torchsummary可视化 简介面向图像分割初学者与PyTorch入门者的一份PDF说明文档重点讲解用PyTorch搭建Unet卷积神经网络并结合torchsummary可视化模型结构。Unet常用于医学图像等场景的像素级分割任务整体呈对称U形左侧四个下采样阶段逐级提取深层语义右侧四个上采样阶段逐步恢复分辨率同时通过特征拼接融合低级细节文档对MaxPool下采样、Up_Sample上采样、卷积与ReLU定义、随机种子设置等关键环节均有说明。内容还包含可直接运行的核心代码片段与torchsummary可视化示例并给出结构图示以帮助对照理解每一层输入输出尺寸变化读者只需设定输入输出通道数即可复现模型搭建与训练验证流程。资源为1个PDF文件压缩包约89KB内容精炼集中目前已有6923人学习下载适合希望快速上手Unet图像分割、理解U形编解码结构并掌握PyTorch模型可视化方法的中初级开发者。1. Unet网络结构在PyTorch里的最小实现比想象中简单Unet图像分割模型在PyTorch里实现起来其实代码量比很多新手想象中少得多。虽然unet网络结构看起来有编码器、解码器、跳跃连接三大部分但真正落地到代码核心结构只有两个卷积块和一个拼接操作其余就是循环打包。标题里说的“可以直接运行”意味着代码不需要额外预处理脚本、不需要自定义数据集类下载下来跑通训练循环就能看到loss变化。这对想快速验证unet模型是干什么的、想看看分割效果长什么样的开发者来说是最省事的路径。用PyTorch实现unet的优势在于张量操作和网络层定义天然贴合Unet的“编码-解码”思路。torchsummary则是顺手解决网络结构可视化问题的工具一行代码就能把每层输出尺寸和参数量打印出来。这篇文章会从头拆Unet的每个结构块给出一份能直接运行的PyTorch实现代码再讲清楚torchsummary参数怎么调、输出怎么看。适合已经会Python基础语法、想从分类网络跨到分割任务的人也适合需要快速搭一个baseline做对比实验的工程人员。2. Unet网络结构拆解编码器、解码器和跳跃连接的PyTorch对应关系2.1 Unet编码器卷积池化的下采样路径Unet的编码器部分就是一个典型的卷积特征提取过程每一层包含两个3x3卷积激活函数用ReLU跟在后面的是一个2x2最大池化。池化把特征图尺寸减半同时下一样卷积层的通道数翻倍这就是Unet“先收缩”的阶段。从64通道开始每经过一次下采样通道变为128、256、512最后到底层变成1024。PyTorch实现编码器时常见做法是把“两个卷积ReLU”封装成一个双卷积块然后重复调用。关键参数上有两个地方容易踩坑一个是conv2d里padding要设为1否则3x3卷积会让特征图缩小一圈叠加多次后尺寸对不上跳跃连接的要求另一个是池化层kernel_size2、stride2这是标准的尺寸减半写法。import torch import torch.nn as nn class DoubleConv(nn.Module): 两个3x3卷积BNReLU的标准unet基础块 def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x)这个DoubleConv块是整条Unet路径的基础我通常把inplaceTrue打开以减少显存占用。BatchNorm放在卷积和激活之间而不是卷积之前是为了贴合原版Unet的实现习惯。后面的down采样层就是在这个块外面套一层MaxPool2d每层输出同时给到下采样分支和跳跃连接分支。2.2 解码器和跳跃连接特征拼接的尺寸对齐问题解码器是Unet和普通自编码器的最大区别所在。每一层上采样先用转置卷积把特征图尺寸翻倍然后把对应的编码器特征图在通道维度上拼接起来再经过DoubleConv融合。转置卷积是常见的上采样选择也可以换双线性插值但转置卷积参数可学习分割效果通常更好。拼接时最大的坑是编码器特征图和解码器特征图尺寸不一致常见原因有两个一是编码器用了奇数尺寸的输入图二是卷积padding没设对。为了从根上避开这个问题常规做法是让输入尺寸满足“能被2整除4次”常见的256x256和512x512都满足128x128也可以。如果输入是奇数尺寸转置卷积输出的尺寸会跟编码器差1个像素torch在拼接时直接报错。class Up(nn.Module): 上采样块转置卷积翻倍尺寸 跳跃连接拼接 双卷积 def __init__(self, in_ch, skip_ch, out_ch): super().__init__() self.up nn.ConvTranspose2d(in_ch, in_ch // 2, kernel_size2, stride2) self.conv DoubleConv(in_ch // 2 skip_ch, out_ch) def forward(self, x, skip): x self.up(x) # 如果尺寸差一个像素用F.pad补齐再拼接 if x.size(2) ! skip.size(2): diff skip.size(2) - x.size(2) x F.pad(x, [diff // 2, diff - diff // 2] * 2) x torch.cat([x, skip], dim1) return self.conv(x)这里用skip_ch单独表示跳跃连接的通道数是因为最底层的跳跃连接通道数跟编码器支路不完全一样。torch.cat([x, skip], dim1)是把通道维度直接拼起来拼完后的通道数是两者之和所以DoubleConv的输入通道要写成in_ch // 2 skip_ch。2.3 整体网络组装Unet类的完整PyTorch实现把编码器和解码器串起来就是完整的Unet网络。编码器依次是64、128、256、512通道最底层是512到1024的DoubleConv。解码器反向操作1024上采样到512拼接跳跃连接后变回512、256、128、64。最后一层用一个1x1卷积把64通道映射到分割类别数二分类就输出1通道多分类输出类别数。class UNet(nn.Module): unet完整实现输入(B,3,H,W) 输出(B,num_classes,H,W) def __init__(self, in_channels3, num_classes1, features[64, 128, 256, 512]): super().__init__() self.downs nn.ModuleList() self.ups nn.ModuleList() self.pool nn.MaxPool2d(kernel_size2, stride2) # 编码器4次下采样 for f in features: self.downs.append(DoubleConv(in_channels, f)) in_channels f # 最底层bottleneck self.bottleneck DoubleConv(features[-1], features[-1] * 2) # 解码器4次上采样 for f in reversed(features): self.ups.append(Up(f * 2, f, f)) self.ups.append(DoubleConv(f * 2, f)) # 其实这里Up里已包含conv保留以兼容某些实现 self.final_conv nn.Conv2d(features[0], num_classes, kernel_size1) def forward(self, x): skip_connections [] for down in self.downs: x down(x) skip_connections.append(x) x self.pool(x) x self.bottleneck(x) skip_connections skip_connections[::-1] # 反转让最浅层先拼 for idx in range(0, len(self.ups), 2): x self.ups[idx](x, skip_connections[idx // 2]) return self.final_conv(x)参数说明features列表控制Unet深度默认[64, 128, 256, 512]是经典配置小数据集可以缩减为[32, 64, 128, 256]来降低显存占用并加速训练。in_channels要根据输入图像通道数修改灰度图是1RGB是3。num_classes1配合sigmoid做二分类分割多类别的医疗分割任务记得改成类别数。编码器部分每层卷积后接池化最后一个下采样层之后不再池化而是直接进bottleneck这个细节很多人会写错。3. 直接运行Unet训练从数据到loss的完整流程3.1 最小可用的数据加载与预处理方案unet模型是干什么的——简单说就是逐像素分类所以训练数据只需要原始图像和对应的掩码图像mask。为了做到“可以直接运行”我通常用公开的分割数据集配合PyTorch的Dataset和DataLoader接口来加载。关键预处理只有三步图像resize到统一尺寸、转Tensor、归一化到0-1区间。mask需要把类别值映射到0到num_classes-1的整数标签。from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as T class SegDataset(Dataset): 最小分割数据集image和mask目录下放同名文件即可 def __init__(self, img_dir, mask_dir, size(256, 256)): self.img_dir img_dir self.mask_dir mask_dir self.size size self.names [f for f in os.listdir(img_dir) if f.endswith(.png) or f.endswith(.jpg)] def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] image Image.open(f{self.img_dir}/{name}).convert(RGB) mask Image.open(f{self.mask_dir}/{name.replace(.jpg, .png)}).convert(L) image T.Resize(self.size)(image) mask T.Resize(self.size, interpolationImage.NEAREST)(mask) image T.ToTensor()(image) mask torch.as_tensor(np.array(mask), dtypetorch.long) # 如果有背景类别把像素值255改成1变成二分类标签 mask (mask 128).long() return image, mask这里的核心参数是interpolationImage.NEAREST。mask是标签图resize时如果用双线性或双三次插值会产生0到255之间的中间值标签就乱了。最近邻插值不会引入新像素值这是分割任务和数据加载的通用要求。另一个参数是256, 256的resize尺寸前面提到过需要能被2整除4次256是比较稳妥的选择显存紧张可以换成128。3.2 训练循环、损失函数和评估指标怎么配Unet的损失函数一般用交叉熵多分类或BCEWithLogitsLoss二分类。医疗分割任务里经常有类别不平衡问题可以给交叉熵加weight参数让稀有类别的梯度贡献更大。训练循环本身跟分类网络没有本质区别但每次迭代要手动把预测结果转成概率再算loss。device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, num_classes2).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-4) criterion nn.CrossEntropyLoss() for epoch in range(30): model.train() total_loss 0.0 for images, masks in data_loader: images, masks images.to(device), masks.to(device) optimizer.zero_grad() outputs model(images) # (B, 2, H, W) loss criterion(outputs, masks) # masks: (B, H, W) loss.backward() optimizer.step() total_loss loss.item() avg_loss total_loss / len(data_loader) print(fEpoch {epoch1:02d} Loss: {avg_loss:.4f})PyTorch的CrossEntropyLoss会自动对outputs的第1维做softmax所以模型最后一层不需要手动接softmax直接输出原始logits即可。masks的shape必须是(B, H, W)而不是(B, 1, H, W)多了一个通道维度会导致loss计算直接报错或结果完全错误。Adam优化器里lr1e-4是我常用的分割任务初始值比分类任务的1e-3更保守因为Unet参数量大太激进的步长容易让训练早期就震荡。3.3 训练过程中最常遇到的3个报错及修复第一个报错是Expected 4D input。网络输入是(B, C, H, W)四维张量很多新手在单张验证时给的是(C, H, W)三维张量。修复方法是images.unsqueeze(0)加一个batch维度或者直接用DataLoader保证四维。这个报错在torchsummary里也会出现原因类似。第二个报错是CUDA显存不足。输入尺寸256x256、batch size 8的默认配置在6GB显卡上勉强能跑但建议先batch size2跑通再往上加。如果实在要调大batch可以把features改成[32, 64, 128, 256]参数量直接降到原来的四分之一。第三个报错是loss不下降且数值很大。常见原因是mask标签里混入了255或其他异常像素值。交叉熵对任意整数类别都接受但类别数超出num_classes范围时loss会变得非常大。排查方法是在训练循环里打印masks.unique()确认标签类别数跟模型输出通道数一致。4. torchsummary可视化Unet网络结构参数含义和输出解读4.1 torchsummary安装和最小调用代码torchsummary和pytorch安装是配套的用pip安装后直接summary(model, input_size)就行。PyTorch官方并不自带这个工具网上搜“pytorch 模型可视化”出来的结果大多数也是torchsummary。它做的事情就是给模型喂一个假输入跑一次前向传播同时把每层输入输出张量的shape和参数量记录下来。# 安装命令和pytorch的安装步骤互相独立 # pip install torchsummary from torchsummary import summary model UNet(in_channels3, num_classes2).to(device) summary(model, input_size(3, 256, 256), batch_size2)input_size必须是一个tuple第一个维度是通道数不是图像宽。batch_size参数可以不传默认是-1输出显示的是-1占位但显存足够建议显式传一个值。summary执行时会真实跑一次前向所以模型必须在正确的device上显存占用也跟真实推理差不多。4.2 读懂torchsummary的输出Output Shape和Param字段torchsummary的输出是一张表格每个layer一行列分别是Layer名称、Output Shape和Param。Output Shape显示的是该层输出的张量维度比如[-1, 64, 256, 256]表示batch维是-1通道数64宽高256x256。从这个列可以直观看到Unet网络结构的尺寸变化路径编码器阶段H和W不断减半通道数翻倍解码器阶段反过来。Param字段是该层的参数量是所有卷积核的权重加bias之和。可以把同颜色层的参数相加得到整个Unet的参数量经典配置大概是31M。我一般用这个数值判断模型是否加载正确——如果加载的预训练权重和模型参数量不一致torchsummary打印的总数会对不上。---------------------------------------------------------------- Layer (type) Output Shape Param # Conv2d-1 [-1, 64, 256, 256] 1,792 BatchNorm2d-2 [-1, 64, 256, 256] 128 ReLU-3 [-1, 64, 256, 256] 0 Conv2d-4 [-1, 64, 256, 256] 36,928 BatchNorm2d-5 [-1, 64, 256, 256] 128 ReLU-6 [-1, 64, 256, 256] 0 MaxPool2d-7 [-1, 64, 128, 128] 0 Conv2d-8 [-1, 128, 128, 128] 73,856 ... ConvTranspose2d-31 [-1, 512, 128, 128] 2,097,152 Total params: 31,033,410Output Shape列里有一类比较特殊——特征图尺寸不变的那些层比如所有ReLU此时Param为0因为激活函数没有可训练参数。BatchNorm的参数量是2乘以通道数对应gamma和beta两个可学习向量。如果看到某个Conv层的Output Shape和你预期不同问题大概率出在padding或stride设置上对照(H - kernel_size 2*padding) / stride 1就能算明白。4.3 一个常见坑总步长和输入尺寸不匹配的解决办法torchsummary本身不会报尺寸错误报错的是模型前向传播。Unet的总步长是164次池化2的4次方意味着输入尺寸必须是16的倍数否则最终输出尺寸会跟输入对不齐或者在跳跃连接阶段拼接失败。RuntimeError: Sizes of tensors must match except in dimension 1. Expected size 65 but got size 64这个报错的修复思路有两个。第一是改输入尺寸把图像resize到16的倍数比如224不是16倍数改成256或240。第二是在Up块里加尺寸补齐逻辑前面代码里的F.pad就是干这个的。我倾向于第二种方案这样训练和推理时即使输入尺寸不固定也能跑通代价是每次拼接多一次pad操作性能损耗可以忽略。5. 验证Unet网络结构和torchsummary输出是否正确的3个技巧写完模型不看任何训练指标先用三个小技巧确认网络结构本身没问题。第一个技巧是在输入全0和全1两个张量上分别做一次前向传播确认输出shape都是(B, num_classes, H, W)且输出值不相同。如果输出shape对但值一样说明模型某个位置把输入丢弃了通常是跳跃连接写成了加法而不是拼接。第二个技巧是验证编解码器的对称通道数跟踪。直接在forward的自定义代码里插入print看每一层x.shape或者用PyTorch的register_forward_hook打印特征图维度变化。重点对比downs列表第i层输出和解码器第i次上采样前的通道数经典Unet中解码器第i层拼接前的通道数应该是编码器第i层的通道数加上上一层传来的通道数。第三个技巧是可视化训练早期的预测结果。训练到第1个epoch后把一张验证集图像的预测mask画出来。如果Unet结构实现正确但还没训练好预测图应该是均匀的噪声块如果整张图全黑或全白说明最后的1x1卷积输出或者loss计算有问题。更细一步统计预测mask中每个类别的像素占比如果某个类别占比从不发生变化优先怀疑数据集里这类样本的mask在resize时被抹掉了。这三个技巧花不到五分钟能省掉后续训练几小时发现白跑的时间。Unet的PyTorch实现绕来绕去核心就那几个结构块验证通过之后就可以放心去调数据增强、学习率策略和损失函数的改进方向了。本文还有配套的精品资源点击获取
返回列表