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 的情况下,padding=1 才能保证输出的空间尺寸和输入一致。如果不加 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_size=3, padding=1), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x)这里我加了 BatchNorm,原始论文里没有,但实际训练时加上它收敛更稳,尤其是 batch size 比较小的时候。如果你要严格复现原论文,可以把 BN 去掉,但实测下来加上更好。
2.2 下采样到底用最大池化还是步长卷积
编码器的下采样,原始 U-Net 用的是 2x2 最大池化,stride 为 2。这个操作把空间尺寸精确地减半,而且不引入任何参数。另一种常见做法是用 stride=2 的卷积来代替池化,好处是下采样的过程也可学习,但会引入额外参数,而且尺寸计算稍微麻烦一点。
我两种都试过,在医学图像这种细节敏感的任务上,最大池化反而更稳,因为它保留了每个 2x2 窗口里的最大值,相当于保留了最显著的特征响应。步长卷积在某些数据集上能涨一点点,但不够稳定。所以下面的实现我用最大池化。
class Encoder(nn.Module): def __init__(self, in_channels, features): super().__init__() self.conv = DoubleConv(in_channels, features) self.pool = nn.MaxPool2d(kernel_size=2, stride=2) 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) # batch=1, channel=1, H=572, W=572 # 第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_size=2, stride=2),输出尺寸的计算公式是:
H_out = (H_in - 1) × stride - 2 × padding + kernel_size
当 kernel_size=2, stride=2, padding=0 时,公式简化为:
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_size=2, stride=2) 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], dim=1) print("concat:", concat.shape) # [1, 256, 286, 286]拼接后通道数变成 128+128=256,然后接一个 DoubleConv 把通道数降回 128,同时融合两部分信息。
class Decoder(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2) 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, size=skip.shape[2:], mode='bilinear', align_corners=True) x = torch.cat([skip, x], dim=1) return self.conv(x)这里我加了一个尺寸对齐的保险逻辑。虽然理论上对称结构尺寸应该完全匹配,但实际中如果输入尺寸不是 2 的幂次方,或者用了不同的 padding 策略,就可能差一两个像素。加上这个判断,能避免运行时直接报错。
注意:
torch.cat的dim=1是通道维度。如果你不小心写成dim=0,那就是在 batch 维度拼接,会直接把 batch size 翻倍,后面所有计算全乱。这个错误我见过不止一次,排查时先看 cat 的 dim 参数。
4. 完整网络组装与端到端维度追踪
4.1 把编码器、瓶颈、解码器串起来
现在把前面的模块组装成完整的 U-Net。我用一个 4 层下采样的版本,输入假设是 1 通道、256x256。
class UNet(nn.Module): def __init__(self, in_channels=1, num_classes=2, 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_size=1) 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_channels=1, num_classes=2) x = torch.randn(1, 1, 256, 256) skips = [] for i, encoder in enumerate(model.encoders): skip, x = encoder(x) skips.append(skip) print(f"编码器{i+1} 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"解码器{i+1}: {x.shape}") out = model.final_conv(x) print(f"最终输出: {out.shape}")运行结果会是这样:
| 阶段 | 操作 | 输出维度 |
|---|---|---|
| 编码器1 | DoubleConv + Pool | skip: [1,64,256,256], down: [1,64,128,128] |
| 编码器2 | DoubleConv + Pool | skip: [1,128,128,128], down: [1,128,64,64] |
| 编码器3 | DoubleConv + Pool | skip: [1,256,64,64], down: [1,256,32,32] |
| 编码器4 | DoubleConv + Pool | skip: [1,512,32,32], down: [1,512,16,16] |
| 瓶颈层 | DoubleConv | [1,1024,16,16] |
| 解码器1 | Up + Concat + Conv | [1,512,32,32] |
| 解码器2 | Up + Concat + Conv | [1,256,64,64] |
| 解码器3 | Up + Concat + Conv | [1,128,128,128] |
| 解码器4 | Up + 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, smooth=1e-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-16;512x512 的话,可能只能开到 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_groups=8, num_channels=out_channels)GroupNorm 把通道分成若干组,在每组内部做归一化,不依赖 batch 维度。实测在小 batch 场景下比 BatchNorm 稳很多。
6.3 转置卷积的棋盘格伪影
ConvTranspose2d有一个已知问题:当 kernel_size 不能被 stride 整除时,输出会出现棋盘格状的伪影。U-Net 里用的是 kernel_size=2、stride=2,正好整除,所以一般不会有这个问题。但如果你改成 kernel_size=3、stride=2,就要小心了。
替代方案是先做最近邻上采样,再用普通卷积。这样没有棋盘格问题,而且参数量更少。
class UpConvAlternative(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up = nn.Sequential( nn.Upsample(scale_factor=2, mode='nearest'), nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1) ) 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 这种维度变化规律性很强的网络,尤其适合用这种方式验证。你把上面那份端到端维度追踪的代码跑一遍,对照表格看一遍,基本就再也不会在维度上翻车了。