ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Pytorch下用Unet训练多类别语义分割数据集的完整指南

2026/9/9 6:09:55 拓冰建站 浏览量
Pytorch下用Unet训练多类别语义分割数据集的完整指南 简介基于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, transformNone, augmentNone): 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), dtypetorch.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(p0.5), A.VerticalFlip(p0.5), A.RandomRotate90(p0.5), A.ColorJitter(brightness0.2, contrast0.2, saturation0.2, p0.3), A.Normalize(meanIMG_MEAN, stdIMG_STD), ToTensorV2() ]) val_transform A.Compose([ A.Resize(256, 256), A.Normalize(meanIMG_MEAN, stdIMG_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(weightclass_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_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)BatchNorm2d放在卷积和ReLU之间可以加速收敛对多类别分割帮助明显。kernel_size设为3并padding为1保证特征图尺寸经过卷积后不变方便跳跃连接直接拼接。编码器部分逐层下采样每层特征通道翻倍class Encoder(nn.Module): def __init__(self, in_channels3, features[64, 128, 256, 512]): super().__init__() self.downs nn.ModuleList() self.pool nn.MaxPool2d(kernel_size2, stride2) 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_size2, stride2)) 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], dim1) x double_conv(x) return x上采样通道数的逻辑要特别说明转置卷积输入是上一步的特征图输出通道是feature但跳跃连接会把encoder那侧对应层的输出拼过来拼完之后通道数变为feature * 2所以紧接着的DoubleConv输入要写成feature * 2。这是Unet实现里最容易搞错的地方一旦写错维度不匹配的报错会直接告诉你有哪些维度对不上。最后把整个模型组装起来class UNet(nn.Module): def __init__(self, in_channels3, num_classes5): super().__init__() self.encoder Encoder(in_channels) self.bottleneck DoubleConv(512, 1024) self.decoder Decoder() self.final_conv nn.Conv2d(64, num_classes, kernel_size1) 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, eps1e-7, num_classes5): pred_softmax torch.softmax(pred, dim1) 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(dim1) 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(), lr1e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size20, gamma0.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_norm1.0) optimizer.step() train_loss loss.item() scheduler.step() avg_loss train_loss / len(train_loader) print(fEpoch {epoch1}/{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, }, fcheckpoint_epoch_{epoch1}.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(dim1)训练集正常验证集mIoU极低数据增强过度或者验证集归一化参数不一致验证时只做Resize和Normalize完全不做几何增强显存不足CUDA out of memorybatch size过大或图像分辨率过高减小batch size或降低输入尺寸必要时开启梯度累积4.2 显存优化与训练速度提升多类别分割任务里图像尺寸和batch size是显存的两大消耗点。如果显存不够第一个方案是把batch size降下来比如从8降到4第二个方案是减小输入分辨率比如从512降到256但要评估对精度的影响第三个方案是开启梯度累积模拟更大的batch sizeaccumulation_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_classes5, save_pathresult.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(dim1).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, cmaptab20) axes[1].set_title(Ground Truth) axes[2].imshow(pred, cmaptab20) 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部署到实际业务里。多类别分割的应用场景很广从遥感解译到医疗影像都离不开这套链路打好基础之后迁移起来会非常顺手。本文还有配套的精品资源点击获取