UNet与RSDDs:小样本轨道缺陷分割的完整实践指南
2026/9/16 20:51:30 网站建设 项目流程

如果让我一句话评价UNet在轨道缺陷检测里的地位,我会说它是“小样本场景里最容易出效果的分割网络”。去年我接轨道缺陷检测项目时,第一反应是把它当成目标检测来做——毕竟当时 Faster R-CNN 用得顺手,结果很快被打脸:裂纹、擦伤、剥落这类缺陷形状太不规矩,用矩形框根本框不干净,两个紧挨在一起的缺陷经常被判定成一个大框,召回率上去了,精确率一塌糊涂。后来老老实实切回UNet做像素级分割,轨道缺陷检测才算真正跑通。这篇文章就围绕UNet、RSDDs数据集和完整复现代码,把从数据集处理到模型落地的全流程拆开讲一遍,包括网络原理、训练细节和那些文档里不写的坑。适合想尽快跑通二值分割、又对小样本场景发愁的读者,也适合准备复现学术分割模型的同学。

1. 为什么轨道缺陷检测要用UNet:从检测逻辑到分割逻辑

1.1 轨道缺陷的难点不在“有没有”,而在“边界在哪”

轨道表面缺陷和普通工业质检里那种“一个部件上有无划痕”不一样,它的形态千奇百怪:裂纹可能是几条细线斜着延伸,擦伤往往是一片不规则的亮斑,剥落则像月球表面的凹坑。传统目标检测输出的矩形框有两个致命问题:

  1. 矩形框会引入大量背景像素,尤其是细长裂纹这种缺陷,框内的有效缺陷像素可能连10%都不到;
  2. 缺陷密集排列时,一个框里可能塞下两三个缺陷,漏检率被变相拉高。

我做实验时对比过用YOLO和UNet在RSDDs上的表现,前者对大面积擦伤效果还凑合,一遇到细裂纹就完蛋:IOU普遍低于0.3,后处理还要靠坐标裁剪把缺陷抠出来,流程绕了一大圈,效果反而更差。分割网络则直接回答“每个像素是不是缺陷”,这就把问题从“物体定位”变成了“像素分类”,恰恰适配轨道缺陷不规则、边界模糊、分布密集的特点。

1.2 UNet在轨道缺陷场景的三个独特优势

有人说DeepLabv3+精度不是更高吗,为什么选UNet?我的理由很简单,三个字:够用、鲁棒、可改。

  • 小样本能训起来。DeepLabv3+这类模型需要充足的训练数据,RSDDs数据集通常只有一百多张到几百张带标注图,即使做数据增强也很难撑起一个深度很大的语义分割模型。UNet参数相对少,结构平滑,没有太多奇怪的trick,几十张训练图也能收敛到可用水平。
  • 跳跃连接天然适合边界分割。轨道缺陷的边界恰恰是评估重点——维修人员关心的是缺陷延伸到了哪里。UNet在解码器每个尺度上都拼接了编码器的低级特征,这让输出mask的边缘能保留更多细节,不会像FCN那样出现大面积糊边。
  • 单独改损失、改输入都很方便。轨道图像通常是灰度图,单通道输入UNet完全没问题;训练数据不足时,也方便把Encoder部分换成预训练权重。如果你想把网络升级成残差块、注意力机制,UNet这种对称结构也非常好改。

下表是我在实际对比中得到的直观感受,供参考:

方案定位框分割边界小样本表现训练成本
Faster R-CNN容易过拟合较高
FCN较粗糙一般
UNet精细
DeepLabv3+精细差,需预训练

2. RSDDs数据集:获取渠道、目录解析与预处理细节

2.1 RSDDs是什么,里面到底有什么

RSDDs是公开的轨道表面缺陷数据集,全称Rail Surface Defect Dataset,很多做轨道视觉检测的论文都用它做benchmark。数据集分成两个子集:I型是真实运营线路上拍到的缺陷,场景复杂,光照不均匀;II型是人工模拟环境下采集的缺陷,背景更干净,缺陷类型更规整。两个子集的分工很清晰:I型负责“现实感”,II型负责“可控性”。

下载下来之后,最核心的东西是两部分:原始灰度/彩色轨道图,以及对应的像素级mask图。mask是二值图,白色像素代表缺陷区域,黑色像素代表正常轨道表面。有些版本还会附缺陷类别标注或者边界框标注,但做UNet训练时,我们只要二值mask就够了。

这里要特别提醒一句:网上流传的RSDDs版本很多,图片数量、尺寸甚至mask格式都不一定一致。常见版本里I型有几十张到上百张不等,II型一百多张到两百张不等,实际以你下到的版本为准。拿到数据后别急着训练,先写脚本把所有图片和标注过一遍,确认数量对得上、mask不是全黑或全白,再开始做预处理,这一步能省下后面排查数据的很多时间。

2.2 获取渠道与目录整理

RSDDs是学术数据集,正规渠道一般是从论文作者公开的项目主页或者相关学术资源站点下载。实际操作中最快的路径是搜索“RSDDs dataset”,大概率能找到一个带有数据集链接或下载脚本的GitHub仓库。很多做轨道缺陷分割的开源项目都会内置一个下载说明,直接跟着走就行。

下载后目录结构可能是这样的:

RSDDs/ ├── RSDDs_I/ │ ├── images/ │ │ ├── 001.png │ │ ├── 002.png │ │ └── ... │ └── masks/ │ ├── 001.png │ ├── 002.png │ └── ... └── RSDDs_II/ ├── images/ └── masks/

需要先检查文件名是否严格对应。我遇到过一批数据,图片叫“001.png”,mask叫“001_label.png”,如果直接按同名匹配就会漏掉全部样本。建议用一个简单的遍历脚本确认原图和mask数量一致,再确认尺寸一致,尺寸不一致的要先统一resize。

这里再强调一个容易被忽略的点:RSDDs的原始图像尺寸通常比较大,有1920x1080甚至更高的,直接整图输入网络会很吃力。一般做法是resize到512x512或256x256,或者先按ROI裁切轨面区域再resize。轨道图像里背景占比很大,直接整图输入会把大量学习能力浪费在道床、碎石这些无关区域上,所以有条件的话最好先裁掉上下两侧的无效区域。

2.3 预处理:尺寸、归一化与图像增强

我在RSDDs上实测下来,输入尺寸512x512是性价比最高的选择。太小(比如128)会丢失裂纹细节,太大(比如1024)会把显存和时间成本翻倍,对于只有一两百张图的数据集来说收益并不明显。

预处理流程我通常是这么写的:

  1. 统一转成灰度图。轨道缺陷本质是和背景的灰度对比变化,RGB图的颜色信息帮助有限,灰度图能减少参数量、加快训练。
  2. 归一化到0到1,或者按ImageNet的均值和方差归一化。二值分割任务里我更推荐直接除以255,简单直接。
  3. mask图resize一定要用最近邻插值,不能用双线性,否则缺陷边界会被插出灰色过渡带,影响二值标签的准确性。
  4. 在线数据增强:随机水平翻转、随机旋转(-15到15度)、随机亮度抖动、随机对比度调整。翻转和旋转是裂纹类缺陷最需要的增强方式,因为裂纹方向在真实场景里是不确定的,但翻转操作不能改变mask的对应关系,所以要用同步随机种子同时对图和mask做变换。

3. UNet结构拆解:编码、解码、跳跃连接各自扮演什么角色

3.1 编码器:特征提取并不神秘

UNet的左边是编码器,也就是一堆卷积和下采样操作的组合。它的作用可以理解成一个多层“特征放大镜”:刚开始只能看到像素级的边缘和灰度变化,越往下采样,感受野越大,能看到的语义信息越丰富——比如这块区域像裂纹、那边是一整块擦伤。

我用的是经典UNet,每个stage包含两个卷积,每个卷积都是3x3、padding=1,后面接BatchNorm和ReLU,然后接一个2x2的最大池化下采样。基础通道数从64开始,每下采样一次通道数翻倍:64 -> 128 -> 256 -> 512 -> 1024。这个设计是有讲究的:下采样会丢失空间分辨率,但通过增加通道数能保留足够的特征表达能力,让网络在“看全局”和“看细节”之间保持平衡。

编码器部分对轨道缺陷来说还有一层特殊意义:裂纹细长、擦伤成片,尺度差异极大。多尺度下采样让网络既能捕捉裂纹这种局部像素级的特征,又不会漏掉大面积的擦伤区域。如果只用原始分辨率一路卷积下去,感受野不够,小裂纹和大擦伤很难同时兼顾。

3.2 解码器与跳跃连接:为什么它决定了分割边界

解码器就是把编码器压缩的特征逐级恢复到原始分辨率的过程。但这里有个关键问题:单纯靠压缩后的特征上采样,输出的边界会很模糊。你想想,一个512x512的图像经过四次下采样变成32x32,这个32x32的特征图里每一点代表原图16x16的一个区域,细节早就不在了,直接上采样恢复出来的边界能准才怪。

UNet的核心设计就在这里:跳跃连接。解码器每一层上采样之后,都会把对应尺度的编码器特征直接拼接过来。比如最底层解码的时候,会把编码器第一次下采样前的64通道特征拼接进来,这样上采样时网络既能看到高层语义信息,又能拿到浅层的高分辨率边界信息,边界定位自然就准了。

这也是UNet结构最有启发性的地方——它不是靠堆深度解决分割问题,而是用结构设计把多尺度信息和边界信息直接融合。在做轨道缺陷这种边界稀疏但极其重要的任务时,这个特性简直是量身定做的。

3.3 网络参数与轻量化调整

如果你只是想在RSDDs上快速跑通,最基础的UNet就够。不过实际使用中有几个值得调的点:

  • 基础通道数可以砍一半。轨道图像是灰度图,特征本身比RGB图少,从32或者48起步就够用。比如把base从64改成32,模型大小直接变成原来的四分之一左右,对小数据集更友好,还不太掉精度。
  • 深度不是越深越好。第四次下采样之后的特征图只有32x32,继续往下采样到16x16甚至8x8,对小块缺陷来说信息损失太大。RSDDs这种尺寸的数据集,四层下采样已经够用。
  • BatchNorm要不要加?我建议加。轨道图像光照不均匀,不同图片整体亮度差异很大,BatchNorm能在每个batch内把数据分布拉平,显著缓解光照干扰。唯一的坑是预测时要记得用训练时累计的均值和方差,PyTorch的eval模式会自动处理,别手动关掉就行。

4. 完整代码落地:数据加载、模型定义、训练与推理

4.1 Dataset的写法:坑点最多的地方

分割任务最容易写错的就是Dataset里的mask处理。很多坑不是网络出的,而是数据加载器和标签对不上导致的。我的写法如下,标注了每个容易出错的地方:

import glob import random import numpy as np import torch from torch.utils.data import Dataset from PIL import Image class RailDefectDataset(Dataset): def __init__(self, img_dir, mask_dir, image_size=(512, 512), train=True): self.img_paths = sorted(glob.glob(img_dir + "/*")) self.mask_paths = sorted(glob.glob(mask_dir + "/*")) self.image_size = image_size self.train = train assert len(self.img_paths) == len(self.mask_paths), \ f"image和mask数量不一致: {len(self.img_paths)} vs {len(self.mask_paths)}" def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img = Image.open(self.img_paths[idx]).convert("L") mask = Image.open(self.mask_paths[idx]).convert("L") img = img.resize(self.image_size, Image.BILINEAR) # 关键:mask不能用双线性,必须用最近邻 mask = mask.resize(self.image_size, Image.NEAREST) if self.train: # 同步翻转 if random.random() > 0.5: img = img.transpose(Image.FLIP_LEFT_RIGHT) mask = mask.transpose(Image.FLIP_LEFT_RIGHT) # 同步旋转 angle = random.uniform(-15, 15) img = img.rotate(angle, resample=Image.BILINEAR) mask = mask.rotate(angle, resample=Image.NEAREST) img = np.array(img).astype(np.float32) / 255.0 mask = np.array(mask).astype(np.float32) / 255.0 # 归一化 img = (img - 0.5) / 0.5 # 输出形状 (1, H, W) / (1, H, W) return torch.from_numpy(img).unsqueeze(0), torch.from_numpy(mask).unsqueeze(0)

这里mask归一化的时候是除以255,所以背景是0、缺陷是1,正好对应二值分割的监督信号。如果你下载的mask不是0/255而是0/1,就不用除以255,直接转float就行。这也是我建议先统计一下mask像素分布的原因,别让这个简单问题浪费一晚上。

4.2 UNet模型定义

模型我用最经典的结构,代码短、好调试、好改。通道数我按上面说的做了一点轻量化,基础通道取32,这样单卡也能训:

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, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size=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=1, num_classes=1, base=32): super().__init__() self.pool = nn.MaxPool2d(2) self.enc1 = DoubleConv(in_channels, base) self.enc2 = DoubleConv(base, base * 2) self.enc3 = DoubleConv(base * 2, base * 4) self.enc4 = DoubleConv(base * 4, base * 8) self.bridge = DoubleConv(base * 8, base * 16) self.up4 = nn.ConvTranspose2d(base * 16, base * 8, kernel_size=2, stride=2) self.dec4 = DoubleConv(base * 16, base * 8) self.up3 = nn.ConvTranspose2d(base * 8, base * 4, kernel_size=2, stride=2) self.dec3 = DoubleConv(base * 8, base * 4) self.up2 = nn.ConvTranspose2d(base * 4, base * 2, kernel_size=2, stride=2) self.dec2 = DoubleConv(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, kernel_size=2, stride=2) self.dec1 = DoubleConv(base * 2, base) self.out_conv = nn.Conv2d(base, 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.bridge(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)

模型输出没有加Sigmoid,所以我训练时直接用BCEWithLogitsLoss,数值稳定性比“Sigmoid+BCELoss”更好。推理时再手动加Sigmoid。

4.3 训练循环与损失函数

训练时我通常用BCE加Dice的混合损失。如果你只用BCE,因为轨道缺陷在整张图里经常只有不到5%的像素,网络很容易学成“全程输出背景”,Dice Loss存在的意义就是缓解这种正负样本极度不平衡的问题。

def dice_loss(pred, target, smooth=1e-6): pred = torch.sigmoid(pred) intersection = (pred * target).sum() union = pred.sum() + target.sum() + smooth return 1 - (2.0 * intersection + smooth) / union def combined_loss(pred, target): bce = nn.functional.binary_cross_entropy_with_logits(pred, target) dice = dice_loss(pred, target) return bce + dice

训练循环本身不复杂,关键在于组织:

device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = UNet().to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50) for epoch in range(50): model.train() train_loss = 0.0 for imgs, masks in train_loader: imgs, masks = imgs.to(device), masks.to(device) preds = model(imgs) loss = combined_loss(preds, masks) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() scheduler.step() print(f"Epoch {epoch + 1}/{50}, Loss: {train_loss / len(train_loader):.4f}")

这里我把学习率设成1e-4,比默认的1e-3要小一些。轨道缺陷小样本场景下,学习率稍低一点不容易训练震荡,后面再用余弦退火慢慢降下去,效果比固定学习率好很多。

4.4 推理与保存结果

训练完做推理时,有个新手很容易踩的坑:训练时图像做了归一化,推理时也要用一模一样的归一化,否则结果会非常奇怪。我的推理片段如下:

model.eval() with torch.no_grad(): img = Image.open(sample_path).convert("L").resize((512, 512), Image.BILINEAR) img = np.array(img).astype(np.float32) / 255.0 img = (img - 0.5) / 0.5 img_tensor = torch.from_numpy(img).unsqueeze(0).unsqueeze(0).to(device) pred = torch.sigmoid(model(img_tensor))[0, 0].cpu().numpy() pred_mask = (pred > 0.5).astype(np.uint8) * 255 # 可选:后处理去掉小连通域 # from scipy import ndimage # label_image, num = ndimage.label(pred_mask) # sizes = ndimage.sum(pred_mask, label_image, range(num + 1)) # pred_mask = np.where(sizes > 100, pred_mask, 0).astype(np.uint8)

阈值0.5算是一个比较稳妥的默认值。如果发现预测出来的缺陷比标注粗一圈,说明阈值可以往上调,比如0.55、0.6;如果发现召回不够,漏检多,阈值就往下降。这个调整比重新训练模型要快得多。

5. 训练策略和实测效果:指标、结果与调参经验

5.1 训练配置与评估指标

在RSDDs这种量级的数据集上,我一般会把数据按7比2比1划分训练、验证、测试。训练集做在线增强,验证集和测试集只做resize和归一化。评估指标用Dice和IoU这两个分割任务标准指标就够了,Dice对“预测和标签的重合度”敏感,IoU更直观。

def iou_score(pred, target, threshold=0.5): pred_bin = (torch.sigmoid(pred) > threshold).float() intersection = (pred_bin * target).sum() union = (pred_bin + target).gt(0).sum() if union.item() == 0: return 1.0 return (intersection / union).item()

我实测的参考结果:基础UNet加BCE加Dice,512x512输入,训练50个epoch,在RSDDs测试集上的Dice大概在0.78到0.86之间,IoU在0.7到0.78左右。这个数字会因为数据划分不同波动,不用太纠结绝对数值,重点观测Dice有没有持续上升,以及验证集IoU有没有随着训练反而下降——那是过拟合信号。

5.2 我实测遇到的三类问题

第一类:缺陷像素占比太少,训练前期loss下降很慢。特别是I型数据集里,裂纹细长,整个512x512图像里缺陷像素可能只有几十个点。我一开始用纯BCE,前十个epoch几乎完全不学。换成BCE加Dice之后,效果立竿见影,这说明class imbalance问题严重时,不能只靠单一loss顶着。

第二类:小样本过拟合来得非常快。训练到第30个epoch时验证IoU往往就停滞了,而训练集loss还在持续降低。解决办法我总结成三句:增强强度不要过猛、epoch不需要太长、weight_decay给一点。我用AdamW加1e-5的weight_decay,配合早停,稳定了很多。

第三类:mask边缘粗糙,有小孔洞和碎点。分割网络天然会在边界产生软输出,不是0就是1的硬判决很容易留下毛刺。我的经验是不要一上来就上复杂的后处理,先用形态学开闭运算把小噪点清掉,再通过连通域面积阈值把零散的小预测块去掉。如果发现这些后处理对IoU提升很大,再回头审视一下是不是增强时mask旋转用错了插值算法,或者用了双线性导致mask变糊。

5.3 几个值得收藏的调参细节

  • BatchSize不要太小。轨道图resize到512后单张显存占用不小,但BatchSize如果只有2,BatchNorm的统计量会非常不稳定,训练很难收敛。我通常设8到16,如果显存不够,优先把base通道数调小,而不是牺牲batch size。
  • 用随机裁剪代替直接resize。如果显存和训练时间允许,随机裁剪256x256或者512x512的patch来训练比整图resize效果更好,相当于给网络看了更多局部细节,而且等于多做了一种数据增强。
  • 监控可视化而不是只盯loss。每周n次训练中把预测mask和原图叠在一起看一眼,比看一百轮loss曲线更容易发现问题。我第一次发现mask边界偏移,就是从可视化里看出来,边界总往上一侧偏,后来排查发现是mask和原图resize时的人工对齐问题,训练本身根本没毛病。

说到最后,给你一个我自己的土办法:训练之前先拿一两张图跑到过拟合,也就是把某个batch反复训练几十轮,看看模型能不能把这一批图完全背下来。如果连过拟合都过不了,大概率是代码有bug——要么是mask加载错了,要么是网络输入输出尺寸不匹配,要么是优化器配置不对。这个排查思路能帮你把“代码问题”和“算法问题”分开定位,在实际做UNet轨道缺陷检测时非常省时间。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询