ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

多类别语义分割避坑指南:Unet数据准备与训练全流程解析

多类别语义分割避坑指南:Unet数据准备与训练全流程解析 简介一套基于PyTorch的Unet多类别语义分割实现代码面向需要构建自定义数据集并完成像素级分类的深度学习开发者可应用于医学影像、遥感图像等分割场景。压缩包共包含四十六个文件其中十九个为Python源码脚本、二十四个为pyc编译缓存文件、两个为文本说明、一个为JSON配置整体仅六十九KB轻量易部署。源码按功能划分为数据加载模块、网络结构模块、训练流程模块、损失与评估指标模块等另有演示脚本和模型保存工具便于快速测试与二次开发。该资源已有超过一万五千人次学习下载适合初学者与中级使用者可以帮助读者快速搭建语义分割实验环境。借助这份代码读者能理解Unet编码器与解码器的跳跃连接设计弄清多类别分割中输出通道数、交叉熵损失及IoU指标的计算方法文本说明和配置还能用于数据集划分与类别映射显著减少从零实现的时间。1. 从一张512×512的Mask开始说多类别语义分割的核心在数据侧两个月前我接手一个遥感地物分割任务六个类别每张影像512×512。当时以为把分割网络输出通道从1改成6就能跑结果第一周全耗在数据上——有人给的是RGB标签图有人给的是灰度索引图还有人把背景存成255而不是0。同一批数据光统一格式就返工了三次。这让我确认了件事用Unet做多类别语义分割网络结构本身是流水线上最省心的环节真正的坑几乎都埋在数据准备、类别定义和评估方式里。这篇笔记把从Mask清洗、DataLoader封装到模型训练和结果验证的完整链路写一遍代码可以直接抄改适合正准备拿Unet训练自己数据集的读者已经跑通过的人建议直接看第5章和第6章的排查思路。2. 数据准备是第一优先级RGB标签清洗、类别映射与Dataset封装多类别分割里最常见的翻车点不是网络是Mask格式。两类分割时你可以用0和1两个像素值甚至直接用布尔矩阵多类别场景下Mask的每个像素值必须是对应类别的索引而且整个数据集对同一个类别的编号必须完全一致。标注工具导出格式各不相同有的导出RGB调色板图有的导出单通道PNG还有的顺手把背景填了255。如果训练前没有做格式统一CrossEntropyLoss会把255当成一个真实类别去学训练直接失控。2.1 RGB标签转像素级class_id调色板映射脚本先分清两种主流格式灰度索引图每个像素直接存类别id看像素值就知道是哪一类RGB调色板图每个类别用一种颜色表示像素值本身是RGB三元组不能直接当类别id用。这两种格式不能混着训练需要在预处理阶段统一。下面这段脚本把RGB调色板图转成单通道灰度索引图import numpy as np from PIL import Image # 类别顺序固定下来中途不要改动否则之前训的权重全部作废 class_color_map { 0: (0, 0, 0), # 背景 1: (255, 0, 0), # 建筑 2: (0, 255, 0), # 植被 3: (0, 0, 255), # 水体 } def rgb_mask_to_class_id(rgb_path, out_path): rgb np.array(Image.open(rgb_path).convert(RGB)) h, w rgb.shape[:2] class_id np.zeros((h, w), dtypenp.uint8) for cls_id, color in class_color_map.items(): mask (rgb np.array(color)).all(axis-1) class_id[mask] cls_id Image.fromarray(class_id, modeL).save(out_path)all(axis-1)用来匹配三个通道完全相等的像素避免某个通道单独相等导致的误匹配dtypenp.uint8足够覆盖255个类别省内存。转换前务必统计一下原图里实际出现了哪些颜色如果有颜色没写进映射表对应像素会静默变成背景0这种错标在训练时肉眼很难发现。所以背景必须显式写进映射表不建议依赖默认置0行为。提示统一格式时顺带检查一遍像素值分布用np.unique(Image.open(mask_path))逐个文件看有哪些值只有0到C-1之间的id才合法。2.2 自定义Dataset同步变换、插值方式与tensor转换Unet训练和分类任务不一样的地方在于图像和Mask必须做完全相同的空间变换但数值处理逻辑完全不同。图像要归一化Mask里面的类别id不能做任何归一化否则类别值变成小数损失函数直接报废。还有一个细节是多类别分割的Mask在resize时只能用最近邻插值import torch from torch.utils.data import Dataset from PIL import Image import torchvision.transforms.functional as F class MultiClassSegDataset(Dataset): def __init__(self, img_paths, mask_paths, img_size(512, 512)): self.img_paths img_paths self.mask_paths mask_paths self.img_size img_size def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img Image.open(self.img_paths[idx]).convert(RGB) mask Image.open(self.mask_paths[idx]) img img.resize(self.img_size, Image.BILINEAR) mask mask.resize(self.img_size, Image.NEAREST) img_t F.to_tensor(img) mask_t torch.from_numpy(np.array(mask)).long() return img_t, mask_tMask用Image.NEAREST是硬性要求双线性插值会在类别边缘产生小数比如第0类和第1类交界处出现0.5这个值既不属于任何类别还会在后续计算损失时产生不可预测的梯度。F.to_tensor会把图像像素从0到255缩放到0到1之间Mask不能走这条通路必须保持原始整数类别id。mask_t转成long是因为PyTorch的CrossEntropyLoss要求target是长整型。2.3 验证集切分大图裁剪后的数据泄漏风险很多遥感或病理数据集是把一张大图切成若干patch来训练的。切分时如果直接对所有patch做随机划分同一个大图切出来的多个patch极可能同时出现在训练集和验证集mIoU会被严重高估换一张全新的图就暴露问题。正确做法是按大图文件名分组切分from sklearn.model_selection import GroupShuffleSplit # 每个patch的group id是它所属大图的文件名例如 scene_01_patch_3 的 group 是 scene_01 patch_names [...] # 所有patch的文件名 group_ids [name.split(_patch)[0] for name in patch_names] gss GroupShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(gss.split(patch_names, groupsgroup_ids))参数里test_size0.2按大图数量切不是按patch数量切这样验证集才能真实反映模型在未见过的场景上的表现。random_state42固定下来保证每次复现结果一致。切完之后务必做一次像素级类别分布统计确认验证集覆盖了全部类别尤其是稀有类别——如果一个罕见类别只出现在训练集里验证时它的IoU直接是0mIoU曲线会忽高忽低。3. Unet模型结构拆解编码器下采样、解码器上采样与跳跃连接3.1 为什么是Unet跳跃连接解决细节与语义的矛盾多类别语义分割的难点在于既要像素级的边缘细节又要有足够大的感受野理解目标语义。单纯加深网络会让浅层细节特征逐层丢失最后一层输出的特征图分辨率极低很难恢复精细边界。Unet通过编码器逐步池化压缩特征图获得语义信息再由解码器逐步恢复分辨率最关键的跳跃连接把编码器每一层的特征直接拼到解码器对应层。浅层特征保留了大量边缘和纹理信息深层特征提供了类别判断依据两者拼接后解码器可以兼顾两头的优势。网上很多Unet改进方案也是从这条链路入手把普通卷积替换成残差块、在跳跃连接处加注意力模块、把顶层换成ASPP例如U-Net重做了跳跃连接的密集结构Attention U-Net在每个跳跃连接上加了门控注意力。理解了这段结构改模型时就知道从哪里下手。3.2 DoubleConv与四层编码器解码器实现下面这个实现是大部分Unet代码的基础结构编码器四层、通道数逐层翻倍解码器先用转置卷积上采样再与编码器对应层拼接import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels3, num_classes4): super().__init__() self.enc1 DoubleConv(in_channels, 64) self.enc2 DoubleConv(64, 128) self.enc3 DoubleConv(128, 256) self.enc4 DoubleConv(256, 512) self.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(512, 1024) self.up4 nn.ConvTranspose2d(1024, 512, kernel_size2, stride2) self.dec4 DoubleConv(1024, 512) self.up3 nn.ConvTranspose2d(512, 256, kernel_size2, stride2) self.dec3 DoubleConv(512, 256) self.up2 nn.ConvTranspose2d(256, 128, kernel_size2, stride2) self.dec2 DoubleConv(256, 128) self.up1 nn.ConvTranspose2d(128, 64, kernel_size2, stride2) self.dec1 DoubleConv(128, 64) self.out_conv nn.Conv2d(64, num_classes, kernel_size1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) e4 self.enc4(self.pool(e3)) b self.bottleneck(self.pool(e4)) d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out_conv(d1)这里的细节值得细看。所有卷积都加了padding1在stride1时特征图尺寸不变所以每经过一次池化尺寸折半解码器拼接后尺寸恢复图能直接回到原输入大小。这个设计和Unet原论文用valid卷积不同原版每次卷积后尺寸会缩小现在主流实现基本都用same卷积省去计算裁剪坐标的麻烦。torch.cat这一步是跳跃连接的核心up4输出的通道数是512与e4的512拼起来变成1024所以dec4的in_ch必须写1024。解码器每个DoubleConv的输入通道都是“上采样输出通道数加编码器对应层通道数”改代码时最容易漏的就是这里。上采样用nn.ConvTranspose2d是常见方案kernel_size和stride都设为2正好把特征图放大两倍。也可以换成nn.Upsample(scale_factor2, modebilinear)加一个卷积参数更少但转置卷积能学习的参数更多效果略好但显存占用也更高。小数据集上两者差距不大。3.3 in_channels和num_classes改通道是最容易翻车的一步换到自己数据集上只需要改两个参数in_channels对应输入图像通道数RGB图像是3灰度图像是1num_classes对应你定义的总类别数包括背景。模型输出层会生成num_classes个特征通道每个通道对应一个类别的预测分数。这里有个常见误用有人把num_classes设成4但Mask里只有1、2、3三个类别id没有背景类。网络会输出4个通道第0通道对应的“背景”在训练数据里没有任何正样本这个通道基本学不出有效特征导致所有像素都被分到背景里mIoU惨不忍睹。要么补上背景样本要么把类别id重映射成从0开始连续编号。4. 训练到收敛损失函数组合、mIoU评估与超参清单4.1 交叉熵与Dice损失组合类别不均衡的两个解法多类别分割几乎都会遇到类别不均衡。遥感影像里背景或植被经常占绝大部分道路、建筑只占几个百分点。只用交叉熵时网络发现全部预测成背景也能把loss压得很低于是输出结果里目标类别要么完全缺失要么只在目标中心出现一小块。两个常用处理手段给交叉熵的每个类别加权或者组合Dice损失。Dice损失只看预测和真实区域的交叠比例天然对类别像素数量不敏感但单独用它训练前期梯度不稳定所以实践中常把两者组合起来import torch import torch.nn as nn def multiclass_dice_loss(pred, target, eps1e-6): # pred: [B, C, H, W] 未过softmax的logits # target: [B, H, W] 类别id C pred.shape[1] pred_softmax torch.softmax(pred, dim1) target_onehot torch.eye(C, devicepred.device)[target] target_onehot target_onehot.permute(0, 3, 1, 2) # [B, C, H, W] intersection (pred_softmax * target_onehot).sum(dim(2, 3)) union pred_softmax.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) dice (2 * intersection eps) / (union eps) return 1 - dice.mean() # 类别权重按像素占比的倒数归一化 def compute_class_weights(mask_paths, num_classes): pixel_counts np.zeros(num_classes, dtypenp.int64) for path in mask_paths: mask np.array(Image.open(path)) for c in range(num_classes): pixel_counts[c] (mask c).sum() total pixel_counts.sum() weights total / (pixel_counts 1e-6) weights weights / weights.sum() * num_classes return torch.tensor(weights, dtypetorch.float32) criterion_ce nn.CrossEntropyLoss(weightclass_weights) criterion_dice multiclass_dice_loss # 训练时: loss criterion_ce(pred, mask) criterion_dice(pred, mask)eps1e-6防分母为零Dice分数越接近1代表两类重合度越高所以损失取1 - dice。compute_class_weights里加1e-6是为了避免某个类别像素数为0时除零。系数方面我习惯先让两类损失各自单独跑几个step看数值量级再决定权重比例默认ce dice已经够用如果某个小目标类死活学不出来可以单独把Dice那项权重从1调到1.5。4.2 mIoU计算逐类IoU表比平均分更能暴露问题mIoU是语义分割的标准指标但只看平均分数会掩盖单类失效。例如四类分割中三个类IoU都在0.85背景IoU只有0.2mIoU平均下来“看起来还行”实际模型已经废了一半。所以评估时必须把每个类别的IoU单独算出来def compute_iou(pred, target, num_classes): # pred和target都是[H, W]的类别id ious [] for cls in range(num_classes): p (pred cls) t (target cls) inter (p t).sum().item() union (p | t).sum().item() if union 0: ious.append(inter / union) else: ious.append(float(nan)) return ious如果某个类别在验证集的像素数为0理论上IoU分母为0返回nan。出现这种情况不要简单忽略它而是要回到第2.3节重新排查验证集覆盖度。逐类IoU表建议打成三行表格看类别名、像素占比、IoU像素占比最低的那类往往就是IoU最低的类如果相反说明网络并没有被多数类带偏问题出在特征本身难分比如光谱相似的植被和农田。4.3 超参配置表学习率、Batch Size与学习率策略下面这组参数是我在四到六类分割任务上的起点配置大多数情况直接能用超参推荐值设置依据优化器AdamW相比Adam增加weight decay解耦BN层多的网络更稳定初始学习率1e-4Unet参数量大学习率调大容易在BN层产生震荡Batch Size4~8取决于显存8是512×512输入下的常见上限Epochs30~60看验证集mIoU是否连续10个epoch不再上升Weight Decay1e-5小数据集上防止过拟合学习率策略CosineAnnealingLR自动退火避免手动踩学习率悬崖关键点是学习率不要上来就按分类任务的1e-3走。Unet编码器预训练权重较少Decoder是从零开始训的1e-3配合小batch会让BatchNorm统计量剧烈抖动常见的后果是loss在几十个step内反复横跳。先用1e-4跑50个step观察loss曲线确认稳定下降后再决定要不要往上加。Batch Size小于4时BN层每个batch的统计量估计偏差很大如果显存受限考虑把输入分辨率降到320×320而不是强行维持512。5. Unet训练与使用时的常见问题排查覆盖环境、数据和推理5.1 CUDA与PyTorch版本不匹配第一个epoch卡死现象torch.cuda.is_available()返回True但训练到第一个batch就卡住或者直接报错CUDA error: no kernel image is available for execution on the device。原因PyTorch安装包内部自带的CUDA runtime版本高于显卡驱动支持的版本典型场景是机器显卡驱动较老却用pip默认装上了最新CUDA 12.x版本编译的PyTorch。GPU硬件和驱动能正常识别但核心计算kernel跑不起来。解决先查驱动支持的CUDA版本命令行执行nvidia-smi看右上角“CUDA Version”字段然后到PyTorch官网安装页选择对应的CUDA版本安装命令。装完后不要急着开训先跑一条自检命令python -c import torch; atorch.zeros(8).cuda(); print(a.sum().item())能输出0.0才说明GPU通路是通的。再训练一个很小的样本batch确认反向传播没问题。这个自检我每次装完环境都会跑一遍就几分钟能省下后面排查半天环境问题的时间。5.2 Mask类别id不连续loss训练中变NaN现象训练几分钟后loss变成NaN或者某个类别的预测概率一直接近0即便验证集里该类别的像素不少。原因Mask里存在0到C-1之外的像素值。不少标注工具把背景填成255或者用户自己画Mask时默认填充了某个过大的整数。CrossEntropyLoss遇到这些异常值时不会直接报错但梯度会变成NaN整个模型参数跟着崩掉。解决训练前对每个Mask执行np.unique()把所有出现的像素值打印出来和白名单[0, 1, ..., C-1]比对。这一步放在数据准备阶段不要拖到训练后。血泪经验是我曾经忽略了一批标注错乱的样本模型训到第10个epoch崩掉最后发现是某一张图里多了一个灰度值128的孤立像素点。5.3 loss下降但mIoU不动全背景预测与错标现象交叉熵loss稳定下降训练看起来一切正常但验证集mIoU始终卡在0.3左右。原因常见两类。一是类别不均衡背景占绝对主导网络把所有像素预测成背景交叉熵依然很低二是Mask本身有错标网络学到的是噪声边界没法泛化到验证集。解决先做全背景检测——把验证集所有预测结果统计一下像素直方图如果90%以上像素集中在第0类说明是类别不均衡问题回到4.1加Dice损失和类别权重。如果预测类别分布是正常的那就随机抽几组图像和Mask叠加可视化用半透明叠加检查边缘是否对齐。遥感数据里经常出现标注时把阴影区域标错类别的情况这种错标不会让loss产生明显异常但mIoU会一直上不去。5.4 推理输出是C个通道argmax加调色板映射现象模型训练完predict输出的shape是[1, C, H, W]直接保存成图片后全是黑的或花的。原因模型输出的C个通道是每个像素的类别分数不是最终类别标签。要取分数最高的通道作为该像素的预测类别再做调色板映射才能保存成可看的RGB图import torch import numpy as np from PIL import Image id_to_color { 0: (0, 0, 0), 1: (255, 0, 0), 2: (0, 255, 0), 3: (0, 0, 255), } def predict_one(model, img_tensor): model.eval() with torch.no_grad(): logits model(img_tensor.unsqueeze(0)) # [1, C, H, W] pred logits.argmax(dim1).squeeze(0).cpu().numpy() # [H, W] h, w pred.shape rgb np.zeros((h, w, 3), dtypenp.uint8) for idx, color in id_to_color.items(): rgb[pred idx] color return Image.fromarray(rgb)argmax(dim1)在通道维度上取最大值索引得到的pred每个像素就是0到C-1的类别id。保存之前还要确认推理时图像的resize方式与训练时完全一致否则预测结果和原图位置对不上。6. 进阶验证用混淆矩阵和预测叠加图定位分割误差模型训练完不是结束真正的调试从验证集分析开始。我每次训练完会先对验证集生成一张混淆矩阵再叠加预测可视化两步配合能快速定位模型失败在哪。混淆矩阵这里按行归一化每行代表一个真实类别被预测成各类别的比例from sklearn.metrics import confusion_matrix import matplotlib.pyplot as plt import seaborn as sns def plot_confusion_matrix(all_target, all_pred, class_names): cm confusion_matrix(all_target, all_pred) cm_norm cm.astype(float) / cm.sum(axis1, keepdimsTrue) fig, ax plt.subplots(figsize(8, 6)) sns.heatmap(cm_norm, annotTrue, fmt.2f, xticklabelsclass_names, yticklabelsclass_names, axax) ax.set_xlabel(Predicted) ax.set_ylabel(True) plt.savefig(confusion_matrix.png, dpi150)看这张图时重点看两个位置对角线之外数值最大的格子——比如第2行第4列是0.35表示真实类别2有35%被预测成了类别4说明这两类特征重叠严重再看是否有一整列数值都偏高说明网络整体偏向预测成某个类别。遥感场景里最典型的就是植被和农田互相污染以及阴影区域被分到水体。混淆矩阵能告诉你哪两类在混淆但看不出混淆发生在图像的哪个区域这时叠加可视化就派上用场了。更进一步用softmax输出的最大值作为每个像素的置信度低于阈值的区域在原图上直接标黑能直观看到模型“不知道”的区域prob_map torch.softmax(logits, dim1).max(dim1)[0] # [H, W] 每个像素的最高类概率 low_conf prob_map 0.5低置信度区域通常集中在类别交界处和图像边缘如果低置信度区域整片出现往往是训练数据里该区域本身标注不一致。从那以后我每次训练完都强制走一遍这套流程先出整图预测叠加再算逐类IoU表最后看混淆矩阵。这套动作帮我避开过三次拿着高mIoU模型却发现目标类别全错的尴尬——mIoU高只是平均值高某个关键类别彻底失效时它照样能及格。希望这些经验能帮你少踩几个坑尤其是数据准备阶段值得多花时间做清洗和校验。本文还有配套的精品资源点击获取
返回列表