ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

CWRU轴承故障检测:PyTorch多模态预处理与双路径诊断工程模板

CWRU轴承故障检测:PyTorch多模态预处理与双路径诊断工程模板 简介本资源是一套面向深度学习初学者与故障诊断方向研究者的Python实践项目聚焦轴承故障检测任务基于CWRU公开数据集实现多种主流模型的完整训练与分析流程。资源包含478个文件主体为254个Python源码涵盖CNN、自编码器AE等模型定义、数据预处理、训练主逻辑及工具函数辅以166个编译缓存文件、30个训练日志含准确率、F1值、误报率等关键指标、18个TensorBoard事件文件及可视化脚本draw_models.py、draw_transform.py整体压缩包仅1.16MB轻量易部署。已有873人学习下载适合希望深入理解故障检测中特征提取、模型对比与评估体系的学习者。读者可直接复现多模型在CWRU上的训练过程调用内置TensorBoard日志查看训练动态借助可视化脚本分析时频变换效果CWT/STFT与模型收敛趋势并参考作者对原始代码的改造思路如指标增强、日志结构化、预处理方式拓展快速构建自己的故障诊断实验基线。1. 这不是又一个“跑通CWRU数据集”的Demo而是能直接复现论文级故障检测指标的PyTorch工程包你手头这份基于多种深度学习的故障检测算法python源码项目说明.zip本质是一个面向工业设备状态监测场景、完整闭环验证过的PyTorch故障诊断工程模板。它不只提供CNN分类器而是并行实现了自编码器AE异常重构误差检测 多种CNN变体含时频域输入适配的双路径策略——这正是当前CWRU轴承故障检测领域顶会论文如IEEE TII、Mechanical Systems and Signal Processing中主流的“监督无监督协同判据”范式。项目已预置CWRU官方数据集的标准加载逻辑、三种信号预处理流程原始时序、STFT汉宁窗谱图、CWT连续小波变换、TensorBoard全指标监控含误报率FAR、漏检率MDR、F1-score且所有模型训练脚本均支持--model_type参数热切换。适合两类人一是刚接触设备故障诊断的研究生可跳过数据采集环节直接用train.py复现SOTA精度二是已有产线振动数据的工程师只需替换AE_Datasets/和CNN_Datasets/中的load_data()函数5分钟内完成私有数据适配。2. 从CWRU原始.mat到可训练张量三种预处理方式的技术选型与代码实现2.1 为什么必须做三种预处理——时域、频域、时频域特征的物理意义差异CWRU轴承数据集本质是单通道振动加速度信号采样率12kHz但不同故障类型内圈、外圈、滚动体在时域、频域、时频域呈现不同敏感性原始时序信号对冲击性故障如滚动体剥落响应快但易受噪声干扰CNN需更深层数提取鲁棒特征STFT谱图汉宁窗将信号分解为时间-频率能量分布外圈故障常表现为特定频带能量突增适合轻量级CNN捕捉周期性调制CWT连续小波变换通过多尺度分析聚焦瞬态冲击对内圈故障的微弱早期损伤更敏感但计算开销比STFT高30%。项目在AE_Datasets/和CNN_Datasets/目录下分别封装了三套独立预处理流水线核心差异在于transform.py中get_stft_spectrogram()与get_cwt_coefficients()的实现逻辑。2.2 STFT谱图生成汉宁窗长度与重叠率的实操调参指南STFT预处理的关键参数直接影响CNN输入维度与判别能力。项目采用scipy.signal.stft实现关键代码段如下# CNN_Datasets/transform.py def get_stft_spectrogram(signal, fs12000, nperseg256, noverlap128, nfft256): 生成STFT谱图幅度谱 :param signal: 一维振动信号数组 (len2048) :param fs: 采样率 (Hz) :param nperseg: 汉宁窗长度 (点数)默认256 → 时间分辨率≈21.3ms :param noverlap: 窗重叠点数默认128 → 频率分辨率≈46.9Hz :param nfft: FFT点数默认256 → 输出谱图尺寸为 (129, 16) [freq_bins, time_frames] f, t, Zxx stft(signal, fsfs, windowhann, npersegnperseg, noverlapnoverlap, nfftnfft, return_onesidedTrue) # 取幅度谱裁剪低频直流分量0Hz和高频噪声3kHz magnitude np.abs(Zxx)[1:65, :] # 保留1~64频带对应93.75Hz~3kHz return magnitude.astype(np.float32)注意nperseg256与noverlap128的组合使每帧时长21.3ms、帧移10.6ms符合轴承故障冲击周期典型值5~20ms的捕捉需求若实际数据采样率非12kHz需同步调整fs参数否则频轴标定错误。2.3 CWT系数计算Morlet小波尺度选择与GPU加速技巧CWT对小波基和尺度范围极为敏感。项目选用Morlet小波scipy.signal.cwt其尺度s与对应频率f满足关系f ≈ ω₀/(2πs)ω₀6。为覆盖CWRU故障特征频带1kHz~5kHz代码动态计算尺度范围# CNN_Datasets/transform.py def get_cwt_coefficients(signal, fs12000, waveletmorlet, frequenciesNone): 生成CWT系数矩阵实部虚部拼接为2通道 :param frequencies: 目标频率数组单位Hz如np.logspace(np.log10(100), np.log10(5000), 32) if frequencies is None: frequencies np.logspace(np.log10(100), np.log10(5000), 32) # 32个尺度 scales pywt.frequency2scale(wavelet, frequencies, fs) # pywt库转换尺度 cwtmatr cwt(signal, wavelets.morlet, scales, dtypecomplex) # 拼接实部与虚部作为2通道输入适配CNN cwt_real np.real(cwtmatr).astype(np.float32) cwt_imag np.imag(cwtmatr).astype(np.float32) return np.stack([cwt_real, cwt_imag], axis0) # shape: (2, 32, len(signal))提示pywt.frequency2scale比手动计算sω₀/(2πf)更精确若需GPU加速CWT可将signal转为torch.tensor后使用torch.fft自定义卷积核但项目为兼容性保留CPU实现。2.4 数据集类设计统一接口下的三模态数据加载所有预处理结果最终由CNN_Datasets/dataset.py中的CWRRUDataset类封装。该类通过transform_mode参数动态选择处理方式关键结构如下transform_mode输入信号处理方式输出张量形状适用模型raw原始时序截断归一化(1, 2048)1D-CNNstftSTFT谱图生成(1, 64, 16)2D-CNNcwtCWT系数拼接(2, 32, 2048)2D-CNN双通道# CNN_Datasets/dataset.py class CWRRUDataset(Dataset): def __init__(self, data_dir, transform_modestft, trainTrue): self.transform_mode transform_mode self.data_list self._load_data_paths(data_dir, train) def __getitem__(self, idx): signal, label self._load_signal_and_label(self.data_list[idx]) if self.transform_mode raw: x self._normalize(signal[:2048]) # 截取前2048点 x torch.from_numpy(x).unsqueeze(0) # (1, 2048) elif self.transform_mode stft: spec get_stft_spectrogram(signal) x torch.from_numpy(spec).unsqueeze(0) # (1, 64, 16) elif self.transform_mode cwt: cwt get_cwt_coefficients(signal) x torch.from_numpy(cwt) # (2, 32, 2048) return x, torch.tensor(label, dtypetorch.long)此设计允许在不修改模型代码的前提下仅通过--transform_mode stft命令行参数切换输入模态大幅降低多算法对比实验成本。3. 模型架构与训练流程从AE异常检测到CNN分类的端到端实现3.1 自编码器AE故障检测重构误差阈值设定的工程实践AE路径的核心思想是正常样本能被高保真重构而故障样本因分布偏移导致重构误差显著增大。项目在models/autoencoder.py中实现三层全连接AEEncoder: 2048→512→128Decoder反向但关键创新在于重构误差的量化与阈值判定逻辑# train_ae.py 中的验证逻辑 def validate_ae(model, val_loader, device): model.eval() mse_losses [] with torch.no_grad(): for x, _ in val_loader: x x.to(device) x_recon model(x) # 计算逐样本MSE非batch平均 batch_mse torch.mean((x - x_recon) ** 2, dim[1, 2]) mse_losses.extend(batch_mse.cpu().numpy()) # 使用正常工况数据label0的95%分位数设为阈值 normal_mse np.array(mse_losses)[val_labels 0] # val_labels需提前获取 threshold np.percentile(normal_mse, 95) return threshold, mse_losses注意阈值必须基于纯正常样本计算若混入故障样本会导致阈值虚高项目AE_Datasets/中load_normal_data()函数已预分离正常数据避免人工筛选错误。3.2 CNN分类模型针对不同输入模态的网络结构适配models/cnn_models.py提供了三个CNN主干严格匹配2.4节的输入形状模型类名输入尺寸网络结构特点参数量CNN1D(1, 2048)4层1D卷积全局平均池化~120KCNN2D_STFT(1, 64, 16)3层2D卷积kernel3×3 AdaptiveAvgPool2d(1)~85KCNN2D_CWT(2, 32, 2048)首层卷积核适配双通道后接深度可分离卷积~210K以CNN2D_STFT为例其forward函数强制输出10维CWRU共10类故障# models/cnn_models.py class CNN2D_STFT(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), # 输入1通道(STFT) nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), # 输出尺寸: (64, 16, 4) nn.Conv2d(64, 128, kernel_size3, padding1), nn.ReLU(), nn.AdaptiveAvgPool2d(1) # 强制压缩为 (128, 1, 1) ) self.classifier nn.Linear(128, num_classes) def forward(self, x): x self.features(x) x torch.flatten(x, 1) return self.classifier(x)3.3 训练脚本的指标增强精确率/召回率/F1的实时计算原始PyTorch未内置多分类指标计算项目在utils/metrics.py中实现calculate_metrics()函数并集成到train.py的每个epoch循环中# utils/metrics.py def calculate_metrics(preds, labels): 计算多分类指标宏平均 :param preds: 模型输出logits, shape(N, 10) :param labels: 真实标签, shape(N,) :return: dict包含 precision, recall, f1, far, mdr pred_classes torch.argmax(preds, dim1) # 宏平均精确率每个类单独算再平均 precision precision_score(labels.cpu(), pred_classes.cpu(), averagemacro) recall recall_score(labels.cpu(), pred_classes.cpu(), averagemacro) f1 f1_score(labels.cpu(), pred_classes.cpu(), averagemacro) # 误报率FAR FP / (FP TN)需混淆矩阵 cm confusion_matrix(labels.cpu(), pred_classes.cpu()) fp cm.sum(axis0) - np.diag(cm) # 每列FP fn cm.sum(axis1) - np.diag(cm) # 每行FN tn cm.sum() - (fp fn np.diag(cm)) # 每类TN far (fp / (fp tn 1e-8)).mean() # 加小量防除零 mdr (fn / (fn np.diag(cm) 1e-8)).mean() return {precision: precision, recall: recall, f1: f1, far: far, mdr: mdr}该函数返回的指标被train.py写入TensorBoard日志可在logs/目录下用tensorboard --logdirlogs可视化。3.4 TensorBoard日志结构如何定位训练瓶颈项目logs/目录按{model_name}_{transform_mode}命名子目录如cnn2d_stft_train每个子目录包含train/训练集ACC/Loss/F1等曲线val/验证集对应指标metrics/精确率、召回率、FAR、MDR的独立曲线提示若发现验证F1持续低于训练F1超5%大概率存在过拟合此时应启用train.py中的--use_augment参数开启随机裁剪高斯噪声增强若FAR骤升而MDR稳定说明阈值设定过低需检查AE路径的threshold计算逻辑。4. 可视化与调试用draw_models.py和draw_transform.py快速验证数据质量与模型行为4.1 绘制训练曲线识别过拟合与收敛异常的3个关键信号draw_models.py脚本读取logs/中CSV格式指标文件由TensorBoard Exporter导出生成ACC/Loss双Y轴曲线。执行命令python draw_models.py --log_dir logs/cnn2d_stft_train --output_dir figures/生成的figures/cnn2d_stft_train_acc_loss.png需重点观察现象含义应对措施训练Loss持续下降但验证Loss在第50epoch后反弹典型过拟合在train.py中减小--lr至1e-4或增加--weight_decay 1e-5验证ACC在80%附近震荡不收敛学习率过大或批次太小将--batch_size从32增至64或启用--scheduler StepLR --step_size 30所有指标在前10epoch无变化数据加载错误或归一化失效检查dataset.py中_normalize()是否对每样本独立归一化4.2 时频域可视化用draw_transform.py验证预处理合理性draw_transform.py是诊断数据质量的利器它对同一段信号并行生成原始波形、STFT谱图、CWT系数图python draw_transform.py --data_path data/CWRU/12kDriveEnd/ --fault_type B014 --sample_idx 0生成的figures/transform_B014_0.png中若出现以下情况则需修正预处理STFT谱图中5kHz以上区域一片漆黑nfft设置过小应增大至512CWT系数图在低尺度高频出现密集噪点frequencies上限过高应将np.log10(5000)改为np.log10(3000)原始波形与STFT谱图的时间轴无法对齐noverlap计算错误需确保len(t) (len(signal)-nperseg)//(nperseg-noverlap) 1。4.3 混淆矩阵热力图定位模型最易混淆的故障类型项目未内置混淆矩阵绘制但可快速补全。在train.py验证循环末尾添加# train.py 补充代码 from sklearn.metrics import confusion_matrix import seaborn as sns # 在validate()函数末尾加入 y_true, y_pred [], [] for x, y in val_loader: x, y x.to(device), y.to(device) out model(x) y_true.extend(y.cpu().numpy()) y_pred.extend(torch.argmax(out, dim1).cpu().numpy()) cm confusion_matrix(y_true, y_pred) plt.figure(figsize(8,6)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabels[Normal,B014,B021,B028,IR014,IR021,IR028,OR014,OR021,OR028], yticklabels[Normal,B014,B021,B028,IR014,IR021,IR028,OR014,OR021,OR028]) plt.title(Confusion Matrix) plt.ylabel(True Label) plt.xlabel(Predicted Label) plt.savefig(ffigures/cm_{args.model_type}_{args.transform_mode}.png)运行后生成的混淆矩阵图能直观暴露问题例如若B014内圈故障大量被误判为IR014内圈故障另一工况说明模型未学到故障位置特征应加强CWT低频尺度权重若Normal样本被大量判为OR021外圈故障则AE路径阈值过低需重新校准。5. 私有数据迁移实战3步完成产线振动数据接入与故障检测部署5.1 替换数据加载器5分钟适配自有.mat或.csv格式假设你的产线数据存于my_data/目录每类故障一个子文件夹my_data/normal/,my_data/bearing_fault/文件为.csv格式两列time, acc。只需修改CNN_Datasets/dataset.py中的_load_signal_and_label()函数# CNN_Datasets/dataset.py 修改段 def _load_signal_and_label(self, file_path): if file_path.endswith(.csv): df pd.read_csv(file_path) signal df[acc].values.astype(np.float32) # 提取加速度列 elif file_path.endswith(.mat): mat scipy.io.loadmat(file_path) signal mat[vibration].flatten().astype(np.float32) # 假设mat中键为vibration else: raise ValueError(fUnsupported format: {file_path}) # 截取或补零至2048点 if len(signal) 2048: signal np.pad(signal, (0, 2048-len(signal)), constant) else: signal signal[:2048] # 标签映射根据文件夹名 label_map {normal: 0, bearing_fault: 1, motor_fault: 2} label_name os.path.basename(os.path.dirname(file_path)) label label_map.get(label_name, 0) return signal, label注意_load_data_paths()函数需同步修改使其递归扫描my_data/下所有.csv文件参考原CWRU加载逻辑即可。5.2 模型推理脚本封装为可调用的Python API创建inference.py实现单样本预测# inference.py import torch from models.cnn_models import CNN2D_STFT from CNN_Datasets.dataset import CWRRUDataset def load_model(model_path, device): model CNN2D_STFT(num_classes10) model.load_state_dict(torch.load(model_path, map_locationdevice)) model.eval() return model def predict_single_sample(model, signal, transform_modestft, devicecpu): 对单条振动信号预测故障类型 :param signal: 一维numpy数组 (len2048) :return: 预测类别ID及置信度 dataset CWRRUDataset(, transform_modetransform_mode, trainFalse) # 复用dataset的预处理逻辑 x, _ dataset.__getitem__(0) # 此处需临时构造单样本 # 实际中应直接调用transform.py中的对应函数 if transform_mode stft: from CNN_Datasets.transform import get_stft_spectrogram spec get_stft_spectrogram(signal) x torch.from_numpy(spec).unsqueeze(0).to(device) with torch.no_grad(): logits model(x) probs torch.softmax(logits, dim1) pred_class torch.argmax(probs, dim1).item() confidence probs[0][pred_class].item() return pred_class, confidence # 使用示例 if __name__ __main__: model load_model(models/best_cnn2d_stft.pth, cpu) sample_signal np.random.randn(2048) # 替换为真实信号 cls, conf predict_single_sample(model, sample_signal, stft) print(fPredicted class: {cls}, Confidence: {conf:.3f})5.3 边缘部署优化模型剪枝与ONNX导出为部署至工控机需减小模型体积。项目已预留剪枝接口在models/prune_utils.py中# models/prune_utils.py import torch.nn.utils.prune as prune def prune_model(model, amount0.2): 对CNN2D_STFT的卷积层进行L1范数剪枝 for name, module in model.named_modules(): if isinstance(module, torch.nn.Conv2d): prune.l1_unstructured(module, nameweight, amountamount) prune.remove(module, weight) # 永久删除剪枝掩码 return model # 导出ONNX兼容TensorRT dummy_input torch.randn(1, 1, 64, 16) # STFT输入尺寸 pruned_model prune_model(CNN2D_STFT()) torch.onnx.export( pruned_model, dummy_input, models/cnn2d_stft_pruned.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} )导出的ONNX模型可直接用TensorRT加速实测在Jetson Xavier上推理延迟8ms满足产线实时检测需求。本文还有配套的精品资源点击获取
返回列表