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

发布时间:2026/9/15 7:51:40

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满足产线实时检测需求。本文还有配套的精品资源点击获取
延伸阅读

更多相关文章

2026/9/15 7:51:40

如何挖掘与分析无标题项目的技术价值

1. 项目概述作为一名从业多年的技术博主,我经常遇到这样的情况:一个看似简单的项目标题背后,往往隐藏着丰富的技术内涵和实践价值。今天我想和大家聊聊,当我们面对一个"无标题"项目时,应该如何挖掘其潜在价值…

2026/9/15 7:51:40

强化学习基础:从马尔可夫决策到深度Q网络

1. 强化学习基础概念解析强化学习(Reinforcement Learning)作为机器学习三大范式之一,与监督学习、无监督学习有着本质区别。它的核心在于让智能体(Agent)通过试错机制与环境(Environment)持续交…

2026/9/15 7:46:39

山东企业AI转型实战:场景落地与政策红利解析

1. 山东企业AI转型的时代背景与政策红利2023年被称为"AI应用元年",山东省工业和信息化厅最新数据显示,全省已有47%的规上工业企业启动AI应用场景建设。在《山东省"十四五"数字强省建设规划》中,人工智能被列为重点突破的…

2026/9/15 8:01:40

Agent记忆管理实战:用总结压缩与Milvus搭建长期记忆

1. 先聊一个扎心场景:Agent 又“失忆”了做 Agent 开发的朋友大概率都撞过这堵墙:对话轮次一多,模型突然开始胡说八道,或者干脆报错,提示说超出了 token 上限,回答被截断。更难受的是,明明用户十…

2026/9/15 8:01:40

Claude API定价模型解析与成本优化实战

1. Claude API 定价模型深度解析作为AI领域从业者,我最近花了大量时间研究Claude API的定价机制。与市面上其他大模型API不同,Anthropic采用了基于"每百万token"的阶梯式计费模式。这种设计既考虑了不同规模用户的使用需求,又能有效…

2026/9/15 8:01:40

「AI Agent 全栈开发 50 讲」——从本地模型部署到多智能体系统,一年省 87 万 第 44 课 | 评估框架:量化 Agent 的真实能力

第 44 课 | 评估框架:量化 Agent 的真实能力没有评估就没有优化。构建科学的评估体系——成功率、步数、Token、耗时、准确率,让 Agent 能力可量化、可追踪、可优化。一、业务痛点:Agent 的「黑盒」困境 在之前的课程中,我们已经构…

2026/9/15 7:56:40

DeFi安全与信任机制:Aave治理危机的技术解析

1. 项目背景与核心矛盾解析"Aave 冬日乱局"这个标题生动揭示了DeFi领域一个经典困境:技术层面的安全防护与社区信任建设的失衡。作为头部借贷协议,Aave确实通过智能合约构建了堪称行业标杆的资金防护体系——其多重签名机制、风险参数动态调整…

2026/9/15 4:54:30

拯救者Y7000黑屏故障排查与维修实战指南

1. 项目概述:一台黑屏的拯救者Y7000,到底卡在哪一步? 联想拯救者Y7000系列笔记本,从2018年第一代搭载i5-8300H开始,到后来的i7-9750H、i7-10750H、i5-11400H,再到2023年款的R7-7840HS,它始终是学…

2026/9/15 0:01:16

AI英语单词APP开发:自适应学习算法与移动端优化实践

1. 项目概述 作为一名在移动应用开发领域摸爬滚打多年的老手,我最近完成了一个AI英语单词APP的开发项目。这个项目将传统单词记忆方法与现代AI技术相结合,打造了一款能够智能适应不同用户学习习惯的英语学习工具。 市面上大多数单词APP都存在一个通病&a…

2026/9/15 0:01:16

Flutter与OpenHarmony结合开发手语学习APP实战

1. 项目背景与核心价值作为一名同时接触过Flutter和OpenHarmony的开发者,最近我完成了一个基于Flutter for OpenHarmony的手语学习APP实战项目。这个项目最大的特点在于实现了跨平台框架与国产操作系统深度结合的创新实践——用Flutter开发的应用能完美运行在OpenHa…

2026/9/15 0:01:16

六个月成为机器人工程师:从ROS2到SLAM的实战路径

1. 六个月的紧迫感从哪来:先搞清楚你要成为哪种机器人工程师说实话,六个月的期限并不是一个宽松的时间线。市面上任何一本正经的机器人学教材都超过五百页,ROS2的官方文档可以翻到你怀疑人生,再加上ABB、KUKA这些工业机器人厂家动…

2026/9/14 11:59:31

USB Type-C PCB布局分区设计:电源、高速信号与PD协议全攻略

做硬件这行,Type-C接口算是典型的“看着简单,做起来全坑”的东西。光引脚就24个,高低速信号、电源、控制线全部塞在一个小小的连接器里,如果PCB布局不做规划,打样回来基本就是“插上没反应”、“高速掉线”、“静电一打…

2026/9/14 13:53:59

系统编程学习原型如何补齐稳定性边界

系统编程学习原型如何补齐稳定性边界预算有限时&#xff0c;我先优化明显多余的复制&#xff0c;而不是猜测性地换容器。用借用传递只读数据通常就能减少分配&#xff1a; fn parse(line: &str) -> Result<Item, Error> { /* ... */ }用基准确认热点确实在分配&am…

2026/9/14 11:22:57

雨花区哪家财务公司代理记账比较好?

在雨花区&#xff0c;企业处理财税事务常常面临诸多挑战&#xff0c;选择一家靠谱的财务公司至关重要。湖南巨勤财务管理咨询有限公司就是本地正规实体财税服务机构&#xff0c;深耕本地工商财税行业多年&#xff0c;熟悉当地工商局、税务局最新政策与申报流程。主营公司注册、…

还想了解更多?直接咨询顾问

免费诊断 + 免费方案 + 透明报价。

全国咨询热线400-8866-253
免费获取方案
咨询二维码