简介:这份资源面向具备一定深度学习基础的开发者与图像分析方向的学习者,提供在PyTorch框架下实现Unet多类别语义分割的完整代码工程,可用于医学影像、遥感图像等场景的像素级分类任务。压缩包共46个文件,约69KB,以19个py源码文件为核心,涵盖网络建模、数据加载与增强、损失函数、学习率调度、评价指标及训练保存等模块,另有24个pyc编译文件、2个txt与1个json配置说明,目录按modeling、dataloaders、utils等分层组织,结构清晰便于二次开发。目前已有15264人学习下载,热度较高。读者可据此掌握编码器与解码器搭建、跳跃连接、多通道输出层设计、交叉熵损失与IoU评估等关键环节,并参考数据增强、注意力机制、多尺度训练等优化思路,快速迁移到自己的多类别分割数据集上。
1. 从二分类到多类别:这套 Unet 代码到底能替你省掉哪些重复劳动
如果你手头有一批自己标注的遥感影像、工业缺陷图或者医学切片,想跑一个能区分背景、目标 A、目标 B 甚至更多类别的分割模型,大概率会卡在同一个地方:网上能找到的 Unet 教程,十篇里有八篇是二分类的。输入一张图,输出一张单通道 mask,sigmoid 一压,阈值一卡,完事。可一旦类别数变成 3、5、8,损失函数要换、输出通道要改、标签编码要重做、评估指标也得跟着变,原来那套代码几乎要重写一遍。
这份资源就是冲着这个痛点来的。它是一套基于 PyTorch 的 Unet 多类别语义分割实现,核心价值不在于 Unet 本身——Unet 的结构早就被讲烂了——而在于它把「多类别」这条链路上的每个环节都打通了:数据读取时怎么处理彩色标签图、模型输出层怎么改、损失函数用 CrossEntropyLoss 还是带权重的版本、训练过程中怎么算 mIoU、预测结果怎么还原成可视化的彩色 mask。适合已经跑通过二分类 Unet、现在要迁移到自己多类别数据集上的从业者,也适合刚接触语义分割但不想在环境配置和标签编码上反复翻车的新手。
我见过太多人在这类项目上耗掉的时间,不是花在调模型上,而是花在「标签图是 RGB 的,读进来变成三通道,跟输出对不上」这种问题上。这套代码把这些坑都填过了,你拿到手之后,主要精力可以放在数据质量和类别设计上,而不是跟张量维度较劲。
2. 多类别 Unet 的数据管线:从彩色标签图到可训练张量
2.1 为什么多类别任务不能直接复用二分类的数据读取
二分类语义分割的标签通常是一张单通道灰度图,像素值 0 或 255,读进来除以 255 就变成 0 和 1,直接当 float 标签用。但多类别数据集的标签往往是 RGB 彩色图,每个类别对应一种颜色,比如背景是黑色 (0,0,0)、类别 A 是红色 (255,0,0)、类别 B 是绿色 (0,255,0)。这种图直接读进来是 H×W×3 的张量,模型输出是 H×W×C(C 是类别数),两者根本对不上。
常见做法是维护一个颜色到类别索引的映射表,把 RGB 标签图逐像素转换成单通道的类别索引图。这个转换必须在 Dataset 的__getitem__里完成,而且要用 numpy 做向量化操作,不能写双重循环——我试过用 for 循环逐像素查表,一张 512×512 的图要跑好几秒,训练时 dataloader 直接成为瓶颈。
import numpy as np import torch from torch.utils.data import Dataset from PIL import Image class MultiClassSegDataset(Dataset): def __init__(self, img_dir, mask_dir, color_map, transform=None): """ img_dir: 原图目录 mask_dir: 彩色标签图目录 color_map: dict, {(R,G,B): class_index} transform: 可选的增强 """ self.img_dir = img_dir self.mask_dir = mask_dir self.color_map = color_map self.transform = transform self.images = sorted(os.listdir(img_dir)) def __getitem__(self, idx): img = Image.open(os.path.join(self.img_dir, self.images[idx])).convert('RGB') mask = Image.open(os.path.join(self.mask_dir, self.images[idx])).convert('RGB') img = np.array(img, dtype=np.float32) / 255.0 mask = np.array(mask, dtype=np.uint8) # 向量化颜色到索引的转换 class_mask = np.zeros(mask.shape[:2], dtype=np.int64) for color, class_idx in self.color_map.items(): match = np.all(mask == np.array(color, dtype=np.uint8), axis=-1) class_mask[match] = class_idx if self.transform: augmented = self.transform(image=img, mask=class_mask) img, class_mask = augmented['image'], augmented['mask'] img = torch.from_numpy(img).permute(2, 0, 1).float() class_mask = torch.from_numpy(class_mask).long() return img, class_mask这段代码的关键点有三个。第一,color_map的键是 RGB 元组,值是从 0 开始的类别索引,0 通常留给背景。第二,np.all(mask == color, axis=-1)生成一个布尔掩码,直接赋值比逐像素循环快两个数量级。第三,标签必须转成long类型,因为后面 CrossEntropyLoss 要求 target 是 int64。
注意:如果你的标签图里存在抗锯齿产生的过渡色,比如红色边缘有 (200,50,50) 这种像素,上面的精确匹配会漏掉它们,导致这些像素被归为背景。解决办法是在标注阶段就关闭抗锯齿,或者用最近邻颜色匹配代替精确匹配。
2.2 数据增强在多类别分割里的特殊约束
分类任务的数据增强可以随便翻转、旋转、调色,但分割任务的增强必须保证图像和标签同步变换,而且涉及颜色变换时要格外小心。比如 ColorJitter 只作用于原图,标签图不能跟着变,否则颜色到类别的映射就乱了。再比如 RandomRotation 如果用了双线性插值,标签图会产生新的颜色值,精确匹配直接失效。
我一般用 albumentations 这个库,它对分割任务的支持比较完善,能保证 image 和 mask 同步。但要注意,传给 albumentations 的 mask 已经是类别索引图了,不是 RGB 图,所以只能用那些不改变像素值的几何变换,比如 HorizontalFlip、VerticalFlip、RandomRotate90。如果要加旋转角度或者缩放,插值方式必须选最近邻。
import albumentations as A train_transform = A.Compose([ A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=15, interpolation=cv2.INTER_NEAREST, p=0.5), ])interpolation=cv2.INTER_NEAREST是必须的,双线性插值会在类别边界产生中间值,这些值不属于任何类别,训练时会被当成噪声。ShiftScaleRotate 的border_mode默认是反射填充,对标签图来说反射填充可能引入不存在的类别,建议改成border_mode=cv2.BORDER_CONSTANT,填充值设为 0(背景)。
2.3 类别不平衡的采样策略
多类别数据集里,背景像素往往占 80% 以上,小目标类别可能只占 1% 到 2%。如果直接按像素算损失,模型会倾向于全部预测成背景,准确率看起来很高,但 mIoU 惨不忍睹。常见做法有两种:一是在损失函数里给每个类别加权,权重和类别频率成反比;二是用 WeightedRandomSampler 让包含稀有类别的图像被更频繁地采样。
我一般先算一遍训练集里每个类别的像素占比,然后取倒数作为权重传给 CrossEntropyLoss。如果某个类别的占比低于 0.5%,我会额外加一个 WeightedRandomSampler,双管齐下。但要注意,采样器的权重是按图像算的,不是按像素算的,所以需要先统计每张图里稀有类别的像素数,再决定这张图的采样权重。
from torch.utils.data import WeightedRandomSampler # 统计每张图的稀有类别像素占比 image_weights = [] for mask_path in mask_list: mask = np.array(Image.open(mask_path)) rare_pixels = np.sum(np.isin(mask, rare_class_indices)) image_weights.append(rare_pixels / mask.size + 1e-6) sampler = WeightedRandomSampler(weights=image_weights, num_samples=len(image_weights), replacement=True)replacement=True表示有放回采样,这样稀有类别的图会被重复抽到。num_samples设成和数据集大小一样,保证每个 epoch 的迭代次数不变。这个采样器传给 DataLoader 的sampler参数,此时shuffle必须设为 False,否则会冲突。
3. Unet 输出层改造与损失函数选型:让模型真正输出多类别概率
3.1 输出通道数与上采样方式的调整
标准 Unet 的二分类版本在最后一层用 1×1 卷积把特征图压成单通道,然后接 sigmoid。多类别版本要把输出通道数改成类别数 C,然后接 softmax(或者不接,因为 CrossEntropyLoss 内部会做 log_softmax)。这一步看起来简单,但有个细节容易翻车:上采样方式。
原始 Unet 论文用的是转置卷积(ConvTranspose2d),但转置卷积容易产生棋盘格伪影,尤其是在类别边界处。我一般把上采样换成nn.Upsample(mode='bilinear', align_corners=True)加一个 3×3 卷积,这样输出更平滑,边界更干净。如果你追求轻量化,也可以用nn.PixelShuffle,但要注意通道数的整除关系。
import torch.nn as nn import torch.nn.functional as F class UNetMultiClass(nn.Module): def __init__(self, in_channels=3, num_classes=5, base_channels=64): super().__init__() # 编码器省略,假设已有 down1~down4 和 bottleneck # 解码器上采样块 self.up4 = nn.Sequential( nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True), nn.Conv2d(base_channels * 8, base_channels * 4, 3, padding=1), nn.BatchNorm2d(base_channels * 4), nn.ReLU(inplace=True) ) # ... 其余上采样块类似 self.final_conv = nn.Conv2d(base_channels, num_classes, kernel_size=1) def forward(self, x): # 编码和解码过程省略 # 最终输出 shape: (B, num_classes, H, W) return self.final_conv(x)final_conv的输出没有接 softmax,这是故意的。PyTorch 的nn.CrossEntropyLoss期望输入是未归一化的 logits,内部会做 log_softmax 和 NLLLoss。如果你在模型里加了 softmax,再传给 CrossEntropyLoss,相当于做了两次归一化,梯度会变得很小,训练几乎不动。这个坑我踩过,loss 从第一个 epoch 就卡在 1.6 左右不降,排查了半天才发现是 softmax 重复了。
3.2 CrossEntropyLoss 的权重参数怎么设
nn.CrossEntropyLoss有一个weight参数,可以给每个类别指定损失权重。权重应该和类别频率成反比,但不要直接用 1/freq,因为频率极低的类别权重会大到让训练不稳定。常见做法是取频率倒数的平方根,或者用中位数频率除以各类别频率。
# 假设 class_pixel_counts 是每个类别的像素总数 freq = class_pixel_counts / class_pixel_counts.sum() weights = 1.0 / (freq + 1e-6) weights = weights / weights.sum() * num_classes # 归一化,均值为1 weights = torch.tensor(weights, dtype=torch.float32).to(device) criterion = nn.CrossEntropyLoss(weight=weights, ignore_index=255)ignore_index=255是给那些「不确定」或「忽略」的像素留的,比如标注边界模糊的区域。如果你的数据集没有这类像素,可以去掉这个参数。权重的归一化不是必须的,但能让 loss 的量级和不同数据集之间可比,方便调学习率。
3.3 Dice Loss 与 CrossEntropyLoss 的组合
如果类别极度不平衡,单用 CrossEntropyLoss 即使加了权重,小类别的梯度仍然可能被背景淹没。这时候可以加一个 Dice Loss 作为辅助。Dice Loss 直接优化预测区域和真实区域的重叠度,对类别不平衡不敏感。但 Dice Loss 的梯度在预测接近 0 或 1 时会变得很小,单独用容易训练不动,所以一般是 CE 和 Dice 按比例加权。
class DiceLoss(nn.Module): def __init__(self, smooth=1e-6): super().__init__() self.smooth = smooth def forward(self, logits, targets): probs = F.softmax(logits, dim=1) targets_one_hot = F.one_hot(targets, num_classes=logits.shape[1]) targets_one_hot = targets_one_hot.permute(0, 3, 1, 2).float() intersection = (probs * targets_one_hot).sum(dim=(0, 2, 3)) union = probs.sum(dim=(0, 2, 3)) + targets_one_hot.sum(dim=(0, 2, 3)) dice = (2. * intersection + self.smooth) / (union + self.smooth) return 1 - dice.mean() # 组合损失 ce_loss = nn.CrossEntropyLoss(weight=weights) dice_loss = DiceLoss() total_loss = ce_loss(logits, targets) + 0.5 * dice_loss(logits, targets)Dice Loss 里的smooth是防止分母为零的平滑项,一般设 1e-6 到 1e-5。F.one_hot要求 targets 是 long 类型,且值在 [0, num_classes) 范围内。如果 targets 里有 ignore_index,需要先把它 mask 掉再算 Dice,否则会报错。
4. 训练循环与 mIoU 计算:怎么判断模型是真的在学
4.1 训练循环里必须记录的几个量
很多人训练分割模型时只看 loss,loss 降了就以为模型在变好,结果一预测发现全是背景。这是因为 loss 和实际分割质量之间不是线性关系,尤其是加了类别权重之后,loss 的绝对值已经没有直观意义了。我一般会在训练循环里同时记录三个量:总 loss、每个类别的像素准确率、以及 mIoU。
mIoU 的计算需要维护一个混淆矩阵,每个 batch 更新一次,epoch 结束时算全局 mIoU。混淆矩阵的大小是 C×C,C 是类别数,对于 10 类以内的任务完全够用。如果类别数上百,混淆矩阵会很大,但那种情况一般也不会用 Unet 了。
def update_confusion_matrix(conf_matrix, preds, targets, num_classes): preds = torch.argmax(preds, dim=1).view(-1) targets = targets.view(-1) mask = (targets >= 0) & (targets < num_classes) preds, targets = preds[mask], targets[mask] indices = targets * num_classes + preds conf_matrix += torch.bincount(indices, minlength=num_classes**2 ).reshape(num_classes, num_classes) return conf_matrix def compute_miou(conf_matrix): intersection = torch.diag(conf_matrix) union = conf_matrix.sum(dim=1) + conf_matrix.sum(dim=0) - intersection iou = intersection / (union + 1e-6) return iou.mean().item(), ioutorch.bincount是算混淆矩阵最快的方式,比用 sklearn 的 confusion_matrix 快很多,而且可以直接在 GPU 上跑。minlength必须设成 num_classes²,否则如果某个类别在 batch 里没出现,bincount 返回的长度会不够,reshape 会报错。
4.2 学习率调度与早停策略
Unet 这类编码器-解码器结构,编码器通常用预训练权重(比如 ResNet34),解码器是随机初始化的。这两部分的理想学习率不一样,编码器应该用更小的学习率(比如 1e-4),解码器可以用大一点(比如 1e-3)。如果统一用一个学习率,要么编码器被破坏,要么解码器学得太慢。
我一般用torch.optim.Adam加参数组,编码器和解码器分开设学习率。调度器用ReduceLROnPlateau,监控验证集的 mIoU,如果连续 5 个 epoch 不提升就把学习率乘以 0.5。早停的 patience 设 10 到 15,具体看数据集大小和训练速度。
optimizer = torch.optim.Adam([ {'params': model.encoder.parameters(), 'lr': 1e-4}, {'params': model.decoder.parameters(), 'lr': 1e-3}, ]) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode='max', factor=0.5, patience=5, verbose=True ) best_miou = 0.0 patience_counter = 0 for epoch in range(num_epochs): train_one_epoch(...) val_miou = validate(...) scheduler.step(val_miou) if val_miou > best_miou: best_miou = val_miou torch.save(model.state_dict(), 'best_model.pth') patience_counter = 0 else: patience_counter += 1 if patience_counter >= 15: breakmode='max'是因为 mIoU 越大越好,如果是监控 loss 就用mode='min'。verbose=True会在学习率变化时打印一行日志,方便确认调度器有没有在工作。早停的 patience 不要设太小,分割任务的 mIoU 波动比较大,偶尔一两个 epoch 不提升很正常。
4.3 验证阶段的内存优化
验证时不需要计算梯度,用torch.no_grad()包起来能省不少显存。但如果验证集很大,一次性把所有图跑完仍然可能 OOM。这时候可以分 batch 跑,每个 batch 算完混淆矩阵后把预测结果丢掉,只保留混淆矩阵。另外,验证时的数据增强要关掉,只保留 resize 和归一化。
@torch.no_grad() def validate(model, dataloader, num_classes, device): model.eval() conf_matrix = torch.zeros(num_classes, num_classes, dtype=torch.int64).to(device) for imgs, masks in dataloader: imgs, masks = imgs.to(device), masks.to(device) logits = model(imgs) conf_matrix = update_confusion_matrix(conf_matrix, logits, masks, num_classes) miou, per_class_iou = compute_miou(conf_matrix) return miou, per_class_ioumodel.eval()会关掉 BatchNorm 的 running stats 更新和 Dropout,这两者在验证时必须关掉,否则结果不可复现。混淆矩阵用 int64 是为了防止像素数太多溢出,int32 在千万级像素时可能不够。
5. 多类别分割的避坑与排查:那些让我重跑过训练的问题
5.1 现象:loss 正常下降但 mIoU 始终在 0.1 以下
原因通常是标签编码错了。比如 color_map 里把背景映射成了 1 而不是 0,或者某个类别的颜色写错了,导致大量像素被归到错误的类别。还有一种可能是模型输出的通道顺序和标签的类别索引不一致,比如模型输出通道 0 对应类别 1,但标签里 0 是背景。
解决方法是先拿一张训练图跑一遍前向,把预测的 argmax 结果和标签图并排可视化出来。如果预测全是某一个类别,检查 color_map 和 final_conv 的输出通道数是否匹配。如果预测看起来有结构但类别全错,检查 color_map 的键值对有没有写反。
5.2 现象:训练到一半 loss 突然变成 NaN
最常见的原因是学习率太大,导致梯度爆炸。尤其是用 Dice Loss 时,如果某个 batch 里某个类别完全没有出现,Dice 的分母会接近 smooth,梯度可能异常大。解决办法是加梯度裁剪,torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0),放在loss.backward()之后、optimizer.step()之前。
另一个可能的原因是 BatchNorm 在 batch size 很小时统计量不稳定。如果显存不够只能跑 batch size 2 或 4,建议把 BatchNorm 换成 GroupNorm,或者用 SyncBatchNorm 跨卡同步。我试过 batch size 2 跑 BatchNorm,训练到第 30 个 epoch 突然 NaN,换成 GroupNorm 后稳定跑完。
5.3 现象:验证集 mIoU 比训练集高很多
这听起来是好事,但通常意味着训练集的增强太强了,或者训练时 Dropout 开得太大,导致训练集上的表现被压低。另一种可能是验证集和训练集的数据分布不一致,比如验证集里简单样本居多。检查一下训练和验证的 transform 是不是差太多,验证集有没有误用了训练集的增强。
如果确认是增强过强,把 RandomRotate90 的概率从 0.5 降到 0.3,或者去掉 ShiftScaleRotate。如果数据分布不一致,重新划分训练验证集,确保两个集合的类别比例接近。
5.4 现象:预测结果在类别边界处有锯齿或空洞
这是上采样方式导致的。转置卷积的棋盘格伪影在边界处表现为锯齿,双线性上采样如果 align_corners 设错会产生偏移。检查nn.Upsample的align_corners参数,设成 True 时像素中心对齐,设成 False 时角点对齐。对于分割任务,一般用 True。
如果还有空洞,可能是感受野不够,小目标在深层特征里已经丢失了。解决办法是在 Unet 的跳跃连接里加注意力模块,或者把输入分辨率提高。我一般先把输入从 256×256 提到 512×512 试试,如果显存不够再考虑改结构。
5.5 现象:推理时单张图预测正常,批量预测结果错乱
这通常是 BatchNorm 在推理时的 running stats 问题。如果模型训练完后没有调用model.eval(),BatchNorm 会用当前 batch 的统计量而不是训练时累积的 running stats,导致批量预测时结果依赖 batch 内其他样本。确保推理前调用了model.eval(),并且用torch.no_grad()包起来。
另一个可能是输入图像的归一化参数不一致。训练时用了 ImageNet 的 mean 和 std,推理时忘了减,或者减错了。把训练时的归一化参数存下来,推理时严格复用。
6. 从训练到落地:ONNX 导出与推理加速的几个实操技巧
训练完模型只是第一步,真正要用起来还得考虑推理速度和部署方式。PyTorch 模型直接推理在 GPU 上还行,但如果要部署到边缘设备或者用 TensorRT 加速,导出 ONNX 是绕不开的一步。多类别分割模型导出 ONNX 时有两个地方容易出问题:一是动态轴的处理,二是 softmax 的位置。
导出时把 batch 维度设为动态,这样同一个模型可以处理单张图和批量图。softmax 不要放在模型里导出,让 ONNX 只输出 logits,后处理时再算 argmax。这样做的原因是 ONNX 的 softmax 在某些推理引擎里支持不好,而且 logits 保留更多信息,方便后续调整阈值。
dummy_input = torch.randn(1, 3, 512, 512).to(device) torch.onnx.export( model, dummy_input, 'unet_multiclass.onnx', input_names=['input'], output_names=['logits'], dynamic_axes={'input': {0: 'batch'}, 'logits': {0: 'batch'}}, opset_version=11 )opset_version=11是比较稳妥的选择,支持大部分算子且兼容性好。如果用了 PixelShuffle 或者自定义算子,可能需要更高的 opset。导出后用onnxruntime跑一遍,对比 PyTorch 和 ONNX 的输出差异,如果 max diff 超过 1e-3,说明某个算子导出有问题,需要逐层排查。
推理加速方面,如果目标平台是 NVIDIA GPU,可以用 TensorRT 把 ONNX 转成 engine,FP16 精度下速度通常能提升 2 到 3 倍,mIoU 掉不到 0.5 个点。转换时注意设置正确的 workspace size,太小会转换失败,太大浪费显存。我一般设 1GB 到 2GB,大部分分割模型够用。
还有一个容易被忽略的点是输入图像的预处理。训练时用的是 PIL 读图加 albumentations 归一化,推理时如果换成 OpenCV 读图,颜色通道顺序会从 RGB 变成 BGR,导致预测结果完全错误。这种问题不会报错,只会让 mIoU 莫名其妙地低,排查起来很费时间。从那以后我每次部署前都强制走一遍「训练预处理 vs 推理预处理」的对比,用同一张图分别过两个管线,确认输出一致再往下走。
希望帮到你。
本文还有配套的精品资源,点击获取