PyTorch从零实现Unet:结构解析与torchsummary可视化
2026/9/18 14:16:38 网站建设 项目流程

简介:U-Net网络结构的PyTorch简明实现与torchsummary可视化说明,以单个PDF文件形式提供,整体仅有89KB。文档面向具备一定深度学习基础、希望快速搭建图像分割网络的开发者,围绕下采样与上采样对称结构、跳跃连接、上采样模块等关键环节,给出可直接运行的实现思路。目前已有6923人学习下载,适合作为U-Net入门或工程落地前的快速参考。借助这份PDF,读者可以对照理解每个卷积层、ReLU与Upsample组件的设计用意,掌握left1至bottom、right1至right4等层级间的通道变化关系,并了解如何利用torchsummary展示模型结构、参数量与中间特征维度。文档内容紧凑,覆盖从基础卷积定义到完整U-Net组装的常用步骤,同时给出随机种子设置与数据预处理相关要点,能帮助使用者快速跑通模型、聚焦网络本身的设计细节,省去翻阅分散资料的时间。

1. 从一张分割图说起:Unet 没那么神秘

Unet 是图像分割领域最常见的 baseline 网络,医学影像、遥感地物分类、工业质检里到处是它的变体。它最初为医学图像分割设计,真正让它出圈的是"编码器-解码器 + 跳跃连接"这个干净结构:前一半不断下采样提取语义,后一半逐步上采样恢复分辨率,再用跳跃连接把空间细节补回来,输出和输入同尺寸的逐像素预测图。

对刚入手 pytorch 的人来说,Unet 是性价比最高的练手项目,只用卷积、池化、转置卷积和拼接四种算子,就能把张量在通道与空间两个维度上的流动讲清楚。下面按"结构原理 → 从零实现 → torchsummary 可视化"的顺序,给出一份能直接运行的完整代码,并说明尺寸对齐、BatchNorm 行为和损失函数搭配这些实际使用时绕不开的注意事项。

2. Unet 网络结构逐层拆解:编码器、瓶颈与解码器的设计逻辑

2.1 为什么是 U 形:下采样与上采样的对称关系

Unet 的 U 形来自左右两条对称路径。左侧编码器反复执行"两次 3×3 卷积 + 一次 2×2 最大池化":每经过一次池化,特征图空间尺寸减半、输出通道翻倍,特征从大而浅逐步变成小而深。右侧解码器做镜像操作,先用转置卷积把尺寸翻倍、通道减半,再与编码器同层输出拼接,最后接两次卷积。这个对称设计不是随手画的,它保证每一层都有明确职责:浅层负责边缘、纹理等细节,深层负责"这是什么物体"的语义。

参数效率上这个结构同样讲究。下采样让感受野快速扩大,最深处的瓶颈层能看到整张图的大范围上下文,而通道数逐层翻倍保证了特征容量不随空间缩小而坍塌。以 256×256 输入为例,编码器从 64 通道一路升到 1024 通道,特征图从 256×256 缩到 16×16,这正是 Unet 与原版 FCN 最大的区别——FCN 直接做全连接式的分类头,Unet 则坚持全卷积的像素级映射。

2.2 跳跃连接传递空间细节,拼接为什么优于相加

如果只有编码器和解码器,网络会退化成自编码器结构,分割结果的边缘一片模糊。原因在于连续下采样把"这个像素属于哪个结构"这种精确位置信息一点点抹掉了。Unet 的解法是在解码第 i 层前,把编码器第 i 层的输出沿通道维拼接上来:解码器每一层同时看到两类信息,来自深层的语义知道"是什么",来自跳跃连接的细节知道"在哪里"。

选 concat 而不是 ResNet 式 add 是有依据的。concat 之后卷积核可以对两路特征分别加权,相当于让网络自己学"这一层语义重要还是细节重要";add 是强制逐元素叠加,两路特征被绑死,灵活性差一截。代价是拼接后通道数翻倍、计算量上升,但对 dense prediction 这类任务,这个代价普遍值得。Unet 改进研究中很大一部分的讨论起点就在这里:哪些层需要跳跃、拼接前要不要加过渡卷积、跳跃连接要不要换成注意力加权。

2.3 用 Anaconda 配置 pytorch 环境,先把代码跑起来

动手写网络之前先搭环境,常见做法是 Anaconda 建独立环境,避免污染系统 Python,同时方便以后切换项目。CPU 版最小安装命令如下:

conda create -n unet python=3.9 -y conda activate unet pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu pip install torchsummary

四条命令分别干什么要清楚:第一条创建名为 unet 的环境并指定 Python 3.9;第二条激活环境;第三条从 pytorch 官方 CPU 源安装 torch 与 torchvision,GPU 版只需把--index-url换成对应 CUDA 的小版本,如cu118cu121;第四条安装 torchsummary 用于网络可视化。装完在 PyCharm 或 VSCode 里把项目解释器切到 unet 环境,"明明装过却 import 不到"的 ModuleNotFoundError,九成是解释器没切对。

提示:先跑通 CPU 版再考虑 GPU。256×256 输入下 Unet 前向计算量不大,CPU 足够完成结构验证和 torchsummary 可视化,训练阶段再换 GPU 加速。 注意:torchsummary 的 summary() 会真实执行一次前向传播,它不是静态代码解析器。模型代码里有维度错误时,这一步会直接抛异常,后面第 4 章会专门讲怎么用它排错。

3. 用 PyTorch 从零写 Unet:每个模块都能直接运行

3.1 DoubleConv:卷积、归一化、激活的标准三件套

Unet 里的"两次卷积"组合出现频率最高,先把它抽成独立模块。每个卷积后接 BatchNorm 和 ReLU,这是现代实现与原始论文的差别——原始 Unet 没有 BatchNorm,加上之后收敛速度和稳定性都有明显提升,特别是输入分布差异大的数据集(医学影像、遥感图像)上效果更明显。

import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): """两次卷积 + 批归一化 + ReLU 的组合,保持空间尺寸不变""" def __init__(self, in_ch, out_ch): super(DoubleConv, self).__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)

两个参数是整套结构的地基:kernel_size=3, padding=1保证卷积前后特征图宽高不变,这是 Unet 尺寸对称的前提,漏掉 padding 会让特征图每层缩小 2 像素,堆到深层时尺寸彻底错乱;第一层卷积把通道从 in_ch 转到 out_ch,第二层保持 out_ch 不变,通道变化只发生在每个 Down/Up 的边界处。

3.2 Down 与 Up:池化下采样与转置卷积上采样的对称实现

Down 实现编码器的"尺寸减半、通道翻倍":先做 2×2 最大池化,再接一个 DoubleConv。Up 是 Unet 的精华,包含转置卷积和跳跃连接拼接两个动作:

class Down(nn.Module): """下采样:最大池化 + DoubleConv""" def __init__(self, in_ch, out_ch): super(Down, self).__init__() self.mpconv = nn.Sequential( nn.MaxPool2d(kernel_size=2, stride=2), DoubleConv(in_ch, out_ch), ) def forward(self, x): return self.mpconv(x) class Up(nn.Module): """上采样:转置卷积 + 跳跃连接拼接 + DoubleConv""" def __init__(self, in_ch, out_ch): super(Up, self).__init__() # 转置卷积把尺寸翻倍,通道减半,in_ch 是拼接前的总通道 self.up = nn.ConvTranspose2d(in_ch, in_ch // 2, kernel_size=2, stride=2) self.conv = DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 = self.up(x1) # 两侧尺寸可能差 1 个像素,先用 F.pad 对齐再拼接 diffY = x2.size()[2] - x1.size()[2] diffX = x2.size()[3] - x1.size()[3] x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x = torch.cat([x2, x1], dim=1) return self.conv(x)

这里梳理一下 Up 的通道逻辑:转置卷积把 in_ch 减半,跳跃连接 x2 恰好也是 in_ch / 2 个通道,拼接后恢复为 in_ch,再送入 DoubleConv(in_ch, out_ch)。这就是为什么调用处 in_ch 必须等于"上采样输出通道 + 跳跃连接通道"之和,写错会在 cat 时报维度对不上。forward 里的 F.pad 处理的是奇数尺寸的边界情况:最大池化对奇数维向下取整,下采样后可能差 1 像素,先把 x1 补齐再拼接,避免不到最后一层就崩。

3.3 完整 Unet 类:按 U 形把模块串起来

主体就是按图拼接,注意 down4 把通道翻到 1024,让解码器每一层都精确满足"in_ch 等于两路之和"的关系:

class UNet(nn.Module): """输入 n_channels 通道图像,输出 n_classes 通道分割图""" def __init__(self, n_channels, n_classes): super(UNet, self).__init__() self.inc = DoubleConv(n_channels, 64) self.down1 = Down(64, 128) self.down2 = Down(128, 256) self.down3 = Down(256, 512) self.down4 = Down(512, 1024) # 瓶颈层用最大通道数 self.up1 = Up(1024, 512) self.up2 = Up(512, 256) self.up3 = Up(256, 128) self.up4 = Up(128, 64) self.outc = nn.Conv2d(64, n_classes, kernel_size=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)

forward 的参数名就是跳跃连接的说明:x1 到 x5 是编码器五个阶段的输出,解码端 up1 接 x4、up2 接 x3、up3 接 x2、up4 接 x1,一一对应不能错位。最后 1×1 卷积只做通道映射不改变空间尺寸,把 64 通道压到类别数,二分类就是 1 个通道。整个网络没有任何全连接层和全局池化,所以对输入尺寸是弱约束——只要高宽是 16 的倍数就能跑。

3.4 尺寸对照表与最小运行验证

用 256×256 的 RGB 输入走一遍各层尺寸(通道×高×宽),对照表如下:

模块输入尺寸输出尺寸
inc3×256×25664×256×256
down164×256×256128×128×128
down2128×128×128256×64×64
down3256×64×64512×32×32
down4512×32×321024×16×16
up11024×16×16 拼接 512×32×32512×32×32
up2512×32×32 拼接 256×64×64256×64×64
up3256×64×64 拼接 128×128×128128×128×128
up4128×128×128 拼接 64×256×25664×256×256
outc64×256×2561×256×256

验证脚本就是构造随机张量跑一次前向,确认输出尺寸与输入一致:

if __name__ == "__main__": model = UNet(n_channels=3, n_classes=1) x = torch.randn(1, 3, 256, 256) y = model(x) print("输出尺寸:", tuple(y.shape)) # 期望 (1, 1, 256, 256)

输出和输入同为 256×256,说明这是一个全分辨率的逐像素预测结构。尺寸对不上时优先检查输入高宽是否被 16 整除——四次下采样每次除以 2,2 的 4 次方是 16。这个问题在后续接真实数据时同样会出现,数据加载阶段就该统一做 resize 或 padding,而不是在模型里打补丁。

4. torchsummary 可视化:一行命令看清每一层输出与参数量

4.1 安装并调用 torchsummary

torchsummary 是最轻量的网络结构可视化工具,不依赖绘图库,安装一条命令完成:pip install torchsummary。调用方式:

from torchsummary import summary model = UNet(n_channels=3, n_classes=1) summary(model, input_size=(3, 256, 256), batch_size=1, device="cpu")

两个参数容易踩坑。input_size只写"通道×高×宽",不包含 batch 维,batch 由batch_size单独控制,设为 1 时输出里的形状第一个数字就是 1,看起来更直观;device参数默认是"cuda",CPU 环境不传"cpu"会直接报错找不到设备,网上很多旧教程不写这个参数,在新版本 torchsummary 上单独跑一定会踩这一下。如果是较新的 PyTorch 项目,也可以考虑torchinfo,接口更现代化,但 torchsummary 的表格在细节展示上更经典,够用就行。

4.2 输出字段逐列解读

summary 输出是一个对齐的表格,下面是截取首尾的示例(中间省略):

---------------------------------------------------------------- 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 ... MaxPool2d-7 [1, 64, 128, 128] 0 ... ConvTranspose2d-35 [1, 512, 32, 32] 2,097,664 ... Conv2d-63 [1, 1, 256, 256] 65 ================================================================ Total params: 31,043,521 Trainable params: 31,043,521 Non-trainable params: 0 ----------------------------------------------------------------

逐列看什么,用一张表说清:

输出字段含义排查价值
Layer (type)层名加编号,编号按前向顺序递增编号能看出总层数,本网络共 63 层
Output Shape每层输出张量,第一维是 batch池化后尺寸减半、转置卷积后翻倍,一眼验证对称性
Param #可训练参数数量,0 表示无参数层卷积参数 = 输入通道×输出通道×核宽×核高+偏置
Total params全部参数总和31M 这个量级与经典 Unet 一致
Non-trainable冻结参数数量迁移学习时看这里确认冻结是否生效

以 Conv2d-1 为例,1792 = 3×64×3×3 + 64,正好是输入通道 3、输出通道 64、3×3 卷积核加偏置的组合。ReLU、MaxPool 显示 0 是正常的,它们不产生权重。params size 大约 118MB(31,043,521×4 字节),这是单精度浮点权重的体积,与输入分辨率无关,只由通道配置决定。

4.3 用 summary 定位维度错误的具体手法

torchsummary 是真实执行 forward,任何张量形状不匹配都会变成运行时异常,这是它优于静态解析工具的核心原因。常见的两类报错,对应的排查路线完全不同。

第一类是torch.cat通道对不上:异常信息会指向 Up 模块里 cat 那一行,回查调用处 Up(in_ch, out_ch) 的 in_ch 是否等于"转置卷积输出通道 + 跳跃连接通道"。第二类是输入尺寸不能被 16 整除:比如 200×200 的输入经过四次池化后变成 12×12 的奇数尺寸,问题往往在解码端才暴露,结论是从数据加载处改成 16 的倍数。还有一种隐蔽情况是参数没传对导致 forward 直接没走到——summary 的报错位置在调用 summary 的那一行,这时先单独用随机张量model(x)测一遍,确认模型本身能跑,再让 summary 介入。把这一步养成习惯,之后替换输入尺寸、改通道数,都能在两分钟内确认结构是否还成立。

5. 让 Unet 稳定跑通的四个实战细节

5.1 输入尺寸必须是 16 的倍数

四次下采样对应 16 的因子约束,在数据加载阶段就要处理。常见做法是transforms.Resize((256, 256))统一缩放;需要保留宽高比时,先缩放到合适尺寸再零填充到最近的 16 的倍数。不要在模型内部插入自适应池化,那会破坏 Unet 的尺寸对称性,torchsummary 的输出也失去参考意义。

5.2 BatchNorm 的 train 与 eval 模式差异

BatchNorm 训练时用当前 batch 的均值方差,推理时用全局运行统计量。很多新手验证效果时忘记model.eval(),导致推理结果不稳定甚至明显变差。训练循环用model.train(),验证和推理用model.eval(),这条要写进训练模板。注意浅层 BN 对小 batch 敏感,batch size 小于 4 时把 BN 换成 GroupNorm 是更稳的替代方案。

5.3 损失函数与输出通道的搭配

二分类分割把n_classes设为 1,forward 最后一层不接 sigmoid,直接配nn.BCEWithLogitsLoss,它内部做了数值稳定的 sigmoid 加交叉熵计算,比手动 sigmoid 加 BCELoss 稳定。多分类分割把n_classes设为类别数 K,用nn.CrossEntropyLoss。注意 mask 的 dtype 要和 loss 期望一致,BCE 需要 float,CrossEntropy 需要 long。

5.4 端到端收敛的最小冒烟测试

确认结构后,用随机数据跑几十步,验证前向、反向、优化器整个链路:

model = UNet(n_channels=3, n_classes=1) criterion = nn.BCEWithLogitsLoss() optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) model.train() x = torch.randn(2, 3, 256, 256) mask = torch.randint(0, 2, (2, 1, 256, 256)).float() loss = criterion(model(x), mask) loss.backward() optimizer.step() print(f"loss = {loss.item():.4f}")

在随机数据上 loss 持续下降是正常的,说明模型有学习能力,可以放心换真实数据集。torchsummary 只验证结构,不能验证收敛;结构正确之后,真正决定分割上限的往往是数据预处理、类别不平衡和 loss 权重,这三项比网络本身更值得花时间。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询