news 2026/9/19 8:39:07

从零实现U-Net:PyTorch逐层维度追踪与跳跃连接详解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零实现U-Net:PyTorch逐层维度追踪与跳跃连接详解

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.catdim=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}")

运行结果会是这样:

阶段操作输出维度
编码器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时尺寸对不上。我的排查顺序是这样的:

  1. 先看报错信息里的两个 shape,确认是空间尺寸不一致还是通道数不一致。
  2. 如果是空间尺寸差 1-2 个像素,大概率是输入尺寸不能被 2 整除多次,或者某层 padding 设置不对。解决办法是把输入 resize 到 2 的幂次方,或者在 cat 前用 interpolate 对齐。
  3. 如果是通道数不一致,检查 ConvTranspose2d 的 out_channels 是否和对应 skip 的通道数相等。解码器的 up 输出通道,必须等于同层 skip 的通道数,否则 cat 后通道数不对,下一层 DoubleConv 的 in_channels 也会错。
  4. 如果是 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 - dice

Dice 系数的直观理解是:预测区域和真实区域的重叠程度。完全重叠是 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 这种维度变化规律性很强的网络,尤其适合用这种方式验证。你把上面那份端到端维度追踪的代码跑一遍,对照表格看一遍,基本就再也不会在维度上翻车了。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 8:37:44

AI文献综述系统:自然语言处理与知识图谱的学术革命

1. 项目背景与核心价值文献综述是学术研究的基石工程,却也是最耗时的环节之一。根据Nature最新调研,科研人员平均花费37%的工作时间在文献检索与整理上。传统工作流程存在三个痛点:信息过载导致关键文献漏读、人工归纳效率低下、文献间关联性…

作者头像 李华
网站建设 2026/9/19 8:34:54

QQ邮箱授权码全解析:获取、SMTP代发邮件与避坑指南

做网站或者小程序开发的人,迟早会遇到一个需求:让系统自动给用户发邮件。注册激活、验证码登录、密码找回、风控通知,后台总得有个能自动发信的通道。我见过很多新手第一反应是拿自己的QQ邮箱去发,结果卡在“授权码”这一步——明…

作者头像 李华
网站建设 2026/9/19 8:34:50

React状态管理与性能优化实战指南

1. React 组件状态机原理与性能优化实战作为一名长期奋战在一线的前端开发者,我见过太多因为状态管理不当导致的性能灾难。React 将组件视为状态机(State Machine)的设计理念,本质上是对 UI 开发范式的革命性改进。今天我想分享的…

作者头像 李华
网站建设 2026/9/19 8:33:27

open-code-review:可编程的开源代码评审范式

1. “open-code-review”不是新工具,而是一套可落地的开源代码评审范式你可能在 GitHub Trending 或某次技术分享里见过这个词——open-code-review,它不像eslint那样有明确的 npm 包,也不像prettier那样带.prettierrc配置文件。它没有官网、…

作者头像 李华