基于深度学习的滚动轴承故障诊断:从数据预处理到模型部署全流程实战
2026/8/28 14:09:16 网站建设 项目流程

简介:深度学习作为人工智能的核心技术,通过模拟人脑神经网络结构,能够从海量数据中自动学习复杂特征与模式。其核心原理在于构建多层非线性变换,逐层提取和组合数据中的抽象表示,最终完成分类、回归等任务。在工业领域,这项技术的价值在于实现从“经验驱动”到“数据驱动”的智能化决策转变,尤其在预测性维护场景中,能够提前预警设备潜在故障,避免非计划停机。针对旋转机械关键部件——滚动轴承的故障诊断,深度学习模型如卷积神经网络(CNN)能够直接从振动信号中自动学习故障特征,替代传统依赖专家经验的信号分析方法。本文以公开的CWRU轴承数据集和实际风电项目为例,系统阐述了从振动信号预处理、数据增强,到1D-CNN、2D-CNN及CNN-LSTM混合模型构建、训练调优,直至模型轻量化与工程化部署的完整技术路径,为工业AI落地提供了一套可复现的实战框架。

1. 从“听声辨位”到“数据驱动”:工业设备故障诊断的范式转变

在工业领域,尤其是旋转机械的维护中,滚动轴承的健康状况直接关系到整条生产线的稳定与安全。传统的故障诊断,很大程度上依赖于老师傅的“听声辨位”或定期拆检,这不仅效率低下,对经验依赖性强,而且往往在故障已经发展到一定程度时才能被发现,容易造成非计划停机,带来巨大的经济损失。随着传感器技术和数据采集系统的普及,我们获得了海量的设备运行数据,如何从这些看似杂乱无章的振动、温度、声音信号中,精准、提前地识别出轴承的早期故障,就成了一个极具价值的课题。

深度学习,作为人工智能领域近年来最耀眼的技术之一,为我们提供了全新的解决方案。它不再需要人工设计复杂的特征提取算法(比如计算峭度、峰值因子、包络谱等),而是能够直接从原始振动信号中,自动学习并构建出最能表征故障状态的特征表示。这就像给机器装上了一双能“透视”设备内部状态的“慧眼”。今天,我就结合自己在一个实际风电齿轮箱轴承故障诊断项目中的经验,手把手地带你走一遍基于深度学习的滚动轴承故障诊断全流程。我们将从数据获取、预处理、模型构建、训练调优,一直讲到模型部署和实际应用中的坑点。无论你是刚接触工业AI的工程师,还是有一定机器学习基础想转向实战的研究者,这篇文章都将提供一条清晰的路径和可直接复现的代码框架。

2. 数据:一切智能诊断的基石与第一道难关

没有高质量的数据,再精巧的模型也只是空中楼阁。在轴承故障诊断中,数据工作往往占据了整个项目70%以上的精力。这部分我们将深入探讨数据的来源、处理以及如何为深度学习模型准备“食材”。

2.1 数据来源与公开数据集的选择

对于初学者和研究者而言,使用公开数据集是快速入门和验证算法性能的最佳途径。最著名、使用最广泛的莫过于凯斯西储大学(CWRU)轴承数据中心的数据。这个数据集几乎成了该领域的“基准测试集”(Benchmark)。它模拟了驱动端和风扇端轴承在不同负载(0HP, 1HP, 2HP, 3HP)下的多种故障状态,包括内圈故障、外圈故障、滚动体故障,每种故障又有不同尺寸(0.007英寸, 0.014英寸, 0.021英寸)。数据采样频率为12kHz和48kHz,提供了丰富的分析维度。

注意:虽然CWRU数据集非常经典,但它是在实验室环境下、单一故障点、恒定转速下采集的。实际工业现场的数据要复杂得多:变转速、变负载、强背景噪声、多种故障耦合、传感器安装位置差异等。因此,在CWRU上表现优异的模型,直接应用到现场可能效果大打折扣。但作为学习和算法验证的起点,它无可替代。

除了CWRU,还有如MFPT(机械故障预防技术学会)数据集PU(帕德博恩大学)轴承数据集等,后者包含了更真实的工况和多种损伤程度,挑战性更大。在我们的项目中,为了模拟更真实的场景,我选择以CWRU数据为基础,但会重点讲解如何通过数据增强来模拟现场的不确定性。

2.2 振动信号的预处理与特征工程“平替”

深度学习号称可以“端到端”学习,但恰当的预处理能极大提升模型收敛速度和最终性能。对于振动信号,标准的预处理流程包括:

  1. 去趋势(Detrending):移除信号中可能存在的线性或缓慢变化的趋势项,这些通常由温度漂移或传感器零点漂移引起。可以使用简单的滑动平均或减去信号均值来实现。
  2. 带通滤波(Bandpass Filtering):轴承故障特征频率通常集中在某个频带内。根据轴承型号和转速,可以估算出故障特征频率的大致范围(如外圈故障频率BPFO、内圈故障频率BPFI等),并设计一个带通滤波器,滤除无关的高频噪声和低频干扰。在实际操作中,我常用5阶巴特沃斯带通滤波器
  3. 标准化(Normalization):将数据缩放到一个固定的范围,如[0, 1]或[-1, 1],或者进行Z-score标准化(减去均值,除以标准差)。这能防止某些维度数值过大而主导模型训练。我通常对每个样本单独进行Z-score标准化,因为不同工况下的信号幅值差异可能很大。

那么,还需要传统的特征工程吗?对于深度学习模型,特别是卷积神经网络(CNN),我们可以将原始预处理后的时域信号直接作为输入。但一种非常有效且常用的“平替”方案是将一维时域信号转换为二维时频图像。这相当于为CNN提供了更直观、信息密度更高的输入。最常用的方法是短时傅里叶变换(STFT)生成频谱图,或者连续小波变换(CWT)生成小波尺度图。

在我的项目中,我对比了直接输入时域信号和输入STFT频谱图的效果。发现对于CWRU这类信噪比较高的数据,直接输入时域信号的简单1D-CNN已经能取得不错的效果(>98%准确率)。但对于更嘈杂的MFPT数据,使用STFT频谱图作为2D-CNN的输入,模型鲁棒性明显更强,准确率能提升3-5个百分点。这是因为时频图像能更好地将故障的周期性冲击特征从背景噪声中分离出来。

实操心得:预处理流程不是一成不变的。你需要根据数据的实际情况(信噪比、是否变速)来调整。一个简单的判断方法是,画出原始信号的时域波形和频谱图,肉眼观察故障特征是否明显。如果时域上就能看到明显的周期性冲击,那么1D输入可能就够了;如果特征淹没在噪声中,但频谱的某个频段有突出谱线,那么时频图像会是更好的选择。

2.3 数据切片与增强:应对样本不足的利器

公开数据集通常提供了长时间序列,我们需要将其切割成固定长度的样本。样本长度是一个超参数,太短可能包含不了一个完整的故障冲击周期,太长则计算开销大且可能包含无关信息。一个经验法则是:样本长度应至少包含2-3个故障特征周期。例如,假设轴承转速为1800 RPM,计算出的BPFO约为107 Hz,那么故障周期约为9.3毫秒。在12kHz采样率下,对应112个采样点。因此,样本长度可取256或512点,以确保信息完整性。

数据增强是解决工业现场标签数据稀缺的关键技术。对于振动信号,我们可以使用以下方法:

  • 加性噪声:添加高斯白噪声或实际采集的背景噪声,模拟传感器噪声和环境干扰。
  • 幅度缩放:对信号幅值进行随机微小的缩放,模拟负载的微小变化。
  • 时间偏移(Time Shift):在样本窗口内随机滚动信号,这不会改变信号的频率成分,但能增加多样性。
  • 频率扭曲(Frequency Warping):轻微改变信号的频率成分,模拟转速的微小波动。

在我的代码中,我定义了一个VibrationDataAugmentor类,在训练时在线(on-the-fly)对每个batch的数据随机应用1-2种增强策略,显著提升了模型在噪声数据上的泛化能力。

3. 模型选型与构建:从1D-CNN到混合架构的演进

选择了合适的数据表示形式后,下一步就是设计或选择神经网络模型。轴承故障诊断本质上是一个时间序列分类问题。以下是几种主流且有效的模型架构。

3.1 一维卷积神经网络(1D-CNN):轻量高效的起点

1D-CNN直接在预处理后的时域信号上滑动卷积核,能够自动提取局部时间模式(如故障冲击)。其结构简单,参数少,训练快,非常适合作为基线模型。

import torch import torch.nn as nn import torch.nn.functional as F class Simple1DCNN(nn.Module): def __init__(self, input_length=1024, num_classes=10): super(Simple1DCNN, self).__init__() self.conv1 = nn.Conv1d(in_channels=1, out_channels=16, kernel_size=64, stride=2, padding=32) self.bn1 = nn.BatchNorm1d(16) self.pool1 = nn.MaxPool1d(kernel_size=2) self.conv2 = nn.Conv1d(16, 32, kernel_size=3, stride=1, padding=1) self.bn2 = nn.BatchNorm1d(32) self.pool2 = nn.MaxPool1d(2) self.conv3 = nn.Conv1d(32, 64, kernel_size=3, stride=1, padding=1) self.bn3 = nn.BatchNorm1d(64) self.pool3 = nn.MaxPool1d(2) # 计算全连接层输入尺寸 # 经过三次池化,长度变为 input_length // (2*2*2) = input_length // 8 fc_input_dim = 64 * (input_length // 8) self.fc1 = nn.Linear(fc_input_dim, 100) self.dropout = nn.Dropout(0.5) self.fc2 = nn.Linear(100, num_classes) def forward(self, x): # x shape: (batch_size, 1, signal_length) x = F.relu(self.bn1(self.conv1(x))) x = self.pool1(x) x = F.relu(self.bn2(self.conv2(x))) x = self.pool2(x) x = F.relu(self.bn3(self.conv3(x))) x = self.pool3(x) x = x.view(x.size(0), -1) # 展平 x = F.relu(self.fc1(x)) x = self.dropout(x) x = self.fc2(x) return x

为什么这样设计?

  • 大尺寸首层卷积核(kernel_size=64):轴承故障冲击在时域上是一个短暂的瞬态事件,较大的卷积核有助于捕捉这个局部形状。
  • 步长(stride=2):在首层使用稍大的步长,可以快速降低序列长度,减少计算量,同时引入一些平移不变性。
  • 批归一化(BatchNorm):加速训练,提供轻微的正则化效果,使模型对初始化和学习率更不敏感。
  • Dropout:在全连接层前使用,防止过拟合,这对于数据量相对较小的故障诊断任务尤为重要。

3.2 二维卷积神经网络(2D-CNN):挖掘时频域深层特征

当使用STFT频谱图作为输入时,2D-CNN是自然的选择。你可以使用经典的图像分类网络作为骨干,如ResNet、VGG的变体,或者设计一个更轻量化的专用网络。

class Spec2DCNN(nn.Module): def __init__(self, num_classes=10): super(Spec2DCNN, self).__init__() # 输入假设为 (batch, 1, freq_bins, time_frames) self.features = nn.Sequential( nn.Conv2d(1, 32, kernel_size=3, padding=1), # [b,32,f,t] nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # [b,32,f/2,t/2] nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # [b,64,f/4,t/4] nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # [b,128,f/8,t/8] ) # 需要根据输入频谱图的尺寸计算这里 self.avgpool = nn.AdaptiveAvgPool2d((4, 4)) # 自适应池化到固定尺寸 self.classifier = nn.Sequential( nn.Dropout(), nn.Linear(128 * 4 * 4, 512), nn.ReLU(inplace=True), nn.Dropout(), nn.Linear(512, num_classes), ) def forward(self, x): x = self.features(x) x = self.avgpool(x) x = torch.flatten(x, 1) x = self.classifier(x) return x

实操心得:对于频谱图,卷积核在频率轴和时间轴上的感受野具有不同的物理意义。在频率轴上的卷积有助于识别故障特征频率所在的频带;在时间轴上的卷积有助于识别特征的周期性。在设计网络时,可以考虑使用非对称卷积核(如kernel_size=(5,3)),给频率维和时间维分配不同的权重。

3.3 混合模型:CNN与LSTM/Transformer的联姻

为了同时利用信号的局部特征和长期依赖关系(如故障冲击的周期),可以将CNN与循环神经网络(如LSTM)或Transformer结合。CNN作为特征提取器,将长序列压缩为高级特征序列,然后送入LSTM或Transformer编码器捕捉时序依赖,最后通过全连接层分类。

class CNN_LSTM(nn.Module): def __init__(self, input_length, num_classes): super(CNN_LSTM, self).__init__() # CNN部分用于特征提取 self.cnn = nn.Sequential( nn.Conv1d(1, 64, kernel_size=80, stride=4), nn.BatchNorm1d(64), nn.ReLU(), nn.MaxPool1d(4), nn.Conv1d(64, 128, kernel_size=3), nn.BatchNorm1d(128), nn.ReLU(), nn.MaxPool1d(4), ) # 计算CNN输出长度 cnn_out_length = self._get_cnn_out_len(input_length) # LSTM部分捕捉时序 self.lstm = nn.LSTM(input_size=128, hidden_size=64, num_layers=2, batch_first=True, bidirectional=True, dropout=0.3) # 双向LSTM,输出维度为 hidden_size * 2 self.fc = nn.Linear(64 * 2, num_classes) def _get_cnn_out_len(self, length): # 模拟计算经过CNN后的序列长度 length = (length - 80) // 4 + 1 # conv1 length = length // 4 # pool1 length = (length - 3) // 1 + 1 # conv2 length = length // 4 # pool2 return length def forward(self, x): # x: [batch, 1, seq_len] cnn_features = self.cnn(x) # [batch, 128, cnn_out_len] # 将通道维变为特征维,以适应LSTM输入: [batch, seq_len, features] cnn_features = cnn_features.permute(0, 2, 1) lstm_out, _ = self.lstm(cnn_features) # [batch, cnn_out_len, hidden_size*2] # 取最后一个时间步的输出,或者使用全局平均/最大池化 last_step_out = lstm_out[:, -1, :] output = self.fc(last_step_out) return output

为什么选择混合模型?在变转速工况下,故障冲击的间隔时间会发生变化。单纯的CNN可能难以适应这种时间尺度上的变化,而LSTM或Transformer能够学习这种动态的时间依赖关系,理论上具有更好的泛化能力。在我的风电项目变工况数据测试中,CNN-LSTM混合模型比纯CNN模型的准确率稳定高出约2%。

4. 模型训练、调优与评估:不仅仅是准确率

构建好模型只是第一步,如何训练和评估它,才是决定项目成败的关键。

4.1 损失函数与优化器选择

这是一个多分类问题,因此损失函数自然选择交叉熵损失(CrossEntropyLoss)。优化器方面,Adam因其自适应学习率特性,在大多数情况下都是安全且高效的首选。对于非常深或非常敏感的网络,SGD with Momentum配合学习率衰减策略,有时能收敛到更优的极小点,但需要更多的调参技巧。

import torch.optim as optim from torch.nn import CrossEntropyLoss model = Simple1DCNN(input_length=1024, num_classes=10) criterion = CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4) # 加入L2正则化 scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=5, verbose=True)

学习率调度器(Scheduler)非常重要。ReduceLROnPlateau会在验证集损失不再下降时自动降低学习率,这能帮助模型在训练后期更精细地调整参数,避免在最优解附近震荡。

4.2 训练循环与早停策略

训练循环的代码框架大同小异,但有几个细节决定效率:

  1. 设备转移:明确使用model.to(device)data.to(device)
  2. 梯度清零:每个batch前必须执行optimizer.zero_grad()
  3. 梯度裁剪:对于RNN/LSTM,使用torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)可以防止梯度爆炸。
  4. 早停(Early Stopping):这是防止过拟合的利器。记录验证集损失,如果连续多个epoch(如10个)没有下降,则停止训练,并回滚到验证损失最低的模型参数。
def train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=100, patience=10): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model.to(device) best_val_loss = float('inf') patience_counter = 0 best_model_state = None for epoch in range(num_epochs): # 训练阶段 model.train() running_loss = 0.0 for signals, labels in train_loader: signals, labels = signals.to(device), labels.to(device) optimizer.zero_grad() outputs = model(signals) loss = criterion(outputs, labels) loss.backward() # 可选:梯度裁剪 # torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() running_loss += loss.item() * signals.size(0) epoch_train_loss = running_loss / len(train_loader.dataset) # 验证阶段 model.eval() val_loss = 0.0 correct = 0 total = 0 with torch.no_grad(): for signals, labels in val_loader: signals, labels = signals.to(device), labels.to(device) outputs = model(signals) loss = criterion(outputs, labels) val_loss += loss.item() * signals.size(0) _, predicted = torch.max(outputs.data, 1) total += labels.size(0) correct += (predicted == labels).sum().item() epoch_val_loss = val_loss / len(val_loader.dataset) val_acc = 100 * correct / total scheduler.step(epoch_val_loss) # 根据验证损失调整学习率 print(f'Epoch {epoch+1}/{num_epochs}: Train Loss: {epoch_train_loss:.4f}, Val Loss: {epoch_val_loss:.4f}, Val Acc: {val_acc:.2f}%') # 早停逻辑 if epoch_val_loss < best_val_loss: best_val_loss = epoch_val_loss best_model_state = model.state_dict().copy() patience_counter = 0 # 可以在这里保存最佳模型 torch.save(...) else: patience_counter += 1 if patience_counter >= patience: print(f'Early stopping triggered at epoch {epoch+1}') break # 训练结束,加载最佳模型 if best_model_state is not None: model.load_state_dict(best_model_state) return model

4.3 超越准确率:更全面的评估指标

对于不平衡数据集(如正常样本远多于故障样本),准确率是欺骗性的。必须引入更细致的评估指标:

  • 混淆矩阵(Confusion Matrix):直观展示每个类别被分对和分错的情况,能立刻发现模型在哪些具体故障类型上识别困难。
  • 精确率(Precision)、召回率(Recall)和F1分数(F1-Score):针对每个类别计算。在故障诊断中,我们通常更关心召回率,即“漏报率”要低,宁可误报也不能漏报一个真实故障。
  • 宏平均(Macro-average)和微平均(Micro-average)F1:宏平均对所有类别平等看待,微平均则考虑样本数量权重。对于类别不平衡的数据,关注宏平均F1更有意义。

在测试集上,应该输出完整的分类报告:

from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns import matplotlib.pyplot as plt def evaluate_model(model, test_loader, class_names): device = next(model.parameters()).device model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for signals, labels in test_loader: signals = signals.to(device) outputs = model(signals) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) # 分类报告 print(classification_report(all_labels, all_preds, target_names=class_names, digits=4)) # 混淆矩阵可视化 cm = confusion_matrix(all_labels, all_preds) plt.figure(figsize=(10,8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.ylabel('True Label') plt.xlabel('Predicted Label') plt.title('Confusion Matrix') plt.tight_layout() plt.show()

实操心得:在工业场景中,模型的可解释性同样重要。可以使用梯度加权类激活映射(Grad-CAM)等技术,可视化模型在做出“内圈故障”判断时,重点关注了原始信号或频谱图的哪些部分。这不仅能增加工程师对模型的信任,还能帮助我们发现数据或模型的问题。例如,如果Grad-CAM显示模型主要根据信号开头的一段噪声做判断,那说明模型可能学到了无关特征,需要检查数据或增加正则化。

5. 从实验室到现场:工程化部署与持续学习

在实验室用Jupyter Notebook跑出99%的准确率只是第一步。真正的挑战在于将模型部署到工业现场,并让其持续稳定地工作。

5.1 模型轻量化与优化

部署在边缘设备(如工控机、嵌入式系统)上时,对模型的大小和推理速度有严格要求。可以采用以下技术:

  • 知识蒸馏(Knowledge Distillation):用一个大模型(教师)指导一个小模型(学生)训练,让小模型获得接近大模型的性能。
  • 剪枝(Pruning):移除网络中不重要的连接或通道。
  • 量化(Quantization):将模型参数从32位浮点数转换为8位整数,可以大幅减少模型体积和提升推理速度,对精度影响通常很小。PyTorch提供了方便的量化工具torch.quantization
  • 使用更高效的架构:如MobileNet、ShuffleNet的1D版本,或专门为边缘计算设计的网络。

5.2 部署模式:云端、边缘端与混合模式

  • 云端部署:将数据通过网络传输到云服务器进行推理。优点是算力强,便于模型集中更新和管理。缺点是对网络稳定性要求高,有数据安全和延迟问题。适用于对实时性要求不高(如分钟级)、数据量大的周期性分析。
  • 边缘端部署:在设备现场的工控机或智能传感器内进行推理。优点是实时性高(毫秒级),数据不出局域网,安全性好。缺点是算力有限,模型更新麻烦。适用于对实时性要求极高的预测性维护场景。
  • 混合部署:轻量级模型在边缘端做实时监测和预警,原始数据或特征定期同步到云端,由更复杂的模型做深度分析和模型再训练,形成闭环。

在我们的风电项目中,采用了混合模式。每个风机机舱内的边缘计算盒运行一个轻量化的1D-CNN模型,每10秒进行一次诊断。同时,每小时的振动数据片段会上传至云端,用于触发更复杂的分析(如故障严重程度评估、剩余寿命预测)和模型迭代。

5.3 持续学习与模型迭代:应对数据分布漂移

工业设备的状态、工况、环境都在缓慢变化,这会导致训练数据(旧数据)和实际 inference 数据(新数据)的分布发生漂移,模型性能会随时间下降。因此,必须建立持续学习(Continual Learning)机制。

一个实用的策略是主动学习(Active Learning)结合人工复核

  1. 边缘端模型对预测置信度低(如softmax概率低于0.9)的样本进行标记。
  2. 这些“不确定”样本连同其原始数据被上传到云端。
  3. 云端系统将其推送给领域专家进行人工标注。
  4. 用新标注的数据定期(如每季度)对模型进行增量训练或微调。

踩坑实录:我们最初忽略了数据漂移问题,一个在夏季数据上训练良好的模型,在冬季出现了大量误报。后来分析发现,环境温度变化影响了轴承的润滑状态,导致振动信号的基线发生了偏移。解决方法是收集了全年不同季节的数据进行重新训练,并在预处理中增加了更鲁棒的标准化方法(如滑动窗口标准化)。

6. 开源代码框架与项目实战建议

为了方便大家复现和在此基础上进行开发,我整理了一个基于PyTorch的轴承故障诊断最小可行项目结构。你可以通过以下方式获取并运行:

(假设项目名为BearingFaultDL

BearingFaultDL/ ├── data/ │ ├── CWRU/ # 存放CWRU数据集 │ ├── preprocess.py # 数据下载、切片、增强、STFT转换等 │ └── dataset.py # 自定义PyTorch Dataset类 ├── models/ │ ├── __init__.py │ ├── cnn_1d.py # 1D-CNN模型定义 │ ├── cnn_2d.py # 2D-CNN模型定义 │ └── cnn_lstm.py # 混合模型定义 ├── utils/ │ ├── config.py # 超参数配置 │ ├── logger.py # 日志记录 │ └── visualize.py # 可视化工具(波形、频谱、混淆矩阵等) ├── train.py # 模型训练主脚本 ├── evaluate.py # 模型评估脚本 ├── inference.py # 单样本推理示例 └── requirements.txt # 项目依赖

核心脚本train.py的简化逻辑:

# train.py import argparse from data.dataset import get_data_loaders from models.cnn_1d import Simple1DCNN from utils.config import Config from utils.logger import setup_logger # ... 其他导入 def main(config): logger = setup_logger() logger.info(f"Using device: {config.device}") # 1. 加载数据 train_loader, val_loader, test_loader, class_names = get_data_loaders(config) logger.info(f"Classes: {class_names}") # 2. 初始化模型、损失函数、优化器 model = Simple1DCNN(input_length=config.sample_length, num_classes=len(class_names)).to(config.device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=config.lr, weight_decay=config.weight_decay) scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=5) # 3. 训练 trainer = Trainer(model, criterion, optimizer, scheduler, config, logger) trainer.train(train_loader, val_loader) # 4. 在测试集上评估最佳模型 evaluator = Evaluator(trainer.best_model, test_loader, class_names, config.device) evaluator.evaluate() # 5. 保存模型 torch.save(trainer.best_model.state_dict(), f'best_model_{config.model_name}.pth') if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--config', type=str, default='config.yaml', help='Path to config file') args = parser.parse_args() config = Config.from_yaml(args.config) main(config)

给初学者的实战建议:

  1. 从CWRU数据集和1D-CNN开始:不要一开始就追求复杂的模型和架构。先用最简单的流程(下载数据 -> 切片 -> 训练1D-CNN)跑通整个Pipeline,确保代码和环境没有问题。把准确率做到95%以上,建立信心。
  2. 可视化一切:在数据预处理、训练过程中,多画图。画出原始信号、频谱图、训练损失/准确率曲线、混淆矩阵。视觉反馈能帮你最快地发现问题。
  3. 构建严谨的评估流程:务必使用独立的测试集来报告最终性能,不要在验证集上调参然后说这是最终结果。采用交叉验证可以获得更稳健的性能估计。
  4. 理解你的错误:分析混淆矩阵,看模型主要把哪两类搞混了。然后回去看这两类故障的原始信号和特征有什么相似之处,思考是数据问题、特征问题还是模型容量问题。
  5. 考虑现实约束:在项目早期就思考部署目标。是需要毫秒级响应?还是模型必须小于10MB?这些约束会直接影响你的模型选型和优化策略。

基于深度学习的故障诊断是一个充满挑战但也回报丰厚的领域。它不仅仅是调参炼丹,更需要对机械原理、信号处理和软件工程都有所了解。希望这份从理论到实践、从实验室到现场的详细指南,能为你点亮一盏灯,助你少走弯路,更快地将这项技术应用到解决实际工业问题的征程中。

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

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

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

立即咨询