
U-Net 这个网络结构我最早是在做医学影像分割的项目里接触到的。当时用现成的分割库跑通不难但一旦要改结构、换损失函数、或者排查维度对不上的报错就发现如果不亲手把每一层的张量形状推一遍根本没法定位问题。后来我干脆找了个下午从零开始用 PyTorch 把 U-Net 一行行敲出来每写一层就打印一次 shape把数据在编码器、瓶颈层、解码器里的维度变化彻底摸清楚。这篇就是那次实践的完整记录适合已经会一点 PyTorch、能看懂卷积和池化、但一遇到 skip connection 拼接就犯迷糊的人。我会把每个模块为什么这么设计、维度怎么算、拼接时到底拼在哪个维度上全部讲透代码可以直接抄下来跑。1. 先把 U-Net 的骨架在脑子里搭起来1.1 它到底解决的是什么问题U-Net 最初是为生物医学图像分割设计的任务本质是逐像素分类输入一张图输出一张同样大小的掩码图每个像素标记它属于哪一类。这跟图像分类完全不同分类最后输出的是一个向量而分割要求输出保留空间分辨率。问题就来了——卷积和池化会不断缩小特征图的空间尺寸最后你怎么把尺寸还原回去同时又不丢失细节U-Net 的答案是一个对称的编码器-解码器结构外加一个关键设计跳跃连接skip connection。编码器负责下采样、提取语义特征解码器负责上采样、恢复分辨率而跳跃连接把编码器里高分辨率的浅层特征直接送到解码器对应层弥补上采样过程中丢失的细节。整个结构画出来像字母 U所以叫 U-Net。1.2 为什么是对称的你去看 U-Net 的原始结构图会发现左边下采样几次右边就上采样几次层数是对称的。这不是为了好看而是有实际原因的。编码器每下采样一次特征图边长减半、通道数翻倍解码器每上采样一次边长翻倍、通道数减半。只有两边对称最后输出的特征图才能恢复到和输入相同的空间尺寸通道数也才能收敛到你要的类别数。如果不对称比如编码器下采样了 4 次解码器只上采样了 3 次那输出尺寸就只有输入的 1/2根本没法做逐像素的损失计算。所以对称性是功能需求不是审美需求。1.3 数据维度变化是理解 U-Net 的主线我强烈建议你在学 U-Net 的时候把维度变化当成一条主线。整个网络里张量的形状遵循一个非常规律的节奏编码器阶段空间尺寸H, W每次减半通道数C每次翻倍瓶颈层空间尺寸最小通道数最多解码器阶段空间尺寸每次翻倍通道数每次减半跳跃连接把编码器某层的特征图在通道维度上和解码器对应层拼接只要你能把这条主线在脑子里跑通任何一层维度对不上你都能立刻定位到是哪一步出了问题。下面我就按这个节奏一层层写代码、一层层验证维度。2. 编码器部分下采样与通道翻倍的实现细节2.1 双重卷积块DoubleConv的设计逻辑U-Net 的每个编码阶段核心是一个双重卷积块也就是连续两个 3x3 卷积每个卷积后面跟一个 ReLU。为什么是两次卷积而不是一次因为两次 3x3 卷积的感受野等价于一次 5x5 卷积但参数量更少2×3²18 对比 5²25而且多了一次非线性激活表达能力更强。这是 VGG 网络验证过的经典结论U-Net 直接沿用了。在写代码时有一个细节必须注意padding 要设为 1。3x3 卷积、stride 为 1 的情况下padding1 才能保证输出的空间尺寸和输入一致。如果不加 padding每卷一次边长就减 2几层下来图就没了跳跃连接拼接时尺寸也对不上。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x)这里我加了 BatchNorm原始论文里没有但实际训练时加上它收敛更稳尤其是 batch size 比较小的时候。如果你要严格复现原论文可以把 BN 去掉但实测下来加上更好。2.2 下采样到底用最大池化还是步长卷积编码器的下采样原始 U-Net 用的是 2x2 最大池化stride 为 2。这个操作把空间尺寸精确地减半而且不引入任何参数。另一种常见做法是用 stride2 的卷积来代替池化好处是下采样的过程也可学习但会引入额外参数而且尺寸计算稍微麻烦一点。我两种都试过在医学图像这种细节敏感的任务上最大池化反而更稳因为它保留了每个 2x2 窗口里的最大值相当于保留了最显著的特征响应。步长卷积在某些数据集上能涨一点点但不够稳定。所以下面的实现我用最大池化。class Encoder(nn.Module): def __init__(self, in_channels, features): super().__init__() self.conv DoubleConv(in_channels, features) self.pool nn.MaxPool2d(kernel_size2, stride2) def forward(self, x): skip self.conv(x) # 保存给跳跃连接用 down self.pool(skip) # 下采样后传给下一层 return skip, down注意这里 forward 返回了两个值一个是卷积后、池化前的特征图skip一个是池化后的特征图down。skip 就是留给解码器拼接用的down 继续往下一层传。这个设计是整个 U-Net 维度管理的关键你必须在编码器每一层都把 skip 存下来。2.3 逐层验证编码器的维度变化假设输入是一张 1 通道、572x572 的灰度图原始论文的输入尺寸我们走一遍编码器看看维度怎么变。为了直观我用一个小的输入来演示x torch.randn(1, 1, 572, 572) # batch1, channel1, H572, W572 # 第1层DoubleConv(1, 64)池化 enc1 Encoder(1, 64) skip1, down1 enc1(x) print(skip1:, skip1.shape) # [1, 64, 572, 572] print(down1:, down1.shape) # [1, 64, 286, 286] # 第2层DoubleConv(64, 128)池化 enc2 Encoder(64, 128) skip2, down2 enc2(down1) print(skip2:, skip2.shape) # [1, 128, 286, 286] print(down2:, down2.shape) # [1, 128, 143, 143]你可以看到规律非常清晰每次经过 DoubleConv通道数翻倍、空间尺寸不变每次经过池化空间尺寸减半、通道数不变。skip 保存的是池化前的尺寸down 是池化后的尺寸。这个规律会一直持续到瓶颈层。提示如果你用的是 572x572 这种不能被 2 整除多次的尺寸池化时向下取整最后解码器上采样回来可能会差几个像素。所以实际项目中我一般把输入 resize 到 2 的幂次方比如 512x512 或 256x256这样维度计算干净拼接时不会出现尺寸不匹配。3. 瓶颈层与解码器上采样的维度还原3.1 瓶颈层为什么通道最多、尺寸最小编码器一路下采样到底就到了瓶颈层。这一层不再池化只做 DoubleConv通道数达到最大原始论文是 1024空间尺寸最小。它的作用是提取全局的、抽象的语义信息。你可以理解为到了这一层网络已经看懂了图像里有什么但还不知道这些东西具体在哪个像素位置。瓶颈层的维度变化很简单输入是编码器最后一层的 down经过 DoubleConv 后通道翻倍尺寸不变。class Bottleneck(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv DoubleConv(in_channels, out_channels) def forward(self, x): return self.conv(x) # 接上面的例子 bottleneck Bottleneck(128, 256) b bottleneck(down2) print(bottleneck:, b.shape) # [1, 256, 143, 143]3.2 ConvTranspose2d 上采样的原理与维度计算解码器的核心是上采样。U-Net 原始论文用的是转置卷积也叫反卷积PyTorch 里对应nn.ConvTranspose2d。很多人对它的维度计算感到困惑我这里把公式讲清楚。对于ConvTranspose2d(in_channels, out_channels, kernel_size2, stride2)输出尺寸的计算公式是H_out (H_in - 1) × stride - 2 × padding kernel_size当 kernel_size2, stride2, padding0 时公式简化为H_out (H_in - 1) × 2 - 0 2 2 × H_in也就是说输出尺寸正好是输入的两倍。这就是为什么解码器每上采样一次空间尺寸就翻倍和编码器的池化正好对称。class UpConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up nn.ConvTranspose2d(in_channels, out_channels, kernel_size2, stride2) def forward(self, x): return self.up(x) # 验证 up UpConv(256, 128) u up(b) print(upsampled:, u.shape) # [1, 128, 286, 286]注意上采样后通道数减半256→128空间尺寸翻倍143→286。这个结果正好和编码器第 2 层的 skip2 尺寸 [1, 128, 286, 286] 一致可以拼接。3.3 跳跃连接拼接拼在通道维度上这是 U-Net 最容易出错的地方。上采样后的特征图和编码器对应层的 skip 特征图要在通道维度上拼接而不是空间维度。拼接后通道数相加空间尺寸不变。# u 的 shape: [1, 128, 286, 286] # skip2 的 shape: [1, 128, 286, 286] concat torch.cat([skip2, u], dim1) print(concat:, concat.shape) # [1, 256, 286, 286]拼接后通道数变成 128128256然后接一个 DoubleConv 把通道数降回 128同时融合两部分信息。class Decoder(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up nn.ConvTranspose2d(in_channels, out_channels, kernel_size2, stride2) self.conv DoubleConv(out_channels * 2, out_channels) def forward(self, x, skip): x self.up(x) # 如果尺寸有细微差异做一次裁剪或填充对齐 if x.shape ! skip.shape: x torch.nn.functional.interpolate( x, sizeskip.shape[2:], modebilinear, align_cornersTrue) x torch.cat([skip, x], dim1) return self.conv(x)这里我加了一个尺寸对齐的保险逻辑。虽然理论上对称结构尺寸应该完全匹配但实际中如果输入尺寸不是 2 的幂次方或者用了不同的 padding 策略就可能差一两个像素。加上这个判断能避免运行时直接报错。注意torch.cat的dim1是通道维度。如果你不小心写成dim0那就是在 batch 维度拼接会直接把 batch size 翻倍后面所有计算全乱。这个错误我见过不止一次排查时先看 cat 的 dim 参数。4. 完整网络组装与端到端维度追踪4.1 把编码器、瓶颈、解码器串起来现在把前面的模块组装成完整的 U-Net。我用一个 4 层下采样的版本输入假设是 1 通道、256x256。class UNet(nn.Module): def __init__(self, in_channels1, num_classes2, features[64, 128, 256, 512]): super().__init__() self.encoders nn.ModuleList() self.decoders nn.ModuleList() # 编码器 for feature in features: self.encoders.append(Encoder(in_channels, feature)) in_channels feature # 瓶颈层 self.bottleneck Bottleneck(features[-1], features[-1] * 2) # 解码器逆序 for feature in reversed(features): self.decoders.append(Decoder(feature * 2, feature)) # 最后的 1x1 卷积输出类别数 self.final_conv nn.Conv2d(features[0], num_classes, kernel_size1) def forward(self, x): skips [] for encoder in self.encoders: skip, x encoder(x) skips.append(skip) x self.bottleneck(x) skips skips[::-1] # 反转让解码器从最深层开始取 for decoder, skip in zip(self.decoders, skips): x decoder(x, skip) return self.final_conv(x)4.2 端到端跑一遍打印每一层维度光看代码不够我们实际跑一遍把每一层的维度打出来。这是理解 U-Net 最有效的方式。model UNet(in_channels1, num_classes2) x torch.randn(1, 1, 256, 256) skips [] for i, encoder in enumerate(model.encoders): skip, x encoder(x) skips.append(skip) print(f编码器{i1} skip: {skip.shape}, down: {x.shape}) x model.bottleneck(x) print(f瓶颈层: {x.shape}) skips skips[::-1] for i, (decoder, skip) in enumerate(zip(model.decoders, skips)): x decoder(x, skip) print(f解码器{i1}: {x.shape}) out model.final_conv(x) print(f最终输出: {out.shape})运行结果会是这样阶段操作输出维度编码器1DoubleConv Poolskip: [1,64,256,256], down: [1,64,128,128]编码器2DoubleConv Poolskip: [1,128,128,128], down: [1,128,64,64]编码器3DoubleConv Poolskip: [1,256,64,64], down: [1,256,32,32]编码器4DoubleConv Poolskip: [1,512,32,32], down: [1,512,16,16]瓶颈层DoubleConv[1,1024,16,16]解码器1Up Concat Conv[1,512,32,32]解码器2Up Concat Conv[1,256,64,64]解码器3Up Concat Conv[1,128,128,128]解码器4Up Concat Conv[1,64,256,256]最终输出1x1 Conv[1,2,256,256]看到没有输出尺寸 [1, 2, 256, 256] 和输入 [1, 1, 256, 256] 的空间尺寸完全一致通道数变成了类别数 2。这就是 U-Net 的完整数据流。4.3 维度不匹配时的排查思路实际写的时候最常见的报错就是torch.cat时尺寸对不上。我的排查顺序是这样的先看报错信息里的两个 shape确认是空间尺寸不一致还是通道数不一致。如果是空间尺寸差 1-2 个像素大概率是输入尺寸不能被 2 整除多次或者某层 padding 设置不对。解决办法是把输入 resize 到 2 的幂次方或者在 cat 前用 interpolate 对齐。如果是通道数不一致检查 ConvTranspose2d 的 out_channels 是否和对应 skip 的通道数相等。解码器的 up 输出通道必须等于同层 skip 的通道数否则 cat 后通道数不对下一层 DoubleConv 的 in_channels 也会错。如果是 batch 维度不一致检查 cat 的 dim 是不是写成了 0。这套排查逻辑我用了很多次基本能覆盖 90% 的维度报错。5. 训练相关的几个实操要点5.1 损失函数的选择分割任务最常用的损失是交叉熵但医学图像里经常有类别极度不平衡的问题比如病灶区域只占图像的百分之几。这时候纯交叉熵会让模型倾向于全预测背景。我一般用 Dice Loss 或者交叉熵和 Dice 的组合。class DiceLoss(nn.Module): def __init__(self, smooth1e-6): super().__init__() self.smooth smooth def forward(self, pred, target): pred torch.sigmoid(pred) pred pred.view(-1) target target.view(-1) intersection (pred * target).sum() dice (2. * intersection self.smooth) / (pred.sum() target.sum() self.smooth) return 1 - diceDice 系数的直观理解是预测区域和真实区域的重叠程度。完全重叠是 1完全不重叠是 0。作为损失就是 1 减去它。这个损失对类别不平衡不敏感因为它是按区域算的不是按像素平均的。5.2 输入尺寸与 batch size 的权衡U-Net 的显存占用和输入尺寸的平方成正比。256x256 的输入batch size 可以开到 8-16512x512 的话可能只能开到 2-4。我的经验是如果显存有限优先保证输入尺寸因为分割任务对分辨率很敏感batch size 小一点可以用梯度累积来补偿。# 梯度累积示例 accumulation_steps 4 optimizer.zero_grad() for i, (images, masks) in enumerate(dataloader): outputs model(images) loss criterion(outputs, masks) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这样等效于把 batch size 放大了 4 倍但显存占用不变。5.3 数据增强的注意事项分割任务的数据增强必须保证图像和掩码做完全相同的变换。旋转、翻转、缩放都可以但要注意插值方式图像用双线性插值掩码必须用最近邻插值否则会出现不存在的类别值。import torchvision.transforms.functional as TF import random def augment(image, mask): # 随机水平翻转 if random.random() 0.5: image TF.hflip(image) mask TF.hflip(mask) # 随机旋转 angle random.uniform(-15, 15) image TF.rotate(image, angle) mask TF.rotate(mask, angle) return image, mask提示掩码旋转时TF.rotate默认用最近邻插值这是对的。如果你手动指定了双线性插值掩码边缘会出现 0.5 这种非整数类别值训练时交叉熵会报错或者算出莫名其妙的结果。6. 我踩过的几个坑和对应的解法6.1 上采样后尺寸差一个像素有一次我用 500x500 的输入编码器池化了 4 次尺寸依次是 250、125、62、31。解码器上采样回来是 62、124、248、496最后和输入的 500 差了 4 个像素。跳跃连接拼接时直接报错。解法有两个一是把输入 resize 到 512x512所有尺寸都是 2 的幂次方干净利落二是在 Decoder 里加尺寸对齐逻辑用 interpolate 把上采样结果拉到和 skip 一样的尺寸。我现在的习惯是两者都做输入尽量规整代码里也保留对齐保险。6.2 BatchNorm 在小 batch 下的问题BatchNorm 在 batch size 小于 4 的时候统计量估计不准训练会不稳定。如果你显存不够只能用很小的 batch建议把 BatchNorm 换成 GroupNorm它对 batch size 不敏感。# 把 BatchNorm2d 换成 GroupNorm nn.GroupNorm(num_groups8, num_channelsout_channels)GroupNorm 把通道分成若干组在每组内部做归一化不依赖 batch 维度。实测在小 batch 场景下比 BatchNorm 稳很多。6.3 转置卷积的棋盘格伪影ConvTranspose2d有一个已知问题当 kernel_size 不能被 stride 整除时输出会出现棋盘格状的伪影。U-Net 里用的是 kernel_size2、stride2正好整除所以一般不会有这个问题。但如果你改成 kernel_size3、stride2就要小心了。替代方案是先做最近邻上采样再用普通卷积。这样没有棋盘格问题而且参数量更少。class UpConvAlternative(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up nn.Sequential( nn.Upsample(scale_factor2, modenearest), nn.Conv2d(in_channels, out_channels, kernel_size3, padding1) ) def forward(self, x): return self.up(x)我对比过两种方式在大多数任务上效果差不多但最近邻上采样卷积更稳定不容易出伪影。如果你对转置卷积的伪影比较敏感可以优先用这个方案。6.4 最后的 1x1 卷积不要忘了解码器最后一层输出的通道数是 features[0]也就是 64。但你的类别数可能是 2 或者更多。所以最后必须接一个 1x1 卷积把通道数映射到类别数。这个 1x1 卷积不改变空间尺寸只改变通道数。我见过有人忘了这一步输出通道是 64然后拿去做交叉熵直接报错。7. 关于 U-Net 改进的一些个人看法现在网上有很多 U-Net 的变体比如 U-Net、Attention U-Net、ResU-Net 等等。我的建议是先把原始 U-Net 吃透把维度变化、跳跃连接、上采样这些基础打牢再去改。因为所有改进本质上都是在某几个环节做文章要么改跳跃连接的方式比如加注意力要么改卷积块的结构比如加残差要么改上采样的策略。你基础不牢改出来的东西维度都对不上更别说调参了。如果你要动手改进我推荐从两个方向入手。第一个是在跳跃连接上加注意力机制让网络自己决定哪些浅层特征更重要。第二个是把 DoubleConv 换成残差块缓解深层网络的梯度消失。这两个改动都不复杂而且效果通常比较明显。至于 Transformer 和 U-Net 的结合那是另一个话题了。Swin-UNet 这类结构确实在部分任务上超过了纯卷积的 U-Net但参数量和计算量也上去了。如果你的数据量不大纯 U-Net 加上合适的数据增强往往比硬上 Transformer 更划算。最后分享一个我自己的习惯每次写完一个新的网络结构我都会用一个随机张量跑一遍 forward把每一层的 shape 打印出来和纸上推导的结果对一遍。这个习惯帮我省了无数调试时间。U-Net 这种维度变化规律性很强的网络尤其适合用这种方式验证。你把上面那份端到端维度追踪的代码跑一遍对照表格看一遍基本就再也不会在维度上翻车了。