news 2026/9/18 21:57:43

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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_size=2、stride=2,这是标准的尺寸减半写法。

import torch import torch.nn as nn class DoubleConv(nn.Module): """两个3x3卷积+BN+ReLU的标准unet基础块""" def __init__(self, in_ch, out_ch): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x)

这个DoubleConv块是整条Unet路径的基础,我通常把inplace=True打开以减少显存占用。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_size=2, stride=2) 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], dim=1) return self.conv(x)

这里用skip_ch单独表示跳跃连接的通道数,是因为最底层的跳跃连接通道数跟编码器支路不完全一样。torch.cat([x, skip], dim=1)是把通道维度直接拼起来,拼完后的通道数是两者之和,所以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_channels=3, num_classes=1, features=[64, 128, 256, 512]): super().__init__() self.downs = nn.ModuleList() self.ups = nn.ModuleList() self.pool = nn.MaxPool2d(kernel_size=2, stride=2) # 编码器: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_size=1) 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要根据输入图像通道数修改,灰度图是1,RGB是3。num_classes=1配合sigmoid做二分类分割,多类别的医疗分割任务记得改成类别数。编码器部分每层卷积后接池化,最后一个下采样层之后不再池化,而是直接进bottleneck,这个细节很多人会写错。

3. 直接运行Unet训练:从数据到loss的完整流程

3.1 最小可用的数据加载与预处理方案

unet模型是干什么的——简单说就是逐像素分类,所以训练数据只需要原始图像和对应的掩码图像(mask)。为了做到“可以直接运行”,我通常用公开的分割数据集,配合PyTorch的DatasetDataLoader接口来加载。关键预处理只有三步:图像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, interpolation=Image.NEAREST)(mask) image = T.ToTensor()(image) mask = torch.as_tensor(np.array(mask), dtype=torch.long) # 如果有背景类别,把像素值255改成1,变成二分类标签 mask = (mask > 128).long() return image, mask

这里的核心参数是interpolation=Image.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_channels=3, num_classes=2).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-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(f"Epoch {epoch+1:02d} Loss: {avg_loss:.4f}")

PyTorch的CrossEntropyLoss会自动对outputs的第1维做softmax,所以模型最后一层不需要手动接softmax,直接输出原始logits即可。masks的shape必须是(B, H, W),而不是(B, 1, H, W),多了一个通道维度会导致loss计算直接报错或结果完全错误。Adam优化器里lr=1e-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 size=2跑通再往上加。如果实在要调大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_channels=3, num_classes=2).to(device) summary(model, input_size=(3, 256, 256), batch_size=2)

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,410

Output Shape列里有一类比较特殊——特征图尺寸不变的那些层(比如所有ReLU),此时Param为0,因为激活函数没有可训练参数。BatchNorm的参数量是2乘以通道数,对应gamma和beta两个可学习向量。如果看到某个Conv层的Output Shape和你预期不同,问题大概率出在padding或stride设置上,对照(H - kernel_size + 2*padding) / stride + 1就能算明白。

4.3 一个常见坑:总步长和输入尺寸不匹配的解决办法

torchsummary本身不会报尺寸错误,报错的是模型前向传播。Unet的总步长是16(4次池化,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实现绕来绕去核心就那几个结构块,验证通过之后就可以放心去调数据增强、学习率策略和损失函数的改进方向了。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 21:55:58

CloddsBot市场索引API:Kalshi/Manifold实时行情语义检索实战指南

CloddsBot市场索引API:Kalshi/Manifold实时行情语义检索实战指南 【免费下载链接】CloddsBot Open Source AI trading agent that operates autonomously across 1000 markets - Polymarket, Kalshi, Binance, Hyperliquid, Solana DEXs, 5 EVM chains. Scans for e…

作者头像 李华
网站建设 2026/9/18 21:52:38

设计自己的小传输协议:导论与概念详解

一、引言:为什么需要理解传输协议在大多数应用开发者的日常工作中,网络通信往往被抽象成一行简单的调用:打开一个 Socket,写入一些字节,再读取一些字节。对于使用 HTTP、gRPC、WebSocket 这类成熟协议的业务系统来说&a…

作者头像 李华
网站建设 2026/9/18 21:49:39

MariaDB 3306 握手失败?让走 TaoToken 的 Codex 对照 pymysql 驱动查

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/18 21:49:32

千笔AI与云笔AI论文写作工具深度对比

1. 论文写作工具的现状与痛点作为一名在学术圈摸爬滚打多年的研究者,我深知论文写作过程中的各种痛苦。从选题构思到文献综述,从实验设计到结果分析,每个环节都让人头疼不已。特别是对于在职攻读学位的专业人士,如何在繁忙工作之余…

作者头像 李华
网站建设 2026/9/18 21:49:17

目标检测实战:溺水检测数据集构建与YOLOv8训练全解析

做溺水检测这个方向,说难不难,说简单也真不简单。难点不在模型——现在的目标检测框架一个比一个成熟,YOLO拉起来就能跑;真正的痛点在数据。COCO、VOC这些公开数据集里根本没有“溺水”这个类别,想从零开始标一套又费时…

作者头像 李华