PyTorch实现UNet图像分割:Shape对齐与卷积递推详解
2026/9/19 0:14:52 网站建设 项目流程

简介:这是一份面向图像分割入门者的Unet网络PyTorch实现资料,主要解决“想看懂Unet结构却不知如何下手”的问题。内容围绕U形对称架构展开:左侧通过卷积与最大池化逐级下采样以提取高层语义特征,右侧利用最近邻上采样与跳跃连接逐步恢复空间分辨率和边缘细节,最终实现端到端的像素级预测。代码中封装了default_conv、default_relu、Up_Sample以及Unet主类,并设置固定随机种子保证实验可复现;下载后无需额外配置即可直接运行,同时可借助torchsummary直观查看每层输出尺寸与参数量。资源包为单个PDF文件,大小仅89KB,适合快速查阅核心代码与结构示意图。目前已有6923人学习浏览,对正在复现图像分割模型或准备论文实验的PyTorch使用者很有参考价值。

1. 一份 572x572 输入直接跑通的 Unet,先解决的是 shape 对齐问题

在把 Unet 从示意图变成能跑的代码时,最先卡人的不是 U 形结构本身,而是 3x3 卷积不加 padding 之后,左侧特征图会比右侧上采样结果大出几个像素,两侧拼接时 shape 对不上。这份代码把答案写在 Up_Sample 里:nearest 上采样放大两倍、1x1 卷积降通道、对左侧特征做中心裁剪,然后再拼到一起。输入单通道 572x572 图,输出 2 通道分割图,torchsummary 会直接把每层 shape 和 28,941,698 个可训练参数打出来。它适合两类人:一类是刚配好 PyTorch 环境、想跑通 unet 图像分割练手的人;另一类是在已有分割模型上做改动、需要快速核对网络尺寸递推的开发者。

2. 编码器下采样链:default_conv 与 MaxPool2d(kernel_size=1, stride=2) 的尺寸递推

2.1 左侧四个 stage 的真实结构

这段代码的编码器没有 BatchNorm、没有 Dropout,每个 stage 就是两个 Conv2d 加一个 ReLU。left1 到 left4 把通道从 64 一路翻到 512,bottom 再翻到 1024。卷积核固定是 3x3,padding 显式设为 0,这意味着每过一个 3x3 卷积,特征图每个方向会缩小 2 个像素。

def default_conv(in_channels, out_channels, kernel_size, bias=True): # padding=0 是刻意为之,两个 conv 后特征图尺寸会减 4 return nn.Conv2d(in_channels, out_channels, kernel_size, padding=0, bias=bias) def default_relu(): return nn.ReLU(inplace=True) left1 = [conv(in_channels, n_feats, 3), relu(), conv(n_feats, n_feats, 3)] left2 = [conv(n_feats, 2 * n_feats, 3), relu(), conv(2 * n_feats, 2 * n_feats, 3)] # left3 / left4 同构,分别输出 4*n_feats 与 8*n_feats 通道

n_feats默认是 64,它控制整个网络的宽度,是 Unet 里最值得调的参数。第一个卷积负责把输入通道映射到 64,之后每次池化前把通道乘 2。这里每个 stage 都是两个 3x3 卷积,而不是像 VGG 那样堆更多层,是因为 Unet 的下采样分支只需要中等表达力,更多卷积层会把感受野扩得过大,反而丢失边缘细节。

2.2 下采样为什么是 MaxPool2d(kernel_size=1, stride=2)

编码器部分最反直觉的写法是下采样,它的实现不是常见的nn.MaxPool2d(2),而是:

down = [] for layer in range(4): down.append(nn.MaxPool2d(kernel_size=1, stride=2)) self.down = nn.Sequential(*down)

kernel_size=1时池化窗口里只有一个元素,所以它并不是在 2x2 区域里取最大值,而是每隔一个像素取一个值,等价于x[:, :, ::2, ::2]。它和MaxPool2d(2, 2)的输出形状完全一样,都是宽高减半,但信息保留策略不同:前者直接丢弃一半像素,后者在每个窗口内保留响应最强的值。这个写法在 PyTorch 里是合法的,也维持了“每下采样一次,分辨率减半”的约束。如果改造成自己的项目,想换成更平滑的下采样,可以直接替换成 stride=2 的卷积,例如nn.Conv2d(64, 128, 3, stride=2, padding=1),输出形状不变,梯度传递会更稳定。

2.3 从 572 到 28:一张表看完整左侧递推

以输入 572x572、单通道为例,左侧编码器的尺寸变化如下:

模块两个 3x3 卷积后输出MaxPool(1, 2) 后输出
left1(64, 568, 568)(64, 284, 284)
left2(128, 280, 280)(128, 140, 140)
left3(256, 136, 136)(256, 68, 68)
left4(512, 64, 64)(512, 32, 32)
bottom(1024, 28, 28)

572 是原论文 overlap-tile 策略里常用的输入尺寸,四次池化后落到 28x28,分辨率缩小约 20 倍,通道数从 64 涨到 1024。这个 28x28 的底层特征图是后续所有上采样分支的起点,右边每一层拼接都要回到这张表里的对应尺寸。

2.4 没有 BN 的卷积块在实际训练里的影响

这份代码刻意省略了 BatchNorm,好处是参数结构一目了然,便于复现和核对;坏处是激活分布完全由输入尺度和初始化决定,学习率稍大,深层 1024 通道的方差就容易失控。我一般在把这个骨架接到自己数据集时,会在每个 conv 对之间插入 BN。

block = nn.Sequential( nn.Conv2d(64, 64, 3, padding=0), nn.BatchNorm2d(64), # 每个通道增加 2 个可学习参数 nn.ReLU(inplace=True), )

加入 BN 后总参数量只增加很小一部分,但收敛稳定性明显改善。另外要明确一点:这份代码没有数据加载和训练循环,它的核心价值是把模型定义、前向路径、参数可视化串成一条完整链路,确认网络可运行之后,再接自己的 DataLoader 和损失函数。

3. Up_Sample 解码器:nearest 上采样、中心裁剪与 left/right 通道对齐

3.1 Up_Sample 内部的三段式变换

解码器里最核心的模块是自定义的 Up_Sample,它不是简单调一个上采样就完事,而是把“放大、降维、激活”打包在一起:

class Up_Sample(nn.Module): def __init__(self, in_channels, conv=default_conv, relu=default_relu): super(Up_Sample, self).__init__() up1 = nn.Upsample(scale_factor=2, mode='nearest') up2 = conv(in_channels, in_channels // 2, 1) self.module_up = nn.Sequential(up1, up2, relu()) def forward(self, input_down, input_left): x = self.module_up(input_down) dif = (input_left.shape[3] - x.shape[3]) / 2 input_left = input_left[:, :, int(dif):int(dif + x.shape[3]), int(dif):int(dif + x.shape[3])] return torch.cat((x, input_left), 1)

nn.Upsample(scale_factor=2, mode='nearest')把宽高各放大一倍,mode='nearest'表示最近邻插值,不产生新的像素值,也不增加参数。之后用一个 1x1 卷积把通道从in_channels降到in_channels // 2。这样做的目的是为拼接做准备:上采样分支降到一半通道后,与左侧分支同分辨率特征拼接,拼接后通道数正好翻倍,落入右侧卷积的输入通道范围。

3.2 中心裁剪 dif 公式:差值为什么要除 2

forward 里最容易被忽略的是裁剪行。由于左侧每个 stage 两个 3x3 卷积都不加 padding,left 路径特征图始终比右侧上采样结果大。以 bottom 与 left4 为例:bottom 输出 28x28,上采样后变成 56x56,而 left4 输出 64x64,差 8 像素,两侧各裁 4 像素。dif算的就是每侧要裁掉多少:

dif = (input_left.shape[3] - x.shape[3]) / 2 input_left = input_left[:, :, int(dif):int(dif + x.shape[3]), int(dif):int(dif + x.shape[3])]

input_left.shape[3]是左侧特征的高或宽,x.shape[3]是上采样结果的高或宽,差值除以 2 得到上下、左右各要裁掉的像素数。int()是为了处理差值为奇数的情况,但最好保证差值本来就是偶数,否则中心会偏移 1 像素。当前 572 输入下,四组差值 64-56、136-104、280-200、568-392 都是偶数,所以没有这个问题。

3.3 forward 顺序里的拼接链与通道规则

Unet 的 forward 顺序是沿着 U 形先下到底,再逐层上采样:

x1 = self.left1(x) x1d = self.down[0](x1) # x2 / x3 / x4 结构相同,x_b 为 bottom 输出 y4d = self.up[3](x_b, x4) y3 = self.right4(y4d) y3d = self.up[2](y3, x3) y2 = self.right3(y3d) y2d = self.up[1](y2, x2) y1 = self.right2(y2d) y1d = self.up[0](y1, x1) y = self.right1(y1d) out = self.tail(y)

每一层拼接的来源和通道变化可以整理成下面这张表:

调用上采样分支left 分支拼接后通道右侧 conv 输入
up[3]512x56x56512x64x64 裁剪为 561024right4
up[2]256x104x104256x136x136 裁剪为 104512right3
up[1]128x200x200128x280x280 裁剪为 200256right2
up[0]64x392x39264x568x568 裁剪为 392128right1

这就是代码里 right1 到 right4 的通道数和 left 的通道数看上去“错开一个位置”的原因:right4 的输入其实是拼接后的 1024 通道,而不是 left4 的 512 通道。理解这张表,之后想调整每层通道数就知道要同步改哪些地方。

3.4 nearest 与 ConvTranspose2d 的取舍

上采样分支选 nearest 而不是转置卷积,是一个很实际的选择。nearest 不引入可学习参数,不会产生转置卷积常见的棋盘格伪影,配合 1x1 卷积降维,是目前复现 Unet 时最稳妥的写法。转置卷积能学习空间插值,但要调的参数和超参数更多,数据量不够时反而容易学出噪声。如果后续要在基础版本上做 unet 模型改进,可以先保留 nearest,把 Up_Sample 里的 1x1 卷积换成 3x3 卷积,用少量参数换更强的上采样表达。

4. torchsummary 可视化:28,941,698 参数与 2.27GB 前向显存的读法

4.1 一行 summary 的调用约定

入口函数非常短,核心就一行:

def main(): model = Unet(in_channels=1, out_channels=2) # 灰度图输入,二分类输出 from torchsummary import summary summary(model.cuda(), (1, 572, 572)) # C, H, W,不含 batch 维

in_channels=1对应灰度图,out_channels=2对应分割任务的两个类别。summary的第二个参数是输入张量的 C、H、W,不包含 batch 维,它内部会构造一个 batch=1 的输入做一次前向。model.cuda()是因为默认在 GPU 上执行,如果机器没有 CUDA,就把model.cuda()去掉,改成summary(model, (1, 572, 572)),CPU 上也可以跑,只是 572x572 输入会稍微慢一点。安装依赖用pip install torchsummary即可,前提是当前环境已经装好支持 CUDA 的 PyTorch 版本。

4.2 参数分布:9.44M 的大头在哪里

torchsummary 的输出里最容易让人注意的是总参数量 28.94M,比常见的 ResNet 分类网络大不少。真正占参数的是几个大卷积层:

输出 Shape参数量说明
Conv2d-19[-1, 1024, 28, 28]9,438,208bottom 第二个 3x3 卷积
Conv2d-17[-1, 1024, 30, 30]4,719,616bottom 第一个 3x3 卷积
Conv2d-24[-1, 512, 54, 54]4,719,104right4 第一个 3x3 卷积
Conv2d-13[-1, 512, 66, 66]1,180,160left4 第一个 3x3 卷积

以 Conv2d-19 为例,输入输出都是 1024 通道,3x3 卷积核加上 bias 的参数个数是1024 * 1024 * 9 + 1024 = 9,438,208,和打印值完全一致。这个量级说明 Unet 的参数大头在通道最宽的编码器底部和解码器入口,而不是浅层。这个数字也可以作为改动网络后的核对基准:改 padding、加 BN、换卷积核大小,最后都会有对应的参数变化。

4.3 显存估算:Forward/backward 2,275.74MB 实际代表什么

torchsummary 倒数几行会打印Input sizeForward/backward pass sizeParams sizeEstimated Total Size。Input size 1.25MB 是输入张量本身,Params size 110.40MB 是权重驻留显存,而 Forward/backward 2,275.74MB 来自中间层激活值——尤其是 392x392、388x388 这些高分辨率特征平面,它们才是显存占用的主要来源。

这个 2.27GB 是 batch=1 的估算值,并不是实测显存,实际训练还要叠加梯度、优化器状态和 CUDA context,因此 6GB 显存的卡建议保持 batch=1,不要贸然加大输入尺寸。如果显存紧张,比较直接的做法是把输入从 572 缩小,同时把模型输出的分割图尺寸变化纳入考虑,激活值会按面积比例明显下降。训练时开 AMP 混合精度也能显著降低激活显存,这在 PyTorch 1.6 之后已经是标准操作。

4.4 torchinfo 是更可读的替代品

torchsummary 打印的信息偏平铺,torchinfo 按模块层级缩进显示,且支持显式指定 batch 维,观感更接近 PyTorch 官方文档:

# pip install torchinfo from torchinfo import summary as ts model = Unet(in_channels=1, out_channels=2) ts(model, input_size=(1, 1, 572, 572)) # batch=1, C=1, H=W=572

除了可视化,还可以用一次真实 forward 验证网络输出:

model.eval() with torch.no_grad(): y = model(torch.randn(1, 1, 572, 572)) print(y.shape) # torch.Size([1, 2, 388, 388])

这一步对后续数据加载很关键:模型输出是 2 通道 388x388,标签 mask 也必须预先处理成同样的尺寸,否则损失函数会直接报 shape 不匹配。如果输出尺寸和预期不一致,优先检查每个 Up_Sample 的裁剪差值是否为偶数,以及输入边长是否满足递推条件。

5. 同尺寸输入输出变体:padding=1 加 16 倍数输入,去掉中心裁剪

5.1 572 输入对应 388 输出,输出尺寸由谁决定

上面的一次真实 forward 已经验证,572 输入最终输出 388,和原论文 overlap-tile 策略一致。也就是说,用这份代码处理 512x512 的原图时,不能直接把模型输出和原尺寸标签算损失,需要先把标签从中心裁剪到模型输出尺寸,或者把输入统一做成 572x572。输出尺寸不是简单线性换算出来的,它由每一层卷积的边界损耗累积决定,所以更换输入尺寸后最可靠的检查方式是直接跑一次 forward 看y.shape,不要凭经验估计。

5.2 把 padding 改成 1,并让输入边长是 16 的倍数

如果不想每次都处理中心裁剪,可以把网络改成输入输出同尺寸。修改点有两处:一是把default_conv的 padding 从 0 改成 1,二是删掉 Up_Sample 里的裁剪逻辑。

def default_conv(in_channels, out_channels, kernel_size, bias=True): # padding=1 后,3x3 卷积不改变特征图宽高 return nn.Conv2d(in_channels, out_channels, kernel_size, padding=1, bias=bias) class Up_Sample(nn.Module): def forward(self, input_down, input_left): x = self.module_up(input_down) # 输入边长为 16 的倍数时,两个分支宽高完全一致,无需裁剪 return torch.cat((x, input_left), 1)

padding=1 后,每个 3x3 卷积前后尺寸保持一致。此时输入边长必须是 16 的倍数,因为编码器有四次下采样,任何一次不能整除都会导致上采样后与 left 分支差 1 像素。以 576 输入为例,池化序列是 576 -> 288 -> 144 -> 72 -> 36,bottom 上采样后正好也是 72,与 left4 完全对齐。572 不是 16 的倍数,会落到 35 -> 70 与 left4 的 71 差 1 像素,int(dif)会把裁剪窗口取偏,所以这个变体不要用 572。

5.3 用 summary 验证变体

改完以后,把 main 里的输入换成 576 再跑一次:

model = Unet(in_channels=1, out_channels=2) summary(model.cuda(), (1, 576, 576)) # 最后一层 Conv2d 的输出 shape 应为 [-1, 2, 576, 576]

输出从 388 变成 576,是因为 padding=1 不再损失任何边界像素,模型从“输入比输出大 184 像素”变成“输入输出同尺寸”,这对监督分割任务来说方便很多:标签图不需要再做中心裁剪,直接和模型输出对齐即可。代价是每一层的激活值面积都比原版大,显存占用会随之上升。想要严格复现论文设计就保留原版 padding=0 和中心裁剪;想在同一个输入输出尺寸上跑分割训练,padding=1 加 16 倍数输入是更省心的组合。

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

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

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

立即咨询