简介:基于PyTorch实现Unet多类别语义分割的工程源码包,面向需要训练自有数据集的开发者与学生,可直接作为项目的代码基底。压缩包共46个文件,以19个Python脚本为核心,覆盖数据加载、模型搭建、损失计算、训练评估与可视化流程,并对应dataloaders、modeling、utils等模块;另含24个pyc编译文件和txt、json配置,便于直接运行与调试,整体仅69KB,轻量易用。资源已有15245人学习,入口脚本train.py和demo.py能帮助快速跑通完整流程,稍作修改即可迁移到医学影像、遥感图像等多类别分割场景,省去从零搭建Unet的重复劳动。 做过分割任务的朋友应该都有体会:二分类分割,比如只分前景背景、只分建筑物和非建筑物,跑通容易;一旦切到多类别,各种问题就跟着冒出来——类别不均衡、标签编码混乱、mIoU上不去、显存不够用。我最早接触Unet是从医学影像开始,当时数据集是二值掩膜,后面转到遥感影像做多类别地物分类,把Unet从二分类一路改成多类别,踩了不少坑,也总结了一套相对稳定的流程。这篇文章就围绕“Pytorch下用Unet训练自己的多类别分割数据集”这个主题,把从数据集整理、模型搭建到训练调优的完整链路拆开讲一遍,内容同时覆盖Unet结构理解、多类别标签处理、损失函数选择、评估指标计算和常见坑位排查,适合准备用Unet入门语义分割、或者已经跑通二分类想升级到多类别的同学。
1. 项目整体思路:为什么Unet适合多类别语义分割
1.1 语义分割任务到底是什么
语义分割的本质是像素级分类。普通图像分类给整张图一个标签,目标检测给物体画一个框,而语义分割要求给每一个像素都预测一个类别,输出结果和原图同尺寸。多类别分割就是类别数大于等于2,通常还会包含一个背景类,比如遥感影像里分水体、植被、建筑、道路,再加一个背景或者其他地物,总共五个类别。
与二分类分割相比,多类别分割的第一个变化就是输出层。二分类常用Sigmoid加BCELoss,多类别必须换成Softmax加CrossEntropyLoss,输出通道数等于类别总数。第二个变化是标签形式,二分类掩膜是单通道0和1,多类别掩膜通常是单通道0到类别总数减1的整数编码,不能直接用One-Hot存,否则文件体积会膨胀好几倍。第三个变化是评估逻辑,二分类看一眼准确率就行,多类别需要逐类别算IoU再取平均,也就是mIoU,这样才能公平反映每个类别的表现。
1.2 Unet结构为什么是中小数据集的优选
Unet的经典结构是编码器-解码器加跳跃连接。编码器通过卷积加下采样逐层提取语义特征,感受野越来越大,能回答“这是什么”;解码器通过上采样逐步恢复空间分辨率,能回答“这个位置在哪”。跳跃连接把编码器同尺度的特征图拼到解码器上,相当于把高分辨率的边缘、纹理细节直接传给解码器,弥补了下采样带来的空间信息损失,这一点对小目标分割特别关键。
Unet在中小规模数据集上表现好,核心原因有两个。第一,参数量相对可控,经典Unet基础版也就三千万左右参数,比Transformer系列动辄上亿轻量很多,单卡能训。第二,跳跃连接本质上是一种隐式的数据增强,它让模型在有限样本下也能学到较强的局部一致性。如果你换Swin-UNet这类Transformer结构,在几千张图的数据集上反而容易欠拟合或者过拟合,调参成本直线上升。所以自己整理的数据集,规模一般有限,Unet是非常务实的起点。
1.3 整体技术路线与选型逻辑
从零开始做这个项目,我的整体路线是:
- 准备数据集,统一为图片加单通道标签掩膜的文件组织方式
- 实现数据加载器,完成归一化、尺寸统一、数据增强
- 搭建Unet模型,重点写好编码器、解码器和跳跃连接
- 选择损失函数和优化器,配置训练循环
- 训练过程中记录损失和mIoU,结束后可视化预测结果
为什么把数据集放第一位?因为分割任务的性能上限很大程度由标签质量决定,模型结构反而是相对成熟的部分。很多同学一上来就搭模型,跑通之后发现loss降不下去,排查半天才发现标签类别编码有错,或者数据加载时标签被插值破坏了,等于白白折腾。经验之谈,数据集阶段多花点时间,后面训练会顺利很多。
2. 多类别数据集准备与预处理
2.1 数据集目录组织与标签格式约定
我最常用的组织方式是VOC风格,目录结构清晰,也方便后续扩展:
data/ ├── images/ │ ├── img_001.jpg │ ├── img_002.jpg │ └── ... ├── masks/ │ ├── img_001.png │ ├── img_002.png │ └── ... └── class_dict.jsonimages目录放原始图像,masks目录放标签掩膜。掩膜必须是PNG格式,因为PNG是无损压缩,能保证像素值不被改动;JPG是有损格式,标签边缘很容易出现压缩伪影,导致训练时类别非常混乱。class_dict.json用来记录类别名和像素值映射,比如:
{ "0": "background", "1": "vegetation", "2": "building", "3": "road", "4": "water" }一个容易忽略的细节:标签掩膜里的类别值必须是连续整数,从0到类别数减1。如果你在标注软件里用了255表示某一类,或者类别跳号了,Softmax的输出通道数和标签对不上,训练必炸。整理数据后第一件事就是把掩膜的像素分布统计一遍,用Python打印所有唯一值,确认没有异物。
2.2 自定义Dataset实现与关键处理点
Pytorch里写一个适合多类别分割的Dataset类,核心要处理的是图像和掩膜同步加载,以及训练时对掩膜做与图像一致的数据增强。看这个精简实现:
import torch from torch.utils.data import Dataset from PIL import Image import os import numpy as np IMG_MEAN = [0.485, 0.456, 0.406] IMG_STD = [0.229, 0.224, 0.225] class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, transform=None, augment=None): self.image_paths = sorted([os.path.join(image_dir, f) for f in os.listdir(image_dir)]) self.mask_paths = sorted([os.path.join(mask_dir, f) for f in os.listdir(mask_dir)]) self.transform = transform self.augment = augment def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image = Image.open(self.image_paths[idx]).convert("RGB") mask = Image.open(self.mask_paths[idx]) if self.augment: image, mask = self.augment(image, mask) mask = torch.as_tensor(np.array(mask), dtype=torch.long) image = self.transform(image) return image, mask这里有几个关键点值得展开。第一,掩膜读进来之后要转成torch.long类型,这是CrossEntropyLoss的硬性要求,它期望的标签是长整型索引,而不是浮点数。第二,convert("RGB")保证图像是统一的三通道,否则灰度图和彩色图混着读进来,通道数不一致,第一个卷积层就报错。第三,mask = Image.open(...)不要做convert("RGB"),否则三通道的掩膜会让后续计算交叉熵时维度混乱,一定要保持单通道。
2.3 数据增强时最容易犯的错误
训练分割模型时数据增强必不可少,但很多人直接在图像上用torchvision的随机翻转、随机旋转,标签掩膜没有同步处理,结果就是增强后的图跟标签对不上,模型越训越差。我这里推荐用albumentations库,它的设计初衷就是图像和掩膜同步增强,非常省心:
import albumentations as A from albumentations.pytorch import ToTensorV2 train_transform = A.Compose([ A.Resize(256, 256), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, p=0.3), A.Normalize(mean=IMG_MEAN, std=IMG_STD), ToTensorV2() ]) val_transform = A.Compose([ A.Resize(256, 256), A.Normalize(mean=IMG_MEAN, std=IMG_STD), ToTensorV2() ])A.Resize用双线性插值处理图像,但标签掩膜会自动切到最近邻插值,不会生成类别之间不存在的像素值,这一点非常关键。如果你自己手写增强逻辑,掩膜缩放一定要用Image.Resampling.NEAREST,否则边缘区域会出现类似1.6、2.3这样的中间值,损失函数计算时直接越界,或者把类别搞脏。
2.4 类别不平衡与权重设置
多类别数据集几乎必然存在类别不平衡,比如遥感影像里水体占比可能不到5%,背景占了大半。如果不做任何处理,模型会倾向把所有像素预测为多数类,因为这样也能把loss压得很低,但实际分割效果惨不忍睹。
应对方法有两个。一个是给CrossEntropyLoss传入类别权重,让少数类的梯度贡献变大:
class_weights = torch.tensor([0.3, 1.0, 1.5, 1.2, 2.5]) criterion = torch.nn.CrossEntropyLoss(weight=class_weights.to(device))权重的设定可以按类别像素占比的倒数做归一化,也可以根据实际效果手动调。常见做法是先统计每个类别在训练集里的像素数total_pixels,然后计算weight_i = median / total_pixels_i,这样多数类权重小于1,少数类权重大于1。另一个方法是使用Dice Loss和CrossEntropyLoss的加权组合,Dice Loss对前景小目标更友好,后面损失函数章节再细说。
3. 核心代码实现:模型搭建、损失函数与训练循环
3.1 Unet模型各模块详解
自己搭建Unet时,我习惯把模型拆成卷积块、下采样模块、上采样模块和跳跃连接四个部分分开写,便于调试和修改。卷积块就是两个卷积加ReLU激活:
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)BatchNorm2d放在卷积和ReLU之间,可以加速收敛,对多类别分割帮助明显。kernel_size设为3并padding为1,保证特征图尺寸经过卷积后不变,方便跳跃连接直接拼接。
编码器部分逐层下采样,每层特征通道翻倍:
class Encoder(nn.Module): def __init__(self, in_channels=3, features=[64, 128, 256, 512]): super().__init__() self.downs = nn.ModuleList() self.pool = nn.MaxPool2d(kernel_size=2, stride=2) for feature in features: self.downs.append(DoubleConv(in_channels, feature)) in_channels = feature def forward(self, x): skip_features = [] for down in self.downs: x = down(x) skip_features.append(x) x = self.pool(x) return x, skip_features这里features列表定义每一层的通道数,经典Unet常用[64, 128, 256, 512],输入3通道输出512通道特征图。skip_features存储每一层下采样前的特征图,之后要拼给解码器。
解码器部分用转置卷积做上采样,通道数逐层减半:
class Decoder(nn.Module): def __init__(self, features=[512, 256, 128, 64]): super().__init__() self.up_convs = nn.ModuleList() self.double_convs = nn.ModuleList() for feature in features: self.up_convs.append(nn.ConvTranspose2d(feature * 2, feature, kernel_size=2, stride=2)) self.double_convs.append(DoubleConv(feature * 2, feature)) def forward(self, x, skip_features): for i, (up_conv, double_conv) in enumerate(zip(self.up_convs, self.double_convs)): x = up_conv(x) skip = skip_features[-i - 1] x = torch.cat([x, skip], dim=1) x = double_conv(x) return x上采样通道数的逻辑要特别说明:转置卷积输入是上一步的特征图,输出通道是feature,但跳跃连接会把encoder那侧对应层的输出拼过来,拼完之后通道数变为feature * 2,所以紧接着的DoubleConv输入要写成feature * 2。这是Unet实现里最容易搞错的地方,一旦写错,维度不匹配的报错会直接告诉你有哪些维度对不上。
最后把整个模型组装起来:
class UNet(nn.Module): def __init__(self, in_channels=3, num_classes=5): super().__init__() self.encoder = Encoder(in_channels) self.bottleneck = DoubleConv(512, 1024) self.decoder = Decoder() self.final_conv = nn.Conv2d(64, num_classes, kernel_size=1) def forward(self, x): x, skip_features = self.encoder(x) x = self.bottleneck(x) x = self.decoder(x, skip_features) x = self.final_conv(x) return x最终输出层的1x1卷积把64维特征映射到num_classes维,得到每个像素在所有类别上的得分。这里没有接Softmax,因为Pytorch的CrossEntropyLoss内部已经包含了Softmax计算,网络输出logits就好。
3.2 多类别分割的损失函数选择
多类别分割最常用的损失函数是CrossEntropyLoss,它对每个像素独立计算交叉熵再取平均:
criterion = nn.CrossEntropyLoss()之所以不手动Softmax加NLLLoss再组合,是因为Pytorch的CrossEntropyLoss在数值稳定性上做了优化,直接用logits计算,能避免Softmax之后概率接近0导致log出现负无穷的情况。如果你使用的是带类别权重的CrossEntropyLoss,还能同时缓解类别不平衡问题。
当类别不平衡严重时,我偏向使用CrossEntropyLoss和DiceLoss的混合损失,比如总损失等于0.5倍交叉熵加0.5倍DiceLoss。DiceLoss的原始形式是针对二分类的,多类别场景一般这样扩展:对每个类别先算损失,再对所有类别取平均。
def multiclass_dice_loss(pred, target, eps=1e-7, num_classes=5): pred_softmax = torch.softmax(pred, dim=1) loss = 0.0 for c in range(num_classes): pred_c = pred_softmax[:, c] target_c = (target == c).float() intersection = (pred_c * target_c).sum() union = pred_c.sum() + target_c.sum() loss += 1 - (2.0 * intersection + eps) / (union + eps) return loss / num_classesDiceLoss的优势在于直接优化区域重叠度,对小目标的类别更敏感,但缺点是训练初期容易震荡,所以最好跟CrossEntropyLoss搭配使用,交叉熵负责稳定收敛方向,Dice负责细抠边界和少样本类别。
3.3 评估指标:mIoU的计算逻辑与实现
多类别语义分割的通用评估指标是mIoU,也就是对每个类别分别计算IoU,再对所有类别取平均。单类别的IoU等于该类别的预测结果和真实标签的交集除以并集,你可以把它理解为两个区域的重叠程度:
def compute_miou(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()) valid_ious = [iou for iou in ious if not math.isnan(iou)] return sum(valid_ious) / len(valid_ious)注意一个细节:当某个类别在整张图上都没出现时,union等于0,直接把这个类别的IoU记为nan然后剔除。这是因为空类别的IoU分母为0,无法计算,强行跳过比算成0更公平。还有一种做法是把这个空类别的IoU记成1,因为“没有目标”可以被认为预测完全正确,但这样会让mIoU偏高,不同论文实现不一致,看自己的评估需求选择。
3.4 训练循环与模型保存的最佳实践
训练循环的写法比较固定,但有几个细节对多类别分割很重要。一个是模型要调用model.train(),另一个是每一步要把优化器梯度清零,检测到梯度爆炸时要做梯度裁剪:
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=20, gamma=0.5) for epoch in range(num_epochs): model.train() train_loss = 0.0 for images, masks in train_loader: images, masks = images.to(device), masks.to(device) outputs = model(images) loss = criterion(outputs, masks) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() train_loss += loss.item() scheduler.step() avg_loss = train_loss / len(train_loader) print(f"Epoch {epoch+1}/{num_epochs}, Loss: {avg_loss:.4f}")Adam的初始学习率建议设1e-4,比2e-4更稳,尤其是BatchNorm层较多、数据规模又不大时,学习率稍大就非常容易震荡。StepLR每20个epoch把学习率降一半,可以保证后期训练更精细。
模型保存有个容易被忽略的坑:不要只保存state_dict,最好把优化器状态和当前轮次一起保存,这样如果训练中断可以恢复:
torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'loss': avg_loss, }, f"checkpoint_epoch_{epoch+1}.pth")有些同学为了方便部署,只保存state_dict,这没问题,但训练过程中建议用上面的方式保存完整的checkpoint。实际跑长训练时,训练中断只能用上一次的保存结果继续,少跑几十个epoch的差距在分割任务上非常明显。
4. 训练排错与调优实战
4.1 多类别分割高频报错排查
我把自己和身边朋友在Unet多类别分割中踩过的常见错误整理成了一个排查表,遇到问题可以对号入座。
| 现象 | 可能原因 | 解决方法 |
|---|---|---|
| 训练时报错“Expected target size [N, C, H, W], got [N, H, W]” | 掩膜被读成了多通道,或者转成了one-hot形式 | 掩膜保持单通道,类型为torch.long |
| 报错“IndexError: Target 255 is out of bounds” | 标签里存在超过类别数的像素值 | 统计掩膜唯一值并修正标签,确认类别从0连续编码 |
| Loss一直不下降,预测结果全是同一类 | 类别极端不平衡,或者学习率过大 | 给损失函数加类别权重,降低学习率 |
| 验证集mIoU高但可视化效果差 | 评估时没有对预测的logits取argmax | 确保可视化时先做pred = outputs.argmax(dim=1) |
| 训练集正常,验证集mIoU极低 | 数据增强过度,或者验证集归一化参数不一致 | 验证时只做Resize和Normalize,完全不做几何增强 |
| 显存不足(CUDA out of memory) | batch size过大或图像分辨率过高 | 减小batch size或降低输入尺寸,必要时开启梯度累积 |
4.2 显存优化与训练速度提升
多类别分割任务里,图像尺寸和batch size是显存的两大消耗点。如果显存不够,第一个方案是把batch size降下来,比如从8降到4;第二个方案是减小输入分辨率,比如从512降到256,但要评估对精度的影响;第三个方案是开启梯度累积,模拟更大的batch size:
accumulation_steps = 4 for i, (images, masks) in enumerate(train_loader): outputs = model(images) loss = criterion(outputs, masks) loss = loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()归一化放在数据加载阶段而不是模型内部,可以节省一部分计算资源。另外,如果训练集图片尺寸不统一,不要用Resize强行拉伸成正方形,这样会破坏物体的长宽比,分割边界会变得奇怪。我自己偏向使用A.Resize(256, 256)统一尺寸,简单省事,但如果你的数据集长宽比差异很大,可以考虑A.PadIfNeeded加A.RandomCrop的组合,在正方形滑窗里训练。
4.3 训练过程观察与可视化验证
训练时除了看loss曲线,还要定期做预测结果可视化,确保模型真的在学东西。我把打印逻辑分成两个层次:每轮结束后打印平均loss和验证集mIoU,每5个epoch保存一批原图、真值、预测结果的三联图。
可视化预测的核心代码:
import matplotlib.pyplot as plt def visualize_prediction(model, val_loader, device, num_classes=5, save_path="result.png"): model.eval() images, masks = next(iter(val_loader)) images, masks = images.to(device), masks.to(device) with torch.no_grad(): outputs = model(images) preds = outputs.argmax(dim=1).cpu().numpy() image = images[0].cpu().numpy().transpose(1, 2, 0) mask = masks[0].cpu().numpy() pred = preds[0] fig, axes = plt.subplots(1, 3, figsize=(12, 4)) axes[0].imshow(image) axes[0].set_title("Image") axes[1].imshow(mask, cmap="tab20") axes[1].set_title("Ground Truth") axes[2].imshow(pred, cmap="tab20") axes[2].set_title("Prediction") plt.savefig(save_path)这里tab20是matplotlib里自带的多类别colormap,支持20种离散颜色,足够展示大多数分割任务的类别。如果你的类别少于20,也可以自定义一个colormap列表,比如水用蓝色、植被用绿色、建筑用红色,这样看起来更直观,调色板固定后也方便跟论文里的渲染图对齐。
4.4 提升分割精度的微调技巧
模型跑通之后,如果想进一步提升精度,我按性价比从高到低排序给几个建议。
第一,换更强的编码器。把Unet的encoder从自定义的双卷积块换成ResNet34或ResNet50的预训练权重,加载在ImageNet上训好的参数做迁移学习,小数据集上效果提升通常非常明显。Pytorch里可以用torchvision.models.resnet34(weights=...)取出特征层替换encoder,解码器保持原来的Unet解码部分。
第二,调整损失函数权重。如果发现小类别漏检,可以适当调大DiceLoss的比重,比如从0.5调到0.7。如果发现边缘粗糙,增加CrossEntropyLoss的比重,因为它逐像素约束,对边缘更敏感。
第三,测试时增强(TTA)。推理时把输入图水平翻转、垂直翻转,分别预测后再对logits取平均,类别分布会更平滑,mIoU一般能涨0.5到1个百分点,代价是推理时间变成原来的好几倍。
第四,增加数据规模,包括同分布数据的采集和强数据增强的扩展。分割模型本质上还是数据驱动,数据量上去了,模型能力才有发挥空间。
5. 从二分类切到多类别,我的实操心得
多类别分割和最常见的项目经验总结起来就一句话:把所有跟“类别”相关的环节都检查一遍,从标签编码到损失函数,从评估逻辑到可视化。我自己踩过最深的一个坑是二分类时用Sigmoid,切到五分类之后忘了改输出层和损失函数,结果训练了20轮都没收敛,最后打印预测发现所有类别得分接近一致,才反应过来是输出通道写错了。
还有一点值得单独说:训练过程中不要只看loss,一定要定期人工查看可视化结果。loss曲线光滑只能说明优化过程正常,但分割效果好不好、边缘是不是锯齿状、小目标有没有漏,肉眼看得最快。我通常每5轮保存一次预测可视化,盯着边界质量和类别完整性做判断。
如果你准备在这个项目上继续深入,可以尝试的方向有几个:把Unet编码器替换成预训练ResNet做迁移学习、引入注意力模块优化边界分割、用变体如Attention UNet或UNet++对比精度,以及把模型导出为TorchScript或ONNX部署到实际业务里。多类别分割的应用场景很广,从遥感解译到医疗影像都离不开这套链路,打好基础之后迁移起来会非常顺手。
本文还有配套的精品资源,点击获取