Transformer+GRU组合模型:滚动轴承故障诊断实战

发布时间:2026/9/20 14:08:25
Transformer+GRU组合模型:滚动轴承故障诊断实战 简介面向Python与深度学习研发人员及工业设备维护工程师这份资源是基于Transformer-GRU组合模型的故障诊断完整项目实例覆盖多源异构时序数据融合、长短期依赖提取、噪声与异常值处理等核心难题可应用于旋转机械、电力设备、汽车轨道、智能制造、能源设备等多种工业场景。项目注重数据质量包含信号滤波、异常值检测与修正等预处理步骤并融合多头自注意力与门控循环单元兼顾全局特征与局部时序变化。资源包内仅1个docx文档体积77KB内含程序源码、GUI设计说明及逐步代码详解便于系统学习与随查随用。目前已有805人浏览学习。读者可从中获得端到端建模流程、基于注意力权重的特征重要性可视化方法以及轻量化网络与多平台部署思路有助于在实际工业环境中高效部署并提升诊断结果的可信度。 很多做设备状态监测的朋友应该都有同感辛辛苦苦在实验环境里跑出一个故障诊断模型准确率看着挺高一到现场就露馅换个工况、换台设备结果就飘了。这两年Transformer架构大火大家都想用它来提特征但直接套在振动时序上训练慢、数据少的时候还容易过拟合。我自己试下来纯Transformer在小样本故障诊断上的表现远没有论文里说的那么神。后面我换了个思路把Transformer和GRU串成一个组合模型——Transformer负责捕捉长距离的全局依赖GRU专门负责压缩时序上的局部变化最后接分类头做故障识别。这个方案在公开的轴承数据集上诊断准确率比单纯的CNN或LSTM模型稳了不少而且整体参数量比纯Transformer小跑起来也快。这一篇我就把这个项目的完整实现拆开讲从数据切分到模型结构再到GUI界面设计代码都是一行行能跑的适合正在做故障诊断、设备健康管理或者毕设相关课题的同学参考。1. 内容整体设计与思路拆解为什么是Transformer-GRU而不是别的1.1 单一模型在小样本故障诊断上的局限先说Transformer。它的核心机制是自注意力能让序列里任意两个位置直接建立联系这对捕捉振动信号里的周期性冲击成分特别有利。但问题是自注意力是位置无关的它本身不感知先后顺序虽然可以加位置编码可处理一维时序信号时它对局部上下文比如相邻几个采样点之间的短时突变的建模能力并不强。GRU相反它是循环结构天然按时间顺序处理数据对短时局部特征的记忆很有效但一遇到长序列前面的信息容易衰减长距离依赖就抓不住了。所以组合模型的逻辑很简单先用Transformer的多头注意力把整个序列里重要的时间点都“看一遍”再把带有全局信息的特征序列交给GRU去提炼局部时序规律融合后的特征既知道这段信号“哪里异常”又知道“异常是怎么随时间变化的”。实测下来这种两级结构对滚动轴承的内圈故障、外圈故障、滚动体故障区分度特别明显。1.2 组合模型的结构与特征流向结构上我分了四段输入嵌入层把原始的一维振动信号切片后通过一维卷积做嵌入把每个时间步的维度映射到d_model。Transformer编码器层叠加两层编码器注意我这里只用了两层不是六层每层包含多头自注意力和前馈网络中间有残差连接和层归一化。GRU层把Transformer输出的序列继续输入单向GRU取最后一个时间步的隐藏状态作为整段序列的压缩向量。分类层通过全局平均池化或者直接用GRU最后一步的隐状态接全连接层用Softmax输出各类故障的概率。有一点需要说明GRU层的输入维度要和Transformer输出维度对齐否则会报维度不匹配的错。代码里我用的是d_model64注意力头数为4GRU隐藏层大小为32。1.3 为什么选择PyTorch来实现选PyTorch没有别的原因就是调试方便、生态成熟。故障诊断项目的特征是每次实验都要调参、改结构、加正则化PyTorch的动态图机制改起来不费劲。其次PyTorch自带的Dataset和DataLoader对大量振动样本的批次加载支持得很好后面做GUI推理时也能直接用TorchScript把模型序列化。2. 数据集准备与预处理全流程2.1 数据来源与故障类型定义这个项目用的是凯斯西储大学CWRU的滚动轴承公开数据集采样频率12kHz驱动端加速度计信号。数据集里的故障类型我用四分类的经典划分正常Normal、内圈故障Inner Race、外圈故障Outer Race、滚动体故障Ball。每类故障还有损伤直径的大小区分这里统一使用0.021英寸直径的数据保证故障特征足够明显。2.2 滑动窗口切分与重叠采样原始振动信号是连续的不能直接丢给模型。因为单条样本需要固定长度我采用了滑动窗口切分的方法。窗口长度设定为1024个采样点重叠率为50%这样可以显著扩充样本数量。以驱动端12kHz的数据为例一个10秒的振动信号有120000个点不重叠切分只能得到117个样本加了50%重叠后能拿到234个样本几乎是翻倍的效果。重叠采样还能缓解样本数不足的问题但要注意验证集和测试集不能和训练集有窗口重叠否则数据泄漏会让指标虚高。我通常的做法是先按连续的信号段划分再在各段内部切窗口。2.3 归一化与数据增强策略振动信号单位是m/s²加速度不同设备、不同负载下幅值范围差异很大直接喂给模型会导致训练不稳定。我先把每条样本做了Z-score标准化mean np.mean(sample) std np.std(sample) sample (sample - mean) / (std 1e-8)然后对训练集加了随机噪声增强在原始样本上叠加一个标准差为0.01的高斯噪声让模型对传感器噪声更鲁棒。注意增强只用于训练集测试集和验证集保持原始数据。2.4 数据集划分与标签设计所有类别共收集了12000个样本按照6:2:2划分为训练集、验证集和测试集。划分前打乱样本顺序保证每个batch里的类别分布是均匀的。标签采用one-hot编码故障类别到数字的映射为类别标签值正常0内圈故障1外圈故障2滚动体故障33. 核心模型架构与代码逐段详解3.1 嵌入层与位置编码Transformer不知道序列顺序所以必须在输入处加位置编码。这里我采用正弦位置编码就是原版Transformer论文里的公式它不用训练泛化能力好。import torch import torch.nn as nn import math class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len1024): super(PositionalEncoding, self).__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) div_term torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) self.register_buffer(pe, pe.unsqueeze(0)) def forward(self, x): # x: [batch_size, seq_len, d_model] return x self.pe[:, :x.size(1)]Embedding层我前文提到用一维卷积实现这里讲一下原因直接对原始1024个采样点做Transformer序列太长注意力的计算复杂度是O(n²)算不动。用conv1d做嵌入可以把1024个点压缩成64个时间步每步的向量维度是64回头再进Transformer就非常快。3.2 Transformer编码器层实现这里没有直接用nn.TransformerEncoderLayer而是自己写了一个单层编码器为了更方便后续做注意力权重的可视化导出。核心结构是多头自注意力 前馈网络 残差连接 层归一化。class TransformerBlock(nn.Module): def __init__(self, d_model, nhead, dim_feedforward128, dropout0.1): super(TransformerBlock, self).__init__() self.multi_head_attn nn.MultiheadAttention(d_model, nhead, dropoutdropout, batch_firstTrue) self.feed_forward nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.ReLU(), nn.Dropout(dropout), nn.Linear(dim_feedforward, d_model), ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x): # 自注意力部分 attn_out, attn_weights self.multi_head_attn(x, x, x) x self.norm1(x self.dropout1(attn_out)) # 前馈网络部分 ff_out self.feed_forward(x) x self.norm2(x self.dropout2(ff_out)) return x, attn_weights注意batch_firstTrue这个参数它让输入张量的维度顺序是(batch, seq, feature)和我习惯的数据流一致。如果不设这个参数默认输入是(seq, batch, feature)新手特别容易踩坑。3.3 GRU层与分类头设计GRU层的作用是承接Transformer输出的序列特征做进一步的局部时序特征压缩。我把GRU的return_sequences设为False相当于直接用最后一步的隐藏状态作为整个序列的总结向量。class GRUClassifier(nn.Module): def __init__(self, input_size, hidden_size, num_classes, num_layers1): super(GRUClassifier, self).__init__() self.gru nn.GRU(input_size, hidden_size, num_layers, batch_firstTrue, bidirectionalFalse) self.classifier nn.Sequential( nn.Linear(hidden_size, 64), nn.ReLU(), nn.Dropout(0.3), nn.Linear(64, num_classes) ) def forward(self, x): # x: [batch_size, seq_len, d_model] _, h_n self.gru(x) # 取最后一层的隐藏状态 out h_n[-1] # [batch_size, hidden_size] return self.classifier(out)这里我特意用单向GRU没有做双向因为双向会让模型在实际部署时没法做流式推理实时性会打折扣。故障诊断如果只是离线分析无所谓但一旦要上线到在线监测系统每一步操作都必须考虑可以在线部署。3.4 完整模型拼装class TransformerGRUModel(nn.Module): def __init__(self, input_size1024, d_model64, nhead4, num_encoder_layers2, gru_hidden32, num_classes4): super(TransformerGRUModel, self).__init__() self.embedding nn.Conv1d(1, d_model, kernel_size32, stride16) self.pos_encoder PositionalEncoding(d_model, max_len64) self.encoder_layers nn.ModuleList([ TransformerBlock(d_model, nhead) for _ in range(num_encoder_layers) ]) self.gru GRUClassifier(d_model, gru_hidden, num_classes) def forward(self, x): # x: [batch_size, 1024] 原始振动信号 x x.unsqueeze(1) # [batch_size, 1, 1024] x self.embedding(x) # [batch_size, d_model, 64] x x.permute(0, 2, 1) # [batch_size, 64, d_model] x self.pos_encoder(x) attn_weights_list [] for layer in self.encoder_layers: x, attn_weights layer(x) attn_weights_list.append(attn_weights) out self.gru(x) # [batch_size, num_classes] return out, attn_weights_list3.5 训练过程的细节与超参数设定训练采用AdamW优化器初始学习率1e-3权重衰减1e-4。损失函数使用交叉熵。训练轮数上限设为80轮但配合早停Early Stopping连续8轮验证集准确率不提升就停止并恢复最优模型参数。criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30)训练过程中的几个重要经验Batch size选64太大的batch在局部小样本数据上容易导致早期过拟合。学习率一定要配合衰减调度器我用的是余弦退火比固定学习率收敛更稳定。如果训练loss降到0.01以下而验证loss开始上升说明模型开始记住训练数据的噪声了早停要果断。4. GUI设计与代码实现4.1 GUI功能需求分析模型训练好了最终还是得让不会写代码的工程师用。所以我用PySide6写了一个图形化诊断工具整个界面分为四个区域数据加载区、模型加载区、诊断显示区、运行日志区。需要实现的操作流程是加载振动信号文件支持csv/txt格式→加载训练好的模型权重→点击“诊断”按钮→界面输出故障类型、置信度以及注意力热力图。4.2 核心界面布局代码from PySide6.QtWidgets import (QApplication, QMainWindow, QWidget, QVBoxLayout, QHBoxLayout, QPushButton, QLabel, QFileDialog, QTextEdit, QComboBox) from PySide6.QtCore import Qt class FaultDiagnosisGUI(QMainWindow): def __init__(self): super().__init__() self.setWindowTitle(Transformer-GRU 轴承故障诊断系统) self.setMinimumSize(900, 600) self.model None self.data None self.init_ui() def init_ui(self): central_widget QWidget() self.setCentralWidget(central_widget) layout QVBoxLayout(central_widget) # 数据加载行 load_layout QHBoxLayout() self.load_data_btn QPushButton(加载振动数据) self.load_model_btn QPushButton(加载模型权重) self.diagnose_btn QPushButton(开始诊断) self.diagnose_btn.setEnabled(False) load_layout.addWidget(self.load_data_btn) load_layout.addWidget(self.load_model_btn) load_layout.addWidget(self.diagnose_btn) layout.addLayout(load_layout) # 结果显示 self.result_label QLabel(等待加载数据...) self.result_label.setAlignment(Qt.AlignCenter) layout.addWidget(self.result_label) # 日志 self.log_box QTextEdit() self.log_box.setReadOnly(True) layout.addWidget(self.log_box)4.3 模型推理与结果展示逻辑推理时要把原始振动信号切窗、归一化、转Tensor然后送给模型。多个窗口的结果用投票法集成每个窗口输出一个故障类别最后取出现次数最多的类别作为最终结果置信度则用各类别的平均概率表示。def run_diagnosis(self): if self.model is None or self.data is None: self.log(请先加载数据和模型) return self.log(开始诊断...) windows self.preprocess_data(self.data) predictions [] probabilities [] self.model.eval() with torch.no_grad(): for w in windows: w_tensor torch.tensor(w, dtypetorch.float32).unsqueeze(0) probs torch.softmax(self.model(w_tensor)[0], dim1) pred torch.argmax(probs, dim1).item() predictions.append(pred) probabilities.append(probs.squeeze().numpy()) final_pred max(set(predictions), keypredictions.count) avg_probs np.mean(probabilities, axis0) self.result_label.setText(f诊断结果: {self.class_names[final_pred]} 置信度: {avg_probs[final_pred]:.4f})这里最关键的一点是预处理逻辑必须和训练时完全一致包括归一化的均值和标准差、窗口长度、重叠率。很多同学独立写GUI时结果不对八成是因为测试时的预处理和训练时不一致。4.4 注意力热力图可视化为了增强诊断结果的可解释性我在界面右侧加了一个matplotlib画布将最后一个Transformer层最后一批样本的注意力权重画成热力图。越亮的位置代表模型在诊断时越关注该时间步。这样工程师拿到诊断结果后还能反查是哪个时间段触发了模型决策对分析实际故障很有参考价值。5. 常见问题与排查技巧实录问题可能原因解决方案训练loss下降很快但验证集准确率不涨过拟合数据增强不足增大dropout、增强噪声、增加样本量模型对某一类故障的召回率特别低类别不平衡或者该类样本特征不明显用类别权重加权损失或对该类做SMOTE过采样Transformer部分参数量太大导致OOM序列长度过长、d_model过大增大卷积嵌入的步幅减少注意力层数GUI中加载模型后推理报维度错误输入数据长度和训练时不一致检查预处理窗口长度和模型输入层长度测试集准确率高但实际运行不稳定训练集和测试集分布不一致加入更多工况数据做域适应再说一个大家很容易忽略的坑CWRU数据集的故障特征其实比较明显模型在上面跑出99%以上的准确率并不稀奇这不代表模型在真实场景就能直接用。真实设备的转速波动、负载变化、噪声干扰都会让特征分布发生偏移。所以项目里我建议大家训练完后在另一个负载条件下的数据上做一个跨域测试看准确率掉多少。如果掉得厉害就需要用迁移学习的方法做适配。还有个调参经验分享一下。组合模型的瓶颈通常不在Transformer层数而在GRU的隐藏层大小。初始阶段我从gru_hidden64调到了32识别准确率反而上升了推测隐藏层太大引入了多余的参数噪声。针对小规模故障数据集模型越复杂越容易过拟合Architecture不是越大越好。6. 项目总结与实际使用心得项目做到这里整体流程已经闭环数据采集 → 预处理 → 模型训练 → 模型导出 → GUI部署。前前后后调试了大概两周时间最大的感受是组合模型在故障诊断场景下确实比单一模型更有优势尤其是特征提取的两段式思路既有全局视角又保留局部细节。想再提一个很多人忽略的点GUI和模型之间的耦合问题。训练时的模型类定义和GUI里加载模型时用的模型类必须保持一致包括相同的d_model、nhead、num_encoder_layers等参数。建议在训练完成后把模型结构参数保存成一个json文件GUI启动时读取这个json来动态构建模型再加载权重这样就算以后改了模型结构旧GUI也能兼容。实际使用中这个诊断工具对单个约3秒的振动文件整个推理加界面刷新的时间在0.5秒以内完全可以做到近实时诊断。如果后续要进一步提高可以考虑把模型转换成ONNX格式用onnxruntime来做推理速度还能再快一些。最后留一个小贴士做GUI时PySide6的信号槽里千万别做耗时的模型推理动作要不然界面会卡死。要么把推理放到QThread线程里要么在推理前禁用按钮并让鼠标等待推理完成后恢复。我因为这个卡顿问题还专门重构过一次界面代码能从坑里省下的时间比从任何理论技巧里省下的都多。本文还有配套的精品资源点击获取