简介这份铁路轨道故障检测图像分类数据集面向计算机视觉学习者、轨道交通智能运维研究者及需要开展缺陷识别实验的开发者用于解决轨道状态二分类任务中样本获取与标注成本高的问题。数据已按正常与故障两类完成标注并预先划分训练集、验证集与测试集各类图片分目录存放便于直接接入CNN分类网络训练与评估。压缩包共803个文件以jpg为主779个另含少量jpeg、webp图像以及1个json标注文件和1个Python可视化脚本整体约278.41MB运行show脚本即可快速浏览样本分布与图像内容。目前已有212人学习下载。借助该数据集读者可省去繁琐的采集与标注环节直接复现分类流程、验证网络改进效果并结合json文件核对类别信息为轨道故障检测模型的调参与对比实验提供可靠的数据基础。1. 铁路轨道故障检测图像分类数据集约800张已标注数据能跑出什么结果铁路轨道巡检的痛点是场景重复、缺陷样本稀少、标注成本高。你手头如果只有几百张现场图想验证一个图像分类模型能不能区分轨道扣件缺失、裂纹、异物侵入这几类故障最直接的办法就是找一份已标注的小规模数据集先跑通全流程。铁路轨道故障检测图像分类数据集约800张已标注图像正好卡在“够用但不富裕”的区间它足够让你完成一次端到端的训练与验证又不至于让你在数据清洗上耗掉一周。这类数据集通常按故障类别分目录存放标注形式以文件夹名或CSV标签为主适合用迁移学习快速出baseline。适合谁用一是做轨道交通智能运维的算法工程师需要快速验证模型选型二是高校做视觉检测课题的研究生缺真实场景数据三是巡检机器人或边缘盒子的开发者想评估模型在有限样本下的泛化边界。800张不是让你刷SOTA的是让你把“数据→训练→部署→现场误报”这条链路先跑通再决定要不要扩标。2. 拿到800张标注图先别急着训练数据体检与划分策略2.1 先做三件事去重、查类别平衡、看图像质量约800张的数据集第一反应不应该是train_test_split而是先体检。常见做法是写一个脚本统计每个类别的文件数、图像尺寸分布、以及用感知哈希找近似重复图。铁路巡检图有个特点连续帧之间高度相似如果随机划分训练集和验证集会出现几乎一样的图验证准确率虚高到95%以上上线就翻车。我一般会先算pHash把汉明距离小于5的图归为一组再按组划分保证同一段轨道的图不会同时出现在训练和验证里。import os import imagehash from PIL import Image from collections import defaultdict def audit_dataset(root): stats defaultdict(list) hashes {} for cls in os.listdir(root): cls_dir os.path.join(root, cls) if not os.path.isdir(cls_dir): continue for fname in os.listdir(cls_dir): fpath os.path.join(cls_dir, fname) try: img Image.open(fpath).convert(RGB) except Exception as e: print(f损坏文件: {fpath}, {e}) continue # 感知哈希用于找近似重复 ph imagehash.phash(img) hashes[fpath] (cls, ph) stats[cls].append(img.size) # 打印类别数量与尺寸分布 for cls, sizes in stats.items(): w [s[0] for s in sizes] h [s[1] for s in sizes] print(f{cls}: {len(sizes)}张, 宽{min(w)}-{max(w)}, 高{min(h)}-{max(h)}) return hashes hashes audit_dataset(./rail_dataset)这段代码做三件事遍历类别目录、用pHash记录每张图的指纹、输出类别数量和尺寸范围。参数上imagehash.phash的hash_size默认8对轨道图足够敏感如果发现某类只有30张而另一类有300张后面训练必须用加权采样或focal loss否则模型会偏向多数类。尺寸范围如果跨度大比如从640×480到1920×1080不要直接resize到224先按短边缩放再中心裁剪保留缺陷纹理。2.2 划分比例不是7:2:1小样本要用分层分组800张按类别分假设4类每类200张。如果按7:2:1随机分验证集只有160张每类40张指标波动极大。更稳的做法是分层分组划分先按pHash分组再在每组内按类别分层抽样。常见比例是训练70%、验证15%、测试15%但测试集要留出“不同轨道区段”的图用来评估跨区段泛化。我一般会写一个GroupShuffleSplitgroup用pHash的聚类ID。import numpy as np from sklearn.model_selection import GroupShuffleSplit def split_by_group(hashes, test_size0.15, val_size0.15): paths list(hashes.keys()) labels [hashes[p][0] for p in paths] # 用pHash的字符串作为分组ID近似图归为一组 groups [str(hashes[p][1]) for p in paths] gss GroupShuffleSplit(n_splits1, test_sizetest_size, random_state42) train_val_idx, test_idx next(gss.split(paths, labels, groups)) # 再从train_val里分验证集 train_val_paths [paths[i] for i in train_val_idx] train_val_labels [labels[i] for i in train_val_idx] train_val_groups [groups[i] for i in train_val_idx] gss2 GroupShuffleSplit(n_splits1, test_sizeval_size/(1-test_size), random_state42) train_idx, val_idx next(gss2.split(train_val_paths, train_val_labels, train_val_groups)) return train_idx, val_idx, test_idx关键参数是groups它保证同一组近似图只出现在一个split里。random_state固定后结果可复现。如果分组后训练集少于500张说明重复图太多需要先做去重把冗余图删掉再分。这一步不做后面调参全是玄学。2.3 标注格式转换从文件夹到CSV再到DataLoader数据集标注常见两种按类别文件夹存放或一个CSV带filename和label。文件夹结构最省事但训练时需要生成索引。我一般先转成统一CSV再写Dataset类。转换脚本要处理路径分隔符和类别映射避免Windows和Linux混用出错。import pandas as pd import os def folder_to_csv(root, out_csvlabels.csv): records [] classes sorted([d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]) cls_to_idx {c: i for i, c in enumerate(classes)} for cls in classes: cls_dir os.path.join(root, cls) for fname in os.listdir(cls_dir): if fname.lower().endswith((.jpg, .jpeg, .png, .bmp)): records.append({ filename: os.path.join(cls, fname), label: cls_to_idx[cls], class_name: cls }) df pd.DataFrame(records) df.to_csv(out_csv, indexFalse) print(f共{len(df)}条, 类别映射: {cls_to_idx}) return df df folder_to_csv(./rail_dataset)输出CSV后用torch.utils.data.Dataset或tf.data读取。注意filename用相对路径换机器时只改root。类别映射要保存成json推理时顺序不能乱。如果数据集自带标注文件先检查有没有漏标、错标尤其是“正常”类里混入缺陷图这种脏数据会让模型学偏。3. 用迁移学习在800张图上跑通baseline模型选型与训练参数3.1 选ResNet18还是ViT小样本下的真实差距800张图4类每类200张。这个规模下从零训练CNN基本没戏必须迁移学习。常见选择是ResNet18、EfficientNet-B0、MobileNetV3以及最近的ViT-B/16。我的血泪经验是ResNet18在800张上微调验证准确率能到85%92%ViT-B/16如果直接微调容易过拟合到70%出头需要更强的数据增强和更低的学习率。如果追求边缘部署MobileNetV3是首选精度掉23个点但推理快3倍。选型理由很简单数据量小归纳偏置强的CNN更稳ViT需要至少几千张才能体现优势。如果非要用Transformer建议用预训练的Swin-Tiny或DeiT-Small并且冻结前几个stage。3.2 训练脚本冻结骨干、分层学习率、早停下面是一个PyTorch训练脚本的核心部分用ResNet18做迁移学习。关键点冻结layer1和layer2只训练layer3、layer4和全连接层分层学习率骨干用1e-4分类头用1e-3早停patience7。import torch import torch.nn as nn import torchvision.models as models from torch.utils.data import DataLoader from torchvision import transforms def build_model(num_classes4): model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 冻结浅层 for name, param in model.named_parameters(): if name.startswith(layer1) or name.startswith(layer2): param.requires_grad False model.fc nn.Linear(model.fc.in_features, num_classes) return model def train_one_epoch(model, loader, criterion, optimizer, device): model.train() total_loss, correct, total 0, 0, 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * imgs.size(0) correct (outputs.argmax(1) labels).sum().item() total imgs.size(0) return total_loss / total, correct / total # 数据增强轨道图需要颜色抖动和随机擦除 train_tf transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.2, 0.2, 0.2), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])参数说明weightsIMAGENET1K_V1加载ImageNet预训练权重不要用pretrainedTrue的旧写法。RandomCrop(224)配合Resize(256)是标准做法。ColorJitter对轨道图有用因为现场光照变化大。RandomHorizontalFlip要谨慎如果缺陷有方向性比如左右扣件不同翻转可能改变类别语义建议先验证。优化器用AdamWweight_decay1e-4学习率用param_groups分层设置。3.3 训练时看什么损失曲线、混淆矩阵、误报样本800张数据训练20个epoch通常够。但不要只看准确率。我一般每5个epoch画一次混淆矩阵重点看“正常”和“扣件缺失”之间的误判。如果正常类被大量判成缺陷说明阈值偏了需要调整类别权重。另一个关键是保存误报样本人工看一遍很多是标注错误或图像模糊。验证集准确率到90%后提升就靠数据清洗和难例挖掘不是调参。from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt def evaluate(model, loader, device, classes): model.eval() all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in loader: imgs imgs.to(device) outputs model(imgs) preds outputs.argmax(1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) sns.heatmap(cm, annotTrue, fmtd, xticklabelsclasses, yticklabelsclasses) plt.savefig(confusion_matrix.png) return cm混淆矩阵保存后对照看哪一类召回低。如果“裂纹”召回只有60%去训练集里找裂纹图看是不是分辨率太低或标注框不准。分类任务里标注质量比模型结构重要得多。4. 避坑与排查800张轨道图训练中最容易翻车的5个点4.1 验证集准确率95%上线后误报一片现象本地验证准确率很高部署到巡检设备后正常轨道被频繁判成故障。原因随机划分导致训练集和验证集有近似图模型记住了背景而不是缺陷。解决按pHash分组划分测试集用不同区段的图上线前用一段未标注视频抽帧做盲测看误报率。4.2 训练loss震荡验证loss上升现象训练loss忽高忽低验证loss在第3个epoch后持续上升。原因学习率太大或者batch size太小比如8BatchNorm统计量不稳定。解决骨干学习率降到1e-4分类头1e-3batch size至少16不够就累积梯度加weight_decay1e-4。4.3 某类样本只有30张模型完全不学现象四类中“异物侵入”只有30张训练后该类召回为0。原因类别不平衡交叉熵损失被多数类主导。解决用加权采样WeightedRandomSampler权重取类别频率的倒数或改用focal lossgamma2。同时对该类做更强增强比如旋转、裁剪、加噪声。4.4 图像尺寸不统一resize后缺陷消失现象原图1920×1080resize到224×224后细小裂纹变成一条模糊线模型分不出来。原因直接resize压缩了高频信息。解决先按短边缩放到256再中心裁剪224或者用滑动窗口切patch每个patch单独分类最后投票。轨道图建议切640×640的patch保留纹理。4.5 标注文件里类别名有空格或大小写不一致现象folder_to_csv生成的类别映射出现Fault和fault两个类实际是一类。原因文件夹命名不规范。解决转换前统一转小写、去空格用os.listdir后先打印所有目录名人工确认。标注工具导出的CSV也要检查label列有没有多余空格。5. 从800张到可部署模型难例挖掘与置信度校准技巧800张数据训完baseline下一步不是急着扩标而是做难例挖掘。具体做法用训练好的模型在验证集和一段未标注视频上推理把置信度在0.40.6之间的样本挑出来人工复核。这些样本通常是光照突变、雨雪遮挡、油污干扰正是现场误报的主要来源。我一般会挑出50100张难例修正标注后加入训练集再微调5个epoch验证集准确率通常能再涨35个点。另一个技巧是置信度校准分类模型输出的softmax概率往往偏高用温度缩放temperature scaling在验证集上拟合一个温度参数T让概率更接近真实置信度。部署时设阈值0.7低于阈值的图转人工复核能大幅降低误报。import torch import torch.nn.functional as F from scipy.optimize import minimize_scalar def calibrate_temperature(model, val_loader, device): model.eval() logits_list, labels_list [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs imgs.to(device) logits model(imgs) logits_list.append(logits.cpu()) labels_list.append(labels) logits torch.cat(logits_list) labels torch.cat(labels_list) nll_criterion torch.nn.CrossEntropyLoss() def eval_temp(T): loss nll_criterion(logits / T, labels) return loss.item() res minimize_scalar(eval_temp, bounds(0.5, 3.0), methodbounded) return res.x # 使用T calibrate_temperature(model, val_loader, device) # 推理时probs F.softmax(logits / T, dim1)温度T一般在1.22.0之间。校准后模型说“90%是故障”时实际准确率更接近90%而不是盲目自信。这个技巧在边缘设备上尤其有用因为误报一次就要派人去现场成本很高。最后说一个习惯每次训练完把配置文件、类别映射、最佳权重、混淆矩阵存到一个带日期的文件夹里别问为什么问就是后悔药。希望帮到你。本文还有配套的精品资源点击获取