我最早接触图像分割时,第一反应是“把分类网络的全连接层换成卷积层,输出每个像素的类别不就行了?”这个思路没错,但效果始终不理想——边缘糊成一团,小目标直接消失。直到我把U-net网络结构完整地复现了一遍,看到编码器-解码器配合跳跃连接的设计,才真正理解了分割任务需要什么样的网络骨架。这篇文章就从零开始,用pytorch一步步搭建完整的U-net结构,让想入门深度学习图像分割的读者能一次性吃透它,而不是只抄个代码就跑。
1. 为什么这个结构能统治图像分割这么多年
1.1 图像分割和分类的本质差异
图像分类只需要回答“这张图是什么”,网络可以在一次次池化中丢掉位置信息,只要最后保留抽象的语义特征就够了。但语义分割要求每个像素都得到标签,所以网络必须同时在两件事上都做好:一是“知道这个东西是什么”,二是“知道它在哪儿”。这里就出现了一个天然的矛盾——想要更抽象的语义,就得不断下采样扩大感受野,而每次下采样都会丢掉空间细节;想要精细的边界,就得保留全分辨率特征,但那样感受野又不够,无法区分相似物体。很多初学者在分割任务上翻车,本质上都是没处理好这个矛盾。
1.2 U-net设计动机:全卷积网络做了什么
U-net的直接前身是FCN(全卷积网络)。FCN把分类网络的分类头替换成卷积和上采样,让网络可以接受任意尺寸输入并输出同尺寸的分割图。FCN提出了“跳跃结构”的雏形:把浅层特征和上采样后的深层特征相加,弥补丢失的细节。但FCN的上采样路径比较浅,设计也比较粗糙,浅层和深层的融合只有加和这一种方式,信息融合不充分。U-net把这条思路做到了极致,用一整条对称的扩展路径替代FCN里简单的上采样,形成了我们今天熟悉的U型结构。
1.3 对称的U型结构到底好在哪里
U-net的核心是一个对称的编码器-解码器结构,左半部分是收缩路径,负责编码语义信息;右半部分是扩展路径,负责恢复空间分辨率。每下采样一次,特征通道翻倍;每上采样一次,特征通道减半。这种对称设计让信息在两条路径中流向非常自然,参数数量也相对可控。真正的点睛之笔是连接左右两半的跳跃连接:扩展路径的每一层,都会把收缩路径同层输出的特征在通道维度上拼接过来。收缩路径的浅层特征保留了丰富的边界、纹理和位置信息,深层特征提供了类别语义,两者拼接后,解码器既知道目标是什么,又知道边界在哪里,这正是分割任务最需要的信息组合。
1.4 一个容易忽略的设计选择:为什么跳连要用拼接而不是相加
FCN里早期版本用的是元素级相加,U-net用了通道维度拼接。相加相当于对两组特征做了一个固定的线性混合,特征维度没有增加,模型只能被动地“挑选”现有信息;而拼接把两组特征的维度直接翻倍,后面的卷积层可以学习到更复杂的组合方式,浅层细节和深层语义之间的交互更充分,也保留了各自的独立性。付出的代价是通道数翻倍后的计算量增加,但换来的是分割精度的显著提升,这个取舍在绝大多数场景下都是值得的。
个人在实际使用中的体会是,当你减少U-net通道数(比如第一层从64改成32),拼接带来的增益会更明显——因为特征容量本身不足时,让网络自己去学习如何融合两组不同粒度的信息,比硬性相加灵活得多。
2. 动手前的准备:环境、数据与整体结构拆解
2.1 环境搭建:Win10下Anaconda配置PyTorch的要点
很多初学者不是被网络结构难住的,而是被环境配置劝退的。这里给一个稳定的配置流程:
- 安装Anaconda,创建独立环境:
conda create -n unet python=3.9 - 激活环境:
conda activate unet - CPU版本直接
conda install pytorch torchvision cpuonly -c pytorch - NVIDIA GPU用户先去官网看自己CUDA版本,然后
pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu118
关键的一点是不要用conda install pytorch这种不带 channel 的写法,它默认装的是 CPU 版,不少人在这一步卡了很久。验证安装成功只需要跑一句torch.cuda.is_available(),返回True说明GPU版本正常,False则说明装的还是CPU版。如果没有独显,用CPU跑小数据集做结构学习完全够用,不需要为了跑 demo 特意买 GPU 机器。
2.2 数据组织:标签图的格式是第一个大坑
U-net训练需要影像和对应的标签图。医学图像领域常用PNG格式的灰度标签图,其中每个像素的灰度值对应一个类别编号,比如0是背景,1是目标区域。这里容易踩坑的地方是:很多人把标签图保存成JPG,JPG是有损压缩,会在物体边缘产生过渡色,像素值不再是干净的0或1,直接影响了Loss的计算。
推荐的数据目录结构很简单:
data/ images/ # 原图,JPG或PNG都可以 masks/ # 标签图,一定是PNGDataset类的核心逻辑就是靠索引同时读入图片和对应名称相同的标签图。如果原图是灰度图,注意要转成三通道输入,或者改网络第一层的输入通道数为1。标签图在输入网络前不需要归一化,直接转成LongTensor供交叉熵损失使用即可。
2.3 U-net结构搭建在整个项目中的定位
我把一个完整的分割项目拆成四块:数据加载、网络结构、训练循环、评估策略。网络结构在整个pipeline中占的分量往往没有初学者想象中那么大,但它决定了模型的上限。数据加载决定了模型能不能学到有效特征,训练循环决定了模型能不能收敛,评估策略决定了你怎么判断模型好坏。U-net结构搭建是这四块里最“死”的部分——它不是靠调参能弥补的,也不是靠玄学可以绕开的,它是一个必须完全吃透、能够手写的基础零件。下面从像素级开始,把整个结构拆开来看。
3. 基于PyTorch搭建U-net的核心流程
3.1 最基础的双卷积模块
原版U-net中,每个路径段都由两个卷积组成(DoubleConv),这是整个网络最频繁使用的重复单元。为什么用两个卷积而不是一个?两个3x3卷积堆叠的感受野等于5x5,但参数量只有5x5卷积的72%左右,而且中间多了一次非线性激活,特征表达能力更强。这个设计可以追溯到VGG,U-net原封不动继承了过来。
import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x)这里有两个点需要说明。第一,原版U-net没有BatchNorm,但那个年代训练技巧不像现在这么成熟。实践下来加了BN训练稳定性明显更好、收敛更快,尤其是网络较深时效果很直观。第二,padding=1是我刻意加的,目的有两个:一个是让输出尺寸和输入保持一致,另一个是避免边缘信息被过早丢弃。原版U-net为了弥补两次valid卷积造成的尺寸缩小,采用镜像填充,现在用padding=1完全可以替代。
3.2 收缩路径:下采样与通道扩展
收缩路径做的事情很直白:通过最大池化把分辨率减半,同时用DoubleConv把通道数翻倍。
class Down(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.pool = nn.MaxPool2d(kernel_size=2, stride=2) self.conv = DoubleConv(in_ch, out_ch) def forward(self, x): x = self.pool(x) x = self.conv(x) return x关于池化的选择,U-net选的是MaxPool2d而不是AvgPool2d。原因在于,分割任务中的激活值稀疏性很强,边缘和关键点的响应往往体现在少数高响应的神经元上,最大池化可以保留这种最强的响应信号,平均池化反而会把它们稀释掉。当然也有研究者用stride=2的卷积替代池化来下采样,效果接近但计算量略大。对于基础结构的复现,用最大池化和原论文保持一致就够了。
3.3 扩展路径:上采样与跳跃连接的拼接
扩展路径是U-net结构中最需要仔细理解的部分,也是初学者最容易写错的地方。每一步先做上采样,把feature map放大一倍,然后取出收缩路径对应层的特征,在通道维度上拼接,最后过一个DoubleConv融合信息。
class Up(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.up = nn.ConvTranspose2d(in_ch, in_ch // 2, kernel_size=2, stride=2) self.conv = DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 = self.up(x1) # 处理可能存在的尺寸不一致 diffY = x2.size()[2] - x1.size()[2] diffX = x2.size()[3] - x1.size()[3] x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x = torch.cat([x2, x1], dim=1) return self.conv(x)这里有个细节容易让人困惑:为什么Up模块的输入通道数是in_ch,而输出通道数是out_ch?因为经过torch.cat之后,当前特征和跳跃特征拼接,通道数变成了上采样后通道数与跳跃特征通道数之和。如果上一层输出是in_ch // 2,跳跃特征也是in_ch // 2,拼接后就是in_ch,此时DoubleConv的输入必须写成in_ch。很多人的代码报错就出在这个通道数没转过来。
如果你用的是循环构建而非手写每一层的写法,要注意在拼接前处理尺寸不一致的问题。经典U-net在结构对称、输入尺寸可被16整除的情况下尺寸始终一致,但实际训练中如果用到了非对称的裁剪,或者输入尺寸不是2的整数倍,这里就会报错。上面的代码用F.pad做了对称填充,健壮性更好,实测下来能省去不少调试时间。
3.4 输出层与完整网络拼装
输出层用的还是1x1卷积,作用是把最后一层的通道数映射到类别数上。1x1卷积在这里相当于一个逐像素的全连接层,它把每个位置的特征向量压缩成类别得分,不具有跨像素的信息交互,这样做是合理的——跨像素的特征交互已经在前面各层完成了。
class OutConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv = nn.Conv2d(in_ch, out_ch, kernel_size=1) def forward(self, x): return self.conv(x)完整拼装起来,就是标准的U-net结构:
class UNet(nn.Module): def __init__(self, in_channels=3, num_classes=1, features=(64, 128, 256, 512)): super().__init__() self.downs = nn.ModuleList() self.ups = nn.ModuleList() self.pool = nn.MaxPool2d(2, 2) # 收缩路径 ch = in_channels for f in features: self.downs.append(DoubleConv(ch, f)) ch = f # 瓶颈层 self.bottleneck = DoubleConv(features[-1], features[-1] * 2) # 扩展路径 for f in reversed(features): self.ups.append(nn.ConvTranspose2d(f * 2, f, kernel_size=2, stride=2)) self.ups.append(DoubleConv(f * 2, f)) self.out_conv = OutConv(features[0], num_classes) def forward(self, x): skips = [] for down in self.downs: x = down(x) skips.append(x) x = self.pool(x) x = self.bottleneck(x) skips = skips[::-1] for idx in range(0, len(self.ups), 2): x = self.ups[idx](x) skip = skips[idx // 2] if x.shape != skip.shape: x = F.interpolate(x, size=skip.shape[2:], mode="bilinear", align_corners=True) x = torch.cat([skip, x], dim=1) x = self.ups[idx + 1](x) return self.out_conv(x)这段代码用features参数控制每层通道数,比手写四层Down和四层Up灵活得多。想复现原始论文的效果,只要把features设置成(64, 128, 256, 512)即可;显存有限的场景,可以改成(32, 64, 128, 256)。输入图像的尺寸只要是2的整数次幂倍数(比如512x512、256x256),forward过程就能顺利跑通。
3.5 上采样方式:转置卷积并非唯一选择
U-net原版用的是转置卷积(ConvTranspose2d),这也是我在上面代码中使用的方式。转置卷积是可学习的上采样,理论上网络可以自己学出一套最优的上采样核。但它有一个知名的副作用——当卷积核尺寸是偶数时,会产生棋盘效应,也就是特征图中出现规则分布的伪影。U-net用的kernel_size=2是偶数,所以这个现象在理论上存在,只是在实际分割任务中,由于跳跃连接带来的细节信息足够强,棋盘效应对最终结果的影响并不明显,很多复现项目并不专门处理它。
如果你在上采样后直接输出分割图,或者做生成类任务,棋盘效应就会变得肉眼可见。替代方案是用双线性插值上采样,再接普通卷积。插值是固定的,没有可学习参数,但不会引入伪影。实测在分割任务里两者的mIOU差距通常在0.5个百分点以内,不必过度纠结,关键是理解各自的取舍:转置卷积多了参数、可能引入伪影,插值方法更稳但表达力稍弱。我自己的习惯是基础复现用转置卷积,实际工程项目优先用插值方法,稳定性排在第一位。
4. 训练流程与损失函数怎么选
4.1 数据增强:小数据集的分割任务靠它续命
U-net最初是为医学图像设计的,这类数据集的典型特点是样本量小、标注成本高。几十张图训练一个分割网络是常态,这种情况下数据增强不是可选项,而是必需品。我在训练中常用的增强手段包括:
- 随机水平翻转和垂直翻转,概率各0.5
- 随机旋转90度、180度、270度
- 随机缩放(0.8到1.2倍)后再随机裁剪到固定尺寸
- 亮度对比度抖动,模拟不同采集条件下的图像差异
有一个细节必须注意:对图像做几何变换时,标签图必须做完全相同的变换。如果图像转了90度,标签没转,模型会学到一团混乱的特征。这里我的经验是使用同一个随机种子分别调用图像和标签的transform函数,或者直接用albumentations库,它可以对image和mask同步变换,能省掉几十行容易出bug的代码。
4.2 损失函数:交叉熵、Dice Loss和混合损失
分割任务最常用的损失函数是交叉熵。二分类用BCEWithLogitsLoss,多分类用CrossEntropyLoss。交叉熵对每个像素独立计算损失,梯度信号稳定,但它有一个很实际的问题——如果目标区域只占整张图的5%,网络只要把所有像素都预测成背景就能拿到95%的准确率,损失值也很低。对此加一个pos_weight参数能部分缓解,但作用有限。
分割领域更常用的指标是Dice系数,它衡量两个集合的重叠程度:
Dice = 2 * |A ∩ B| / (|A| + |B|)Dice Loss的定义是1 - Dice,它天然对前景和背景的数量差异不敏感,类别不平衡问题处理得比交叉熵好。但Dice Loss的梯度在预测极端值时可能不太稳定,单独使用时也容易陷入局部最优。实操中效果比较稳定的组合是混合损失:
Loss = BCEWithLogitsLoss + DiceLoss也就是把逐像素的交叉熵和整体区域的Dice Loss加在一起。这样既保留了逐像素监督的稳定梯度,又加入了区域级别的约束,前景过小的场景也能得到有效训练。实现Dice Loss时注意——它一定会返回一个标量,实现在纯pytorch中只需要几行代码:
def dice_loss(pred, target, smooth=1.0): pred = torch.sigmoid(pred) pred = pred.contiguous().view(-1) target = target.contiguous().view(-1) intersection = (pred * target).sum() return 1 - (2.0 * intersection + smooth) / (pred.sum() + target.sum() + smooth)pred在进入这个函数前是没有经过sigmoid的原始logits,所以需要先做激活,和BCEWithLogitsLoss的输入保持一致。多分类场景则用softmax激活,并对每个类别分别计算Dice后取平均。
4.3 优化器与学习率策略
U-net属于标准的卷积网络,Adam和SGD都是常见选择。我的经验是:Adam收敛快,前期效果提升明显,适合快速验证网络结构是否正确;SGD+动量收敛更稳,最终精度通常略高,适合调最终的模型。如果只是项目复现,从Adam开始最省心,learning rate设为1e-4基本不会出大问题。训练到一半如果验证集指标不再提升,可以改用余弦退火或ReduceLROnPlateau,让学习率逐步降下来,这一步常常能把验证集的Dice从0.85拉到0.87以上。
这里忍不住多说一句:不要迷信任何固定学习率参数。同一个学习率在32通道的轻量U-net上可能跑得好好的,换到原始64通道版本就可能发散。每次改网络结构,都先用一个很小的子集做一次10个epoch的快速测试,学习率先从1e-4开始,观察loss曲线有没有下降趋势,再决定要不要调。
4.4 标准训练循环模板
def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss = 0 for images, masks in loader: images = images.to(device) masks = masks.to(device) preds = model(images) loss = criterion(preds, masks) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader) def validate(model, loader, criterion, device): model.eval() total_loss = 0 with torch.no_grad(): for images, masks in loader: images = images.to(device) masks = masks.to(device) preds = model(images) loss = criterion(preds, masks) total_loss += loss.item() return total_loss / len(loader)注意训练阶段必须调用model.train(),验证阶段必须调用model.eval(),否则BatchNorm和Dropout在推理时仍会使用训练模式的行为,验证指标会严重失真。这个错误在我见过的新手代码里出现频率极高。
5. 验证模型与踩坑排查记录
5.1 评估指标:Dice和mIOU
训练过程中只能看到loss在下降,但loss下降并不代表分割效果好,尤其是类别不平衡时。我习惯每个epoch保存一次验证集的Dice系数和mIOU,作为模型好坏的判断依据。
mIOU的计算方式:对每个类别,计算预测值和真实值的交集除以并集,取所有类别的平均值。在pytorch中,可以先对logits取argmax得到预测类别,再逐类计算IOU:
def iou_score(pred, target, num_classes): ious = [] pred = pred.argmax(dim=1) for cls in range(num_classes): pred_cls = (pred == cls) target_cls = (target == cls) intersection = (pred_cls & target_cls).sum().float() union = (pred_cls | target_cls).sum().float() if union == 0: ious.append(float("nan")) else: ious.append((intersection / union).item()) return ious这里有一个实际操作的细节:如果某个类别在整张验证集里都没有出现,union为0,直接跳过它,否则算出来inf或nan会污染整体指标。
5.2 从踩坑到排错的完整链路
我在初次复现U-net时,遇到过一个持续了半天的报错:训练到第二个epoch,loss突然变成nan。排查过程如下。
第一反应是学习率太大,把学习率从1e-4降到1e-5,问题依旧。然后怀疑是数据问题,打印了输入的min和max,图片归一化到0到1之间,数值没有问题。接着怀疑标签问题,发现标签类型是torch.uint8,和CrossEntropyLoss要求的torch.long不匹配,转成long之后还是没有解决问题。最后打印了每一层输出的标准差,发现是瓶颈层的输出出现了巨大的激活值,一路追下去,定位到问题根源出在模型参数初始化上——我手工初始化了转置卷积层,初始化值过大,使特征经过多级传播后数值爆炸。去掉了那行不合理的初始化,网络恢复正常。
这个排查过程给我一个非常重要的经验:网络结构代码跑通只是第一步,loss不收敛或不稳定时,优先检查输入范围、标签类型、输出层是否带激活、初始化是否合理,不要一上来就动学习率。
5.3 几条高频异常情况的快速对照
| 现象 | 可能原因 | 排查方向 |
|---|---|---|
| loss一直不降 | 标签和图像不配对、增强没同步 | 可视化一个batch的输入输出 |
| 验证集指标为0 | 预测时忘记sigmoid/softmax | 检查模型输出的后处理 |
| 边界轮廓很粗 | 损失函数只用了逐像素损失 | 加上Dice Loss或边界损失 |
| 小目标全部丢失 | 输入尺寸太小或下采样丢失细节 | 增大输入尺寸、加深跳跃连接特征 |
| 训练和验证gap很大 | 数据量太少、增强不足 | 加增强、加dropout、减小模型容量 |
| 空洞和棋盘格 | 转置卷积产生的伪影 | 换插值上采样 |
5.4 关于BatchNorm的一个隐藏坑
BatchNorm在batch size较大时效果很好,但如果你的显存只允许batch size为2或者4,BN的均值和方差估计会非常抖,训练会变得不稳定。这也是我在轻量U-net复现时踩过的坑。解决办法有三个:一是用GroupNorm替代BatchNorm,它对batch size不敏感;二是用梯度累积模拟更大的batch;三是换成InstanceNorm。如果你只是做结构学习、跑通流程,把batch size设成8以上就暂时不用管这个问题,但真正在真实数据集上训练时,这个细节影响非常大。很多人复现论文效果不稳定,问题往往就出在BN的batch size上。
6. 在基础U-net上还能怎么改
6.1 注意力与密集连接的两个经典变体
U-net结构搭建完成后,如果你希望在具体任务上进一步提升精度,有两个被验证过很多次的改进方向。
第一个是Attention U-net,它在跳跃连接之前加了一个注意力门控,让网络自动学习哪些位置的特征需要被强调,哪些位置应该被抑制。核心动机是:并不是所有浅层特征都对当前像素的分割有用,有些位置是无关背景,强行拼接反而引入噪声。注意力模块在浅层和深层特征拼接前,学习一个注意力权重图(取值范围0到1),对浅层特征进行加权。这个改动只增加少量参数,对分割边界清晰度的提升很直接。
第二个是U-net++,它把原来单一的跳跃连接改成了密集的嵌套连接,每个解码器层不仅接收编码器同层的输出,还接收之前所有解码器层的中间结果,再用深监督让每一层的输出都参与损失计算。这样做的效果是网络收敛更快、精度更高,但显存占用也上涨得明显。如果你的任务是医学图像分割,U-net++通常比基础U-net高出1到2个点的Dice,这在很多场景下是质变。
6.2 轻量化思路:让U-net能跑在实时场景
U-net的标准版本有约3100万参数,单张512x512图片的推理需要的时间在CPU上很难接受。如果你的场景要求实时性,可以从三个方向压缩:把第一层通道数从64减到16或8,参数量会大幅降低;将普通卷积替换为深度可分离卷积,参数量和计算量都降到原来的十分之一左右;把上采样从转置卷积改为插值,同时减少解码器的通道数。上述操作组合起来,可以把模型压缩到原有体积的五分之一以下,精度损失通常在2到3个点以内,能够接受的话,这套轻量版U-net在边缘设备上非常实用。
6.3 关于这个结构还能走多远的个人看法
虽然Transformer类结构这两年占据了大量视野,但U-net在现代视觉任务中并没有过时。它的编码器-解码器加跳跃连接这套骨架,被灵活地继承到了Transformer分割模型中——很多结构只是把编码器和解码器中的卷积块换成了自注意力模块,整体依然保持U型架构。所以先把U-net结构吃透,再去看ViT、Swin Transformer这些变体,你会发现很多概念是相通的。我建议动手复现时不要只满足于跑通代码,试着回答三个问题:为什么下采样要四次而不是三次?为什么通道数要随深度翻倍?如果去掉跳跃连接,模型的输出会发生什么变化?把这几个问题在实验里逐一验证后,你对分割网络的理解就真正到位了。