简介:这是一份面向图像分割入门者的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] | 512x56x56 | 512x64x64 裁剪为 56 | 1024 | right4 |
| up[2] | 256x104x104 | 256x136x136 裁剪为 104 | 512 | right3 |
| up[1] | 128x200x200 | 128x280x280 裁剪为 200 | 256 | right2 |
| up[0] | 64x392x392 | 64x568x568 裁剪为 392 | 128 | right1 |
这就是代码里 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,208 | bottom 第二个 3x3 卷积 |
| Conv2d-17 | [-1, 1024, 30, 30] | 4,719,616 | bottom 第一个 3x3 卷积 |
| Conv2d-24 | [-1, 512, 54, 54] | 4,719,104 | right4 第一个 3x3 卷积 |
| Conv2d-13 | [-1, 512, 66, 66] | 1,180,160 | left4 第一个 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 size、Forward/backward pass size、Params size和Estimated 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 倍数输入是更省心的组合。
本文还有配套的精品资源,点击获取