news 2026/9/29 3:17:16

Pytorch Unet多类别语义分割实战:数据管线、训练与避坑指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Pytorch Unet多类别语义分割实战:数据管线、训练与避坑指南

简介:本资源面向具备一定深度学习基础的开发者与图像分析方向学习者,聚焦在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]的 logits,N 是类别数。交叉熵损失要求 target 是[B, H, W]的 LongTensor,每个像素值是 0 到 N-1 的整数索引。而你在 labelme、labelimg 或自研工具里标出来的图,通常是 RGB 三通道 PNG,每个类别对应一种颜色,比如路面是 (128,64,128)、车辆是 (0,0,142)。直接拿这种图当 target,Pytorch 会把它当成三通道浮点,维度对不上,就算强行 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), dtype=np.int64) # 逐类别做全图相等判断,向量化,比逐像素快几个数量级 for idx, color in enumerate(CLASS_COLORS): match = np.all(mask_rgb == np.array(color, dtype=np.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), interpolation=TF.InterpolationMode.BILINEAR) mask = TF.resize(Image.fromarray(mask), (new_h, new_w), interpolation=TF.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(..., interpolation=cv2.INTER_NEAREST)和A.HorizontalFlip这类同时接受 image 和 mask 的接口,不要分开调用。

2.3 类别不均衡:先统计像素频率再决定权重

多类别数据集几乎必然不均衡,背景和天空可能占 80% 像素,行人只占 1%。不处理的话模型很快学会全预测背景,准确率看着很高但 IoU 惨不忍睹。我一般先跑一遍统计脚本,把每个类别的像素占比打出来,再决定用加权交叉熵还是 Dice 组合。

def compute_class_freq(dataset, num_classes): counts = np.zeros(num_classes, dtype=np.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(f"class {c}: {f:.4%}") return freq # 权重取频率倒数并归一化,背景权重可再压低 freq = compute_class_freq(train_ds, num_classes=6) weights = 1.0 / (freq + 1e-6) weights = weights / weights.sum() * num_classes weights[0] *= 0.5 # 背景通常不需要那么高权重 weights = torch.tensor(weights, dtype=torch.float32)

这段统计跑一次就够,结果存下来。权重不是越极端越好,如果某个类频率是 0.01%,倒数权重会大到让训练震荡,这时更适合用 Dice loss 或 Focal loss 兜底。参数上weights[0] *= 0.5是我自己的经验值,背景权重压低能逼模型关注小类,但压太狠会让边界变毛糙,需要看验证集 IoU 微调。

3. Unet 结构改造:从单通道输出到多类别 logits

3.1 输出通道数、上采样方式和 skip connection 的三个必改点

原始 Unet 论文是二分类,输出 1 通道加 sigmoid。多类别要改三处:第一,最后 1x1 卷积输出通道改成 N,不要接 sigmoid,直接输出 logits 给CrossEntropyLoss;第二,上采样用nn.ConvTranspose2d或nn.Upsample(mode='bilinear')加卷积,前者可学习但容易产生棋盘格,后者更平滑,我一般用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, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch=3, num_classes=6, base=32): 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_factor=2, mode='bilinear', align_corners=False) self.dec4 = DoubleConv(base*16 + base*8, base*8) self.up3 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) self.dec3 = DoubleConv(base*8 + base*4, base*4) self.up2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) self.dec2 = DoubleConv(base*4 + base*2, base*2) self.up1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) 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], 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.head(d1) # [B, num_classes, H, W]

参数说明:base=32是通道基数,显存够可以调到 64,小数据集 32 足够且不容易过拟合;align_corners=False在 Pytorch 新版本里是推荐值,和TF.resize的默认行为更一致,能减少上采样错位;bias=False配合 BatchNorm 是标准做法,省参数且不影响表达。如果输入尺寸不是 16 的倍数,四次下采样后 concat 会因尺寸差 1 报错,所以训练和推理的输入尺寸统一 resize 到 16 的倍数,比如 512x512。

3.2 损失函数与忽略标签:让模型不学无标注区域

自己数据集常有未标注区域,比如图像边缘或难标的目标,这些像素不该参与 loss。做法是在掩码里给它们一个固定索引,比如 255,然后CrossEntropyLoss(ignore_index=255)。同时把类别权重传进去。

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(weight=weights, ignore_index=255) # 训练一步 model = UNet(in_ch=3, num_classes=num_classes).cuda() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-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_index=255让这些像素的梯度为 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, ignore=255): 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, minlength=num_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

参数上minlength=num_classes**2保证混淆矩阵尺寸固定;ignore=255和损失函数保持一致。每个 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 scaling,NaN 常从 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, window=512, stride=384): model.eval() _, _, H, W = img_tensor.shape prob = torch.zeros(num_classes, H, W, device=img_tensor.device) count = torch.zeros(1, H, W, device=img_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, dim=1)[0] count[:, y1:y2, x1:x2] += 1 prob = prob / count return prob.argmax(0)

参数上window要和训练尺寸一致,stride取 window 的 0.75 倍左右,重叠越多越平滑但越慢。推理完还可以做一步类别后处理:对每个类别二值掩码做开运算去噪点,再取最大连通域,能去掉零散误检。验证方法上,我习惯留 10% 数据完全不参与训练,推理后算 mIoU 并可视化三张最差样本,看是标注问题还是模型问题。这套流程跑通后,换数据集只需要改CLASS_COLORS和权重,Unet 主体不用动。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/29 3:17:03

PyCharm高效使用指南:从环境配置到调试重构的完整攻略

1. 先别急着写代码&#xff0c;把PyCharm调教成趁手的兵器每次看到刚装好PyCharm就噼里啪啦开写的新人&#xff0c;我都替他着急。编辑器这东西&#xff0c;默认配置顶多算个“能用的状态”&#xff0c;离“好用”差着十万八千里。PyCharm跟别的IDE不一样&#xff0c;它的逻辑和…

作者头像 李华
网站建设 2026/9/29 3:15:13

从零手搓AI工程:数据管道、模型训练与服务化全链路实践

1. 从零手搓AI工程&#xff1a;为什么我不建议你直接调包很多人一听到“AI工程”这四个字&#xff0c;第一反应就是打开某个云平台&#xff0c;拖几个组件&#xff0c;调几个API&#xff0c;然后跑通一个Demo&#xff0c;就觉得自己已经入门了。我刚开始接触这个方向的时候也是…

作者头像 李华
网站建设 2026/9/29 3:14:15

Module Builder——Gem200之Command模块

GEM200 Remote Command&#xff08;远程命令&#xff09;是SEMI E30标准定义的核心功能&#xff0c;允许上位机&#xff08;Host&#xff09;向设备下发指令以控制运行状态或执行特定操作。这是实现半导体产线全自动远程控制的基础。核心通信机制S2F41 Host Command SendHost通…

作者头像 李华
网站建设 2026/9/29 3:12:22

具身智能协同演化动力学(18):原生底座的标准化与模块化演进

前沿技术探索&#xff1a;TVA智能体&#xff08;简称TVA&#xff09;TVA智能体&#xff08;亦称“AI智能体视觉”&#xff09;是依托Transformer架构与“因式智能体”理论构建的新型工业视觉系统&#xff0c;也是当前最具代表性的具身视觉技术之一。它有机融合深度强化学习&…

作者头像 李华
网站建设 2026/9/29 3:11:52

数哈多应用授权系统V1.0.1版本更新介绍

数哈多应用授权系统 2026.07.14 V1.0.1 1.系统完全重构 2.前端UI更换 3.已知问题修复 4.防字典爆破 5.防绕授权检测 6.支持应用等级区分 7.可限制实名购买授权

作者头像 李华