PyTorch从零实现Unet:结构解析与torchsummary可视化
简介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 python3.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 与 torchvisionGPU 版只需把--index-url换成对应 CUDA 的小版本如cu118或cu121第四条安装 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_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)两个参数是整套结构的地基kernel_size3, padding1保证卷积前后特征图宽高不变这是 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_size2, stride2), 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_size2, stride2) 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], dim1) 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_size1) 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×256down164×256×256128×128×128down2128×128×128256×64×64down3256×64×64512×32×32down4512×32×321024×16×16up11024×16×16 拼接 512×32×32512×32×32up2512×32×32 拼接 256×64×64256×64×64up3256×64×64 拼接 128×128×128128×128×128up4128×128×128 拼接 64×256×25664×256×256outc64×256×2561×256×256验证脚本就是构造随机张量跑一次前向确认输出尺寸与输入一致if __name__ __main__: model UNet(n_channels3, n_classes1) x torch.randn(1, 3, 256, 256) y model(x) print(输出尺寸:, tuple(y.shape)) # 期望 (1, 1, 256, 256)输出和输入同为 256×256说明这是一个全分辨率的逐像素预测结构。尺寸对不上时优先检查输入高宽是否被 16 整除——四次下采样每次除以 22 的 4 次方是 16。这个问题在后续接真实数据时同样会出现数据加载阶段就该统一做 resize 或 padding而不是在模型里打补丁。4. torchsummary 可视化一行命令看清每一层输出与参数量4.1 安装并调用 torchsummarytorchsummary 是最轻量的网络结构可视化工具不依赖绘图库安装一条命令完成pip install torchsummary。调用方式from torchsummary import summary model UNet(n_channels3, n_classes1) summary(model, input_size(3, 256, 256), batch_size1, devicecpu)两个参数容易踩坑。input_size只写通道×高×宽不包含 batch 维batch 由batch_size单独控制设为 1 时输出里的形状第一个数字就是 1看起来更直观device参数默认是cudaCPU 环境不传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 大约 118MB31,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设为 1forward 最后一层不接 sigmoid直接配nn.BCEWithLogitsLoss它内部做了数值稳定的 sigmoid 加交叉熵计算比手动 sigmoid 加 BCELoss 稳定。多分类分割把n_classes设为类别数 K用nn.CrossEntropyLoss。注意 mask 的 dtype 要和 loss 期望一致BCE 需要 floatCrossEntropy 需要 long。5.4 端到端收敛的最小冒烟测试确认结构后用随机数据跑几十步验证前向、反向、优化器整个链路model UNet(n_channels3, n_classes1) criterion nn.BCEWithLogitsLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-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(floss {loss.item():.4f})在随机数据上 loss 持续下降是正常的说明模型有学习能力可以放心换真实数据集。torchsummary 只验证结构不能验证收敛结构正确之后真正决定分割上限的往往是数据预处理、类别不平衡和 loss 权重这三项比网络本身更值得花时间。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →