ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

PyTorch实现UNet图像分割:Shape对齐与卷积递推详解

PyTorch实现UNet图像分割:Shape对齐与卷积递推详解 简介这是一份面向图像分割入门者的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_size1, stride2) 的尺寸递推2.1 左侧四个 stage 的真实结构这段代码的编码器没有 BatchNorm、没有 Dropout每个 stage 就是两个 Conv2d 加一个 ReLU。left1 到 left4 把通道从 64 一路翻到 512bottom 再翻到 1024。卷积核固定是 3x3padding 显式设为 0这意味着每过一个 3x3 卷积特征图每个方向会缩小 2 个像素。def default_conv(in_channels, out_channels, kernel_size, biasTrue): # padding0 是刻意为之两个 conv 后特征图尺寸会减 4 return nn.Conv2d(in_channels, out_channels, kernel_size, padding0, biasbias) def default_relu(): return nn.ReLU(inplaceTrue) 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_size1, stride2)编码器部分最反直觉的写法是下采样它的实现不是常见的nn.MaxPool2d(2)而是down [] for layer in range(4): down.append(nn.MaxPool2d(kernel_size1, stride2)) self.down nn.Sequential(*down)kernel_size1时池化窗口里只有一个元素所以它并不是在 2x2 区域里取最大值而是每隔一个像素取一个值等价于x[:, :, ::2, ::2]。它和MaxPool2d(2, 2)的输出形状完全一样都是宽高减半但信息保留策略不同前者直接丢弃一半像素后者在每个窗口内保留响应最强的值。这个写法在 PyTorch 里是合法的也维持了“每下采样一次分辨率减半”的约束。如果改造成自己的项目想换成更平滑的下采样可以直接替换成 stride2 的卷积例如nn.Conv2d(64, 128, 3, stride2, padding1)输出形状不变梯度传递会更稳定。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, padding0), nn.BatchNorm2d(64), # 每个通道增加 2 个可学习参数 nn.ReLU(inplaceTrue), )加入 BN 后总参数量只增加很小一部分但收敛稳定性明显改善。另外要明确一点这份代码没有数据加载和训练循环它的核心价值是把模型定义、前向路径、参数可视化串成一条完整链路确认网络可运行之后再接自己的 DataLoader 和损失函数。3. Up_Sample 解码器nearest 上采样、中心裁剪与 left/right 通道对齐3.1 Up_Sample 内部的三段式变换解码器里最核心的模块是自定义的 Up_Sample它不是简单调一个上采样就完事而是把“放大、降维、激活”打包在一起class Up_Sample(nn.Module): def __init__(self, in_channels, convdefault_conv, reludefault_relu): super(Up_Sample, self).__init__() up1 nn.Upsample(scale_factor2, modenearest) 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_factor2, modenearest)把宽高各放大一倍modenearest表示最近邻插值不产生新的像素值也不增加参数。之后用一个 1x1 卷积把通道从in_channels降到in_channels // 2。这样做的目的是为拼接做准备上采样分支降到一半通道后与左侧分支同分辨率特征拼接拼接后通道数正好翻倍落入右侧卷积的输入通道范围。3.2 中心裁剪 dif 公式差值为什么要除 2forward 里最容易被忽略的是裁剪行。由于左侧每个 stage 两个 3x3 卷积都不加 paddingleft 路径特征图始终比右侧上采样结果大。以 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 裁剪为 561024right4up[2]256x104x104256x136x136 裁剪为 104512right3up[1]128x200x200128x280x280 裁剪为 200256right2up[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_channels1, out_channels2) # 灰度图输入二分类输出 from torchsummary import summary summary(model.cuda(), (1, 572, 572)) # C, H, W不含 batch 维in_channels1对应灰度图out_channels2对应分割任务的两个类别。summary的第二个参数是输入张量的 C、H、W不包含 batch 维它内部会构造一个 batch1 的输入做一次前向。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 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 是 batch1 的估算值并不是实测显存实际训练还要叠加梯度、优化器状态和 CUDA context因此 6GB 显存的卡建议保持 batch1不要贸然加大输入尺寸。如果显存紧张比较直接的做法是把输入从 572 缩小同时把模型输出的分割图尺寸变化纳入考虑激活值会按面积比例明显下降。训练时开 AMP 混合精度也能显著降低激活显存这在 PyTorch 1.6 之后已经是标准操作。4.4 torchinfo 是更可读的替代品torchsummary 打印的信息偏平铺torchinfo 按模块层级缩进显示且支持显式指定 batch 维观感更接近 PyTorch 官方文档# pip install torchinfo from torchinfo import summary as ts model Unet(in_channels1, out_channels2) ts(model, input_size(1, 1, 572, 572)) # batch1, C1, HW572除了可视化还可以用一次真实 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. 同尺寸输入输出变体padding1 加 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, biasTrue): # padding1 后3x3 卷积不改变特征图宽高 return nn.Conv2d(in_channels, out_channels, kernel_size, padding1, biasbias) 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)padding1 后每个 3x3 卷积前后尺寸保持一致。此时输入边长必须是 16 的倍数因为编码器有四次下采样任何一次不能整除都会导致上采样后与 left 分支差 1 像素。以 576 输入为例池化序列是 576 - 288 - 144 - 72 - 36bottom 上采样后正好也是 72与 left4 完全对齐。572 不是 16 的倍数会落到 35 - 70 与 left4 的 71 差 1 像素int(dif)会把裁剪窗口取偏所以这个变体不要用 572。5.3 用 summary 验证变体改完以后把 main 里的输入换成 576 再跑一次model Unet(in_channels1, out_channels2) summary(model.cuda(), (1, 576, 576)) # 最后一层 Conv2d 的输出 shape 应为 [-1, 2, 576, 576]输出从 388 变成 576是因为 padding1 不再损失任何边界像素模型从“输入比输出大 184 像素”变成“输入输出同尺寸”这对监督分割任务来说方便很多标签图不需要再做中心裁剪直接和模型输出对齐即可。代价是每一层的激活值面积都比原版大显存占用会随之上升。想要严格复现论文设计就保留原版 padding0 和中心裁剪想在同一个输入输出尺寸上跑分割训练padding1 加 16 倍数输入是更省心的组合。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进