
简介本资源面向具备一定深度学习基础的开发者与图像分析方向学习者聚焦在PyTorch框架下用Unet完成多类别语义分割任务可应用于医学影像、遥感图像等场景。压缩包共46个文件以19个py脚本和24个pyc缓存文件为主另含少量txt与json配置整体约69KB涵盖数据加载、自定义变换、模型定义、损失与指标计算、学习率调度、训练与可视化等模块目录结构清晰便于按功能查阅。目前已有15264人学习下载热度较高。读者可据此搭建从数据预处理、网络构建到训练评估的完整流程理解多类别输出通道设计、交叉熵损失与IoU等指标的使用并参考注意力机制、多尺度训练等进阶思路快速迁移到自己的数据集上实践。1. 从一张标注图到可训练掩码Unet 多类别语义分割到底在做什么你手里有一批自己拍的或标注的图片每张图里可能有路面、车辆、行人、天空、建筑等若干类别你想让模型对每个像素都给出一个类别标签——这就是多类别语义分割。Pytorch 下实现 Unet 做这件事核心链路只有四步把标注图转成单通道的类别索引掩码、搭一个输入三通道输出 N 通道的 Unet、用交叉熵加忽略背景的损失训练、推理时对 N 通道取 argmax 还原成彩色掩码。听起来简单但真正卡住大多数人的不是网络结构而是数据管线标注颜色和类别索引对不上、掩码被双线性插值插出小数、类别极度不均衡导致模型只学会预测背景。这篇笔记就按我实际跑通自己多类别数据集的顺序把每一步的参数、代码和翻车点讲清楚适合已经会写 Pytorch 训练循环、但第一次拿 Unet 上自己数据的人。2. 数据管线把彩色标注图变成 Unet 能吃的类别索引掩码2.1 为什么不能直接把 RGB 标注图喂给交叉熵Unet 做多类别分割输出是[B, N, H, W]的 logitsN 是类别数。交叉熵损失要求 target 是[B, H, W]的 LongTensor每个像素值是 0 到 N-1 的整数索引。而你在 labelme、labelimg 或自研工具里标出来的图通常是 RGB 三通道 PNG每个类别对应一种颜色比如路面是 (128,64,128)、车辆是 (0,0,142)。直接拿这种图当 targetPytorch 会把它当成三通道浮点维度对不上就算强行 reshape 也会把颜色值当成类别号训练必然发散。所以第一步是建立一张颜色到索引的映射表把 RGB 标注图逐像素查表转成单通道索引图。常见做法是维护一个class_colors列表顺序就是类别索引顺序转换时用向量化查表而不是 Python 循环否则几千张图会慢到怀疑人生。import numpy as np from PIL import Image # 类别顺序即索引顺序背景放 0 CLASS_COLORS [ (0, 0, 0), # 0 背景 (128, 64, 128), # 1 路面 (0, 0, 142), # 2 车辆 (220, 20, 60), # 3 行人 (70, 130, 180), # 4 天空 (119, 11, 32), # 5 建筑 ] def rgb_to_index(mask_rgb: np.ndarray) - np.ndarray: mask_rgb: [H, W, 3] uint8 - [H, W] int64 h, w, _ mask_rgb.shape index np.zeros((h, w), dtypenp.int64) # 逐类别做全图相等判断向量化比逐像素快几个数量级 for idx, color in enumerate(CLASS_COLORS): match np.all(mask_rgb np.array(color, dtypenp.uint8), axis-1) index[match] idx return index # 使用 rgb np.array(Image.open(label.png).convert(RGB)) idx rgb_to_index(rgb) Image.fromarray(idx.astype(np.uint8)).save(label_index.png)这段代码的逻辑是对每个类别颜色用np.all(..., axis-1)生成一个布尔掩码把该类别位置赋成对应索引。参数上要注意CLASS_COLORS的顺序必须和后面模型输出通道顺序、损失函数权重顺序完全一致一旦错位训练 loss 会正常下降但预测结果全乱这是最隐蔽的坑之一。另外如果标注图里有抗锯齿边缘产生的过渡色这些像素不会被任何类别匹配到会留在索引 0相当于被当成背景需要在标注规范里明确禁止羽化。2.2 同步增强图像和掩码必须用同一组随机参数自己数据集通常样本少必须做增强。但图像可以做双线性插值、颜色抖动掩码绝对不行——掩码一旦被插值就会产生 1.5 这种小数类别交叉熵直接报错或静默出错。正确做法是图像和掩码共享几何变换参数且掩码统一用最近邻插值。import random import torch from torch.utils.data import Dataset import torchvision.transforms.functional as TF class SegDataset(Dataset): def __init__(self, img_paths, mask_paths, size(512, 512)): self.img_paths img_paths self.mask_paths mask_paths self.size size def __len__(self): return len(self.img_paths) def __getitem__(self, i): img Image.open(self.img_paths[i]).convert(RGB) mask np.array(Image.open(self.mask_paths[i])) # 已是单通道索引图 # 同步随机缩放裁剪先算同一组参数 scale random.uniform(0.8, 1.25) new_h int(self.size[0] * scale) new_w int(self.size[1] * scale) img TF.resize(img, (new_h, new_w), interpolationTF.InterpolationMode.BILINEAR) mask TF.resize(Image.fromarray(mask), (new_h, new_w), interpolationTF.InterpolationMode.NEAREST) # 同步随机裁剪 top random.randint(0, max(0, new_h - self.size[0])) left random.randint(0, max(0, new_w - self.size[1])) img TF.crop(img, top, left, self.size[0], self.size[1]) mask TF.crop(mask, top, left, self.size[0], self.size[1]) # 同步水平翻转 if random.random() 0.5: img TF.hflip(img) mask TF.hflip(mask) img TF.to_tensor(img) # [3,H,W] float 0~1 mask torch.from_numpy(np.array(mask)).long() # [H,W] int64 return img, mask关键参数说明TF.resize对掩码必须显式指定NEAREST默认的 BILINEAR 会毁掉类别索引裁剪的top/left对图像和掩码用同一组值不能各自随机to_tensor只对图像做掩码保持整数。如果用了 albumentations对应的是A.Resize(..., interpolationcv2.INTER_NEAREST)和A.HorizontalFlip这类同时接受 image 和 mask 的接口不要分开调用。2.3 类别不均衡先统计像素频率再决定权重多类别数据集几乎必然不均衡背景和天空可能占 80% 像素行人只占 1%。不处理的话模型很快学会全预测背景准确率看着很高但 IoU 惨不忍睹。我一般先跑一遍统计脚本把每个类别的像素占比打出来再决定用加权交叉熵还是 Dice 组合。def compute_class_freq(dataset, num_classes): counts np.zeros(num_classes, dtypenp.int64) for _, mask in dataset: m mask.numpy() for c in range(num_classes): counts[c] (m c).sum() freq counts / counts.sum() for c, f in enumerate(freq): print(fclass {c}: {f:.4%}) return freq # 权重取频率倒数并归一化背景权重可再压低 freq compute_class_freq(train_ds, num_classes6) weights 1.0 / (freq 1e-6) weights weights / weights.sum() * num_classes weights[0] * 0.5 # 背景通常不需要那么高权重 weights torch.tensor(weights, dtypetorch.float32)这段统计跑一次就够结果存下来。权重不是越极端越好如果某个类频率是 0.01%倒数权重会大到让训练震荡这时更适合用 Dice loss 或 Focal loss 兜底。参数上weights[0] * 0.5是我自己的经验值背景权重压低能逼模型关注小类但压太狠会让边界变毛糙需要看验证集 IoU 微调。3. Unet 结构改造从单通道输出到多类别 logits3.1 输出通道数、上采样方式和 skip connection 的三个必改点原始 Unet 论文是二分类输出 1 通道加 sigmoid。多类别要改三处第一最后 1x1 卷积输出通道改成 N不要接 sigmoid直接输出 logits 给CrossEntropyLoss第二上采样用nn.ConvTranspose2d或nn.Upsample(modebilinear)加卷积前者可学习但容易产生棋盘格后者更平滑我一般用Upsample加 3x3 卷积第三skip connection 的通道数要保证编码器和解码器对应层一致否则 concat 时维度报错。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.net nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch3, num_classes6, base32): super().__init__() # 编码器 self.enc1 DoubleConv(in_ch, base) self.enc2 DoubleConv(base, base*2) self.enc3 DoubleConv(base*2, base*4) self.enc4 DoubleConv(base*4, base*8) self.pool nn.MaxPool2d(2) # 瓶颈 self.bottleneck DoubleConv(base*8, base*16) # 解码器上采样后 concat再双卷积 self.up4 nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.dec4 DoubleConv(base*16 base*8, base*8) self.up3 nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.dec3 DoubleConv(base*8 base*4, base*4) self.up2 nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.dec2 DoubleConv(base*4 base*2, base*2) self.up1 nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.dec1 DoubleConv(base*2 base, base) self.head nn.Conv2d(base, num_classes, 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], 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.head(d1) # [B, num_classes, H, W]参数说明base32是通道基数显存够可以调到 64小数据集 32 足够且不容易过拟合align_cornersFalse在 Pytorch 新版本里是推荐值和TF.resize的默认行为更一致能减少上采样错位biasFalse配合 BatchNorm 是标准做法省参数且不影响表达。如果输入尺寸不是 16 的倍数四次下采样后 concat 会因尺寸差 1 报错所以训练和推理的输入尺寸统一 resize 到 16 的倍数比如 512x512。3.2 损失函数与忽略标签让模型不学无标注区域自己数据集常有未标注区域比如图像边缘或难标的目标这些像素不该参与 loss。做法是在掩码里给它们一个固定索引比如 255然后CrossEntropyLoss(ignore_index255)。同时把类别权重传进去。import torch num_classes 6 weights torch.tensor([0.5, 1.2, 1.5, 2.0, 0.8, 1.3]) # 按 2.3 统计结果填 criterion torch.nn.CrossEntropyLoss(weightweights, ignore_index255) # 训练一步 model UNet(in_ch3, num_classesnum_classes).cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) img, mask img.cuda(), mask.cuda() logits model(img) # [B, 6, H, W] loss criterion(logits, mask) # mask 里 255 的位置被忽略 loss.backward() optimizer.step()逻辑上ignore_index255让这些像素的梯度为 0不影响参数更新。参数上weight长度必须等于num_classes顺序和CLASS_COLORS一致ignore_index不能设成 0 到 N-1 之间的值否则会误伤真实类别。如果发现 loss 一直不降先检查 mask 的 dtype 是不是 long再检查 mask 最大值是否超过 num_classes-1这两个错误最常见。3.3 训练循环里必须打印的指标mIoU 而不是像素准确率像素准确率在不均衡数据上会骗人必须算 mIoU。实现方式是用混淆矩阵累加每个 epoch 结束后算每个类别的 IoU 再平均。def update_confusion(conf, pred, target, num_classes, ignore255): pred pred.argmax(1).view(-1) target target.view(-1) valid target ! ignore pred, target pred[valid], target[valid] idx target * num_classes pred conf torch.bincount(idx, minlengthnum_classes**2).reshape(num_classes, num_classes) def compute_miou(conf): iou [] for c in range(conf.shape[0]): tp conf[c, c].item() fp conf[:, c].sum().item() - tp fn conf[c, :].sum().item() - tp if tp fp fn 0: continue iou.append(tp / (tp fp fn)) return sum(iou) / len(iou), iou参数上minlengthnum_classes**2保证混淆矩阵尺寸固定ignore255和损失函数保持一致。每个 epoch 打印 mIoU 和各类 IoU能立刻看出模型是不是只学了背景。我一般还会存一份验证集预测可视化每 5 个 epoch 存一张肉眼比数字更早发现问题。4. 避坑与排查多类别 Unet 训练里最常见的五类翻车4.1 现象loss 正常下降但预测全是背景原因类别极度不均衡背景权重或样本量压倒其他类模型找到局部最优就是全预测背景。解决先确认weights是否生效把背景权重压到 0.3 以下同时引入 Dice loss 联合训练Dice 对小类更敏感。另外检查验证集 mIoU如果背景 IoU 接近 1 其他接近 0基本就是这个原因。4.2 现象训练中途 loss 突然变 NaN原因学习率过大或某批数据里 mask 含非法值比如 255 没被 ignore或索引超过 num_classes。解决先把 lr 降到 1e-4 试再在 Dataset 里加断言assert mask.max() num_classes or mask.max() 255把非法样本挡在训练前。混合精度训练时还要注意 loss scalingNaN 常从 fp16 溢出开始。4.3 现象验证集 mIoU 比训练集低很多且预测边界抖动原因过拟合加掩码插值错误。解决增强里确认掩码只用 NEAREST加 Dropout 或减小 base 通道如果标注本身边界就毛糙考虑在 loss 里对边界像素降权或者用 3x3 形态学开运算后处理预测掩码。4.4 现象concat 时报尺寸不匹配原因输入尺寸不是 16 的倍数四次下采样后奇数尺寸除不尽。解决Dataset 里统一 resize 到 512x512 或 480x480 这类 16 的倍数如果必须保持原尺寸用F.interpolate把解码器特征对齐到编码器尺寸再 concat。4.5 现象推理时 argmax 出来的类别整体偏移一位原因CLASS_COLORS顺序和训练时weights、模型输出通道顺序不一致或者转换脚本里背景没放 0。解决把类别映射表写成一个独立配置文件训练、转换、推理三处都 import 同一份禁止各写各的。这个坑我踩过loss 曲线完全正常但预测颜色全错排查了一下午。5. 进阶技巧用滑动窗口推理大图并做类别后处理自己数据集里常有超过显存的大图直接 resize 会丢小目标。我一般用滑动窗口加重叠推理再对拼接后的概率图做 argmax。窗口 512、步长 384重叠区域取平均概率能显著减少拼接缝。torch.no_grad() def sliding_inference(model, img_tensor, num_classes, window512, stride384): model.eval() _, _, H, W img_tensor.shape prob torch.zeros(num_classes, H, W, deviceimg_tensor.device) count torch.zeros(1, H, W, deviceimg_tensor.device) for y in range(0, H, stride): for x in range(0, W, stride): y2 min(y window, H) x2 min(x window, W) y1 max(0, y2 - window) x1 max(0, x2 - window) patch img_tensor[:, :, y1:y2, x1:x2] logits model(patch) prob[:, y1:y2, x1:x2] F.softmax(logits, dim1)[0] count[:, y1:y2, x1:x2] 1 prob prob / count return prob.argmax(0)参数上window要和训练尺寸一致stride取 window 的 0.75 倍左右重叠越多越平滑但越慢。推理完还可以做一步类别后处理对每个类别二值掩码做开运算去噪点再取最大连通域能去掉零散误检。验证方法上我习惯留 10% 数据完全不参与训练推理后算 mIoU 并可视化三张最差样本看是标注问题还是模型问题。这套流程跑通后换数据集只需要改CLASS_COLORS和权重Unet 主体不用动。希望帮到你。本文还有配套的精品资源点击获取