简介:一套基于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), dtype=np.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, mode="L").save(out_path)all(axis=-1)用来匹配三个通道完全相等的像素,避免某个通道单独相等导致的误匹配;dtype=np.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_splits=1, test_size=0.2, random_state=42) train_idx, val_idx = next(gss.split(patch_names, groups=group_ids))参数里test_size=0.2按大图数量切,不是按patch数量切,这样验证集才能真实反映模型在未见过的场景上的表现。random_state=42固定下来,保证每次复现结果一致。切完之后务必做一次像素级类别分布统计,确认验证集覆盖了全部类别,尤其是稀有类别——如果一个罕见类别只出现在训练集里,验证时它的IoU直接是0,mIoU曲线会忽高忽低。
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, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels=3, num_classes=4): 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_size=2, stride=2) self.dec4 = DoubleConv(1024, 512) self.up3 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2) self.dec3 = DoubleConv(512, 256) self.up2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2) self.dec2 = DoubleConv(256, 128) self.up1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2) self.dec1 = DoubleConv(128, 64) self.out_conv = nn.Conv2d(64, num_classes, kernel_size=1) 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], dim=1)) d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1)) d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1)) return self.out_conv(d1)这里的细节值得细看。所有卷积都加了padding=1,在stride=1时特征图尺寸不变,所以每经过一次池化,尺寸折半;解码器拼接后尺寸恢复,图能直接回到原输入大小。这个设计和Unet原论文用valid卷积不同,原版每次卷积后尺寸会缩小,现在主流实现基本都用same卷积,省去计算裁剪坐标的麻烦。torch.cat这一步是跳跃连接的核心:up4输出的通道数是512,与e4的512拼起来变成1024,所以dec4的in_ch必须写1024。解码器每个DoubleConv的输入通道都是“上采样输出通道数加编码器对应层通道数”,改代码时最容易漏的就是这里。
上采样用nn.ConvTranspose2d是常见方案,kernel_size和stride都设为2,正好把特征图放大两倍。也可以换成nn.Upsample(scale_factor=2, mode="bilinear")加一个卷积,参数更少,但转置卷积能学习的参数更多,效果略好但显存占用也更高。小数据集上两者差距不大。
3.3 in_channels和num_classes:改通道是最容易翻车的一步
换到自己数据集上只需要改两个参数:in_channels对应输入图像通道数,RGB图像是3,灰度图像是1;num_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, eps=1e-6): # pred: [B, C, H, W] 未过softmax的logits # target: [B, H, W] 类别id C = pred.shape[1] pred_softmax = torch.softmax(pred, dim=1) target_onehot = torch.eye(C, device=pred.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, dtype=np.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, dtype=torch.float32) criterion_ce = nn.CrossEntropyLoss(weight=class_weights) criterion_dice = multiclass_dice_loss # 训练时: loss = criterion_ce(pred, mask) + criterion_dice(pred, mask)eps=1e-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.2,mIoU平均下来“看起来还行”,实际模型已经废了一半。所以评估时必须把每个类别的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-4 | Unet参数量大,学习率调大容易在BN层产生震荡 |
| Batch Size | 4~8 | 取决于显存,8是512×512输入下的常见上限 |
| Epochs | 30~60 | 看验证集mIoU是否连续10个epoch不再上升 |
| Weight Decay | 1e-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; a=torch.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(dim=1).squeeze(0).cpu().numpy() # [H, W] h, w = pred.shape rgb = np.zeros((h, w, 3), dtype=np.uint8) for idx, color in id_to_color.items(): rgb[pred == idx] = color return Image.fromarray(rgb)argmax(dim=1)在通道维度上取最大值索引,得到的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(axis=1, keepdims=True) fig, ax = plt.subplots(figsize=(8, 6)) sns.heatmap(cm_norm, annot=True, fmt=".2f", xticklabels=class_names, yticklabels=class_names, ax=ax) ax.set_xlabel("Predicted") ax.set_ylabel("True") plt.savefig("confusion_matrix.png", dpi=150)看这张图时重点看两个位置:对角线之外数值最大的格子——比如第2行第4列是0.35,表示真实类别2有35%被预测成了类别4,说明这两类特征重叠严重;再看是否有一整列数值都偏高,说明网络整体偏向预测成某个类别。遥感场景里最典型的就是植被和农田互相污染,以及阴影区域被分到水体。
混淆矩阵能告诉你哪两类在混淆,但看不出混淆发生在图像的哪个区域,这时叠加可视化就派上用场了。更进一步,用softmax输出的最大值作为每个像素的置信度,低于阈值的区域在原图上直接标黑,能直观看到模型“不知道”的区域:
prob_map = torch.softmax(logits, dim=1).max(dim=1)[0] # [H, W] 每个像素的最高类概率 low_conf = prob_map < 0.5低置信度区域通常集中在类别交界处和图像边缘,如果低置信度区域整片出现,往往是训练数据里该区域本身标注不一致。
从那以后,我每次训练完都强制走一遍这套流程:先出整图预测叠加,再算逐类IoU表,最后看混淆矩阵。这套动作帮我避开过三次拿着高mIoU模型却发现目标类别全错的尴尬——mIoU高只是平均值高,某个关键类别彻底失效时它照样能及格。希望这些经验能帮你少踩几个坑,尤其是数据准备阶段,值得多花时间做清洗和校验。
本文还有配套的精品资源,点击获取