尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

CNN与Transformer融合的轴承故障诊断实战:基于CWRU数据集的跨工况方案

发布时间:2026/9/29 18:50:07

资讯中心
01
ARTICLE

CNN与Transformer融合的轴承故障诊断实战:基于CWRU数据集的跨工况方案

CNN与Transformer融合的轴承故障诊断实战:基于CWRU数据集的跨工况方案
简介面向轴承故障诊断与深度学习交叉领域的研究者和初学者这份资料提供基于CNN与Transformer融合模型的完整实验方案依托CWRU标准数据集覆盖数据预处理、模型构建、训练验证与结果分析等环节适合学习智能运维、信号处理及预测性维护的读者。压缩包共27个文件大小约16.65MB主要包含10个mat格式振动信号数据、3个py源码脚本、1个ipynb交互式notebook、1个keras训练模型及csv标签文件另有说明文档和备份文件便于对照复现与二次开发。目前已有84人学习下载。资料以CWRU十类故障分类为切入点提供可直接运行的app.py、信号生成脚本和项目说明可快速搭建完整诊断流程理解CNN局部特征提取与Transformer长程依赖捕捉的互补优势同时模型文件与备份资源也有助于排查实验问题对开展轴承故障识别研究具有实用参考价值。1. 基于CNN与Transformer融合的轴承故障诊断先弄明白这条路子到底能解决什么很多人在CWRU数据集上用单一一套CNN模型跑出来的准确率能打到99%以上就以为轴承故障诊断已经到头了。但把同一个模型挪到其他转速、其他负载的工况上准确率会明显掉下去。这个现象不只在CWRU上出现在实验室自采数据里更严重。原因不复杂CNN善于抓局部形态却很难把振动信号里相隔较远的时序依赖关系串起来而Transformer恰好擅长建模长程依赖但对局部细小的冲击特征又不如CNN敏感。于是“CNN与Transformer融合”成了一个实际的工程方向——用CNN提炼局部冲击特征用Transformer捕捉信号内部的时序关联最后再融合分类。这篇文章就围绕这个思路讲清楚融合模型怎么搭、CWRU数据集怎么用、参数怎么调、坑在哪里。适用人群很明确刚入门故障诊断、想在公开数据集上复现一个能用的深度学习方案的研究生或算法工程师已经在用纯CNN或纯Transformer做振动诊断、想进一步提升跨工况鲁棒性的从业者。这篇文章不含学术空谈只讲能落地的步骤。2. 为什么要融合CNN和Transformer各自擅长什么又各自缺什么2.1 CNN处理振动信号的强项和边界振动信号本质是一维时间序列。CNN在这一领域的常见做法是用一维卷积核沿时间轴滑动提取局部波形形态。例如一个kernel size为64的卷积核一次能“看到”64个采样点内的波形起伏对于轴承故障信号中的周期性冲击这种局部卷积能很有效地捕捉冲击形态、幅值变化和频带能量分布。但CNN的局限也很具体它的感受野是逐层扩大的想覆盖一个完整旋转周期的信号往往需要堆很多层。以12kHz采样率、转速1797 r/min的CWRU数据来说转一圈约400个采样点附近要覆盖整圈信号至少需要数千采样点的感受野。堆深网络可以做到但参数量变大训练变慢而且深层CNN容易把局部噪声也当成特征学进去。另一个实际问题是CNN对输入顺序不敏感。卷积核在信号任意位置提取到的模式是“平移等变”的这既是优点也是缺点——它不会主动建模“先出现冲击、后出现共振衰减”这种顺序关系。而轴承故障信号里故障冲击的发生顺序和间隔恰恰包含重要信息。2.2 Transformer建模长程依赖的机制与代价Transformer的核心是自注意力机制。对一段振动信号它计算每个时间位置与其他所有位置的相关性权重从而直接建模任意距离的依赖关系。在热词里高频出现的“Transformer架构”“注意力机制”讨论其关键点就在这里不需要像CNN那样逐层堆叠扩大感受野一步到位。具体到轴承故障诊断一段信号里如果包含多个故障冲击周期Transformer能直接让模型关注“上一个冲击位置”和“下一个冲击位置”之间的关系从而把旋转周期、故障特征频率这类全局信息纳入特征表达。这对变工况诊断很有价值——不同负载下转频变了但冲击间隔与转频的倍数关系是稳定的。但纯Transformer在振动信号上也有明显问题。注意力机制对位置编码敏感而振动信号不像自然语言有明确的词边界如果直接对原始采样点做注意力计算量极大且难以捕捉局部冲击形态。另外Transformer需要大量数据才能训练充分CWRU单类样本数通常几千条纯Transformer很容易过拟合。2.3 融合策略左右分支并行还是串行级联工程上最常见的做法是双分支并行融合。左边一个CNN分支负责提取局部冲击特征右边一个Transformer分支负责建模长程时序关系两个分支的特征在分类前拼接或加权融合。还有一种做法是串行级联先用CNN把原始信号压缩成特征序列再把特征序列送入Transformer。这种做法参数量通常更小但CNN输出的特征如果已经丢失了时序细节Transformer能补救的空间有限。我一般倾向并行分支——两个分支各自独立提取特征融合时信息互补性更强虽然参数量略大但诊断准确率和鲁棒性都更好。融合位置也很关键。常见选择是在全局池化之后拼接两个分支的特征向量再接全连接分类层。这个位置操作简单、可解释性强也方便观察两个分支各自贡献了多少。3. 搞定CWRU数据集下载、加载、切窗、划分工况3.1 CWRU数据长什么样怎么选文件CWRUCase Western Reserve University轴承数据中心提供的是.mat格式的振动信号文件采样频率主要有12kHz和48kHz两种故障类型包括滚动体故障B、内圈故障IR、外圈故障OR以及正常状态N。每类故障按损伤直径分为0.007英寸、0.014英寸、0.021英寸三档负载从0到3马力不等。实操中我的选数习惯是先固定采样频率12kHz选0.014英寸损伤直径这一档负载选0、1、2、3马力四组都用上但要注意划分方式。如果只选0负载训练、其他负载测试跨工况难度最大准确率下降也最明显如果四个负载混合打乱再划分测试集会隐含相同工况的样本结果会偏乐观。3.2 数据切窗与标签构建原始CWRU文件里每段信号长达十几万采样点不能整段送入网络。常见做法是滑窗切分每个样本是一个固定长度的窗口。我用Python的滑动窗口实现窗口长度设为1024重叠率50%import numpy as np from scipy.io import loadmat def load_cwru_mat(filepath, keyX098_DE_time): mat loadmat(filepath) sig mat[key].flatten().astype(np.float32) return sig def sliding_window(sig, win_len1024, overlap0.5): step int(win_len * (1 - overlap)) if step 1: step 1 n_windows (len(sig) - win_len) // step 1 windows np.stack([sig[i*step : i*stepwin_len] for i in range(n_windows)]) return windows这段代码先用scipy.io.loadmat读入.mat文件取出驱动端加速度计信号键名一般为X098_DE_time不同文件需按实际键名调整。滑窗时overlap0.5意味着相邻窗口有一半采样点是重叠的这样能增加样本量同时保留故障冲击的连续性。窗口长度1024在12kHz采样率下对应约0.085秒转速1797r/min时约为1.25个旋转周期Transformer分支能在这个长度内看到完整的冲击间隔。窗口太短看不到完整周期太长则单样本计算量大、训练慢。重叠率也可以调小到0.25来降低样本间的相关性减少过拟合风险。3.3 训练集、验证集、测试集怎么划分才不“穿帮”这是CWRU应用里最容易翻车的地方。常见错误是把滑窗切出来的所有样本放进一个大池子随机打乱后按比例划分训练和测试。这种做法下同一个原始信号切出的相邻窗口可能一条在训练集、一条在测试集而它们共享大量重叠采样点测试准确率会虚高到接近100%实际部署时立刻掉点。我一般用按文件划分的方式也就是“同一段连续信号的窗口只能进同一个集合”from sklearn.model_selection import train_test_split # segments是每个原始文件的编号feature_matrix是所有窗口样本 train_idx, temp_idx train_test_split( np.arange(len(segments)), test_size0.3, stratifylabels, random_state42 ) val_idx, test_idx train_test_split( temp_idx, test_size0.5, stratifylabels[temp_idx], random_state42 )注意这里train_test_split的第一个参数是文件编号列表而不是样本个数这样划分出来的是文件级别的索引。stratifylabels保证每个文件里各类故障的比例大致相同避免某个类别全被分到测试集。random_state固定下来保证实验结果可复现。如果不打算做跨工况验证而只是想快速跑通模型那么混合划分也能用但论文或报告里必须如实说明划分方式。跨工况才是这个融合方法的真正用武之地。4. 搭建CNN与Transformer融合模型结构设计、PyTorch实现与关键参数4.1 整体结构与分支设计融合模型整体分三段输入端归一化与升维中间双分支特征提取末端特征融合与分类。输入是一维振动信号长度1024通道数1。CNN分支采用三层一维卷积每层后面接批归一化和ReLU激活最后一层做全局平均池化输出128维向量。Transformer分支先把输入信号通过一个线性投影映射成序列嵌入然后经过两层标准TransformerEncoderLayer再做全局平均池化输出128维向量。两个分支的向量拼接成256维送进全连接分类层。选择128维作为分支输出维度是平衡考虑CWRU分类任务类别数在4到10之间128维特征已经足够表达再大参数量增长明显但准确率提升有限。4.2 完整PyTorch代码import torch import torch.nn as nn import torch.nn.functional as F class CNNBranch(nn.Module): def __init__(self, in_ch1, out_dim128): super().__init__() self.conv1 nn.Conv1d(in_ch, 32, kernel_size8, stride2, padding4) self.bn1 nn.BatchNorm1d(32) self.conv2 nn.Conv1d(32, 64, kernel_size8, stride2, padding4) self.bn2 nn.BatchNorm1d(64) self.conv3 nn.Conv1d(64, 128, kernel_size8, stride2, padding4) self.bn3 nn.BatchNorm1d(128) self.pool nn.AdaptiveAvgPool1d(1) self.fc nn.Linear(128, out_dim) def forward(self, x): # x: (batch, 1, win_len) x F.relu(self.bn1(self.conv1(x))) x F.relu(self.bn2(self.conv2(x))) x F.relu(self.bn3(self.conv3(x))) x self.pool(x).squeeze(-1) return self.fc(x) class TransformerBranch(nn.Module): def __init__(self, in_len1024, d_model128, nhead4, num_layers2, dropout0.1): super().__init__() self.proj nn.Linear(1, d_model) self.pos_embed nn.Parameter(torch.randn(1, in_len, d_model) * 0.02) encoder_layer nn.TransformerEncoderLayer( d_modeld_model, nheadnhead, dim_feedforward256, dropoutdropout, batch_firstTrue ) self.encoder nn.TransformerEncoder(encoder_layer, num_layersnum_layers) self.pool nn.AdaptiveAvgPool1d(1) def forward(self, x): # x: (batch, 1, win_len) - (batch, win_len, 1) x x.transpose(1, 2) x self.proj(x) self.pos_embed x self.encoder(x) x x.transpose(1, 2) x self.pool(x).squeeze(-1) return x class CNNTransformerFusion(nn.Module): def __init__(self, num_classes4, in_len1024, fusion_dim256): super().__init__() self.cnn_branch CNNBranch(in_ch1, out_dim128) self.tf_branch TransformerBranch(in_lenin_len, d_model128, nhead4, num_layers2) self.fusion nn.Sequential( nn.Linear(fusion_dim, 128), nn.BatchNorm1d(128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, num_classes) ) def forward(self, x): cnn_feat self.cnn_branch(x) tf_feat self.tf_branch(x) feat torch.cat([cnn_feat, tf_feat], dim1) return self.fusion(feat)4.3 代码逻辑说明与关键参数解释CNN分支三个卷积层的stride都是2输入长度1024经过三次下采样变成128卷积核宽度8能覆盖局部冲击波形的细粒度形态。kernel_size8在12kHz采样率下对应0.67ms的时域宽度适合捕捉冲击的起始沿如果采样率换成48kHzkernel size要适当放大到16或32保持等效时间长度不变。Transformer分支的输入被看成1024个时间步每个时间步是一个标量振动幅值。线性投影把单通道幅值映射成128维向量加上可学习的位置编码让模型知道每个时间步在序列中的位置。nhead4表示多头注意力分成4个头每个头关注128维中的32维子空间分别建模不同类型的依赖关系。dim_feedforward256是前馈网络的隐藏层维度通常是d_model的两倍。位置编码使用的是可学习参数而非三角函数这在振动信号上我实测差别不大但可学习参数在小数据集上更容易拟合。如果你觉得训练波动大换成正弦位置编码也可以。融合部分先拼接256维特征经过一层全连接降到128维再过分类头输出类别概率。中间的Dropout设0.5目的是防止两个分支在训练中共同过拟合到训练集的噪声模式上。5. 避坑与排查训练和部署中常见的5个实际问题5.1 CWRU文件键名不一致导致程序崩溃现象用同一套读取代码处理不同.mat文件时有的文件能读有的报KeyError。原因CWRU数据文件里键名不统一有的文件是X098_DE_time有的是X109_DE_time还有的包含DE_time等不同字段。不同文件名对应不同采样频率和测试位置键名随之变化。解决读取时先打印文件的所有键名或者用glob匹配包含DE_time的键def find_key(mat): for k in mat.keys(): if DE_time in k: return k raise KeyError(未找到DE_time字段)5.2 窗口重叠率过高导致验证集虚高现象模型在验证集上准确率99%但换到另一段新采集的信号上准确率只有80%左右。原因重叠率设为50%甚至更高时相邻窗口大量采样点重复训练集和验证集之间信息重叠严重。模型实际上“背住了”重复的波形片段。解决要么把重叠率降到0到25%之间要么严格按文件级别划分数据集。按文件划分后即使重叠率高训练和验证数据也来自不同原始信号段信息泄漏会小很多。5.3 Transformer分支在小数据集上反复过拟合现象训练集loss持续下降验证集loss先降后升准确率曲线明显分离。原因CWRU单类样本一般几千条Transformer参数量相对较大2层编码器加4个注意力头单分支参数量已接近百万级别很容易记住训练集的细节。解决先增加正则化——Dropout从0.1提到0.3位置编码的初始化方差调小再考虑减小d_model到64或者降低num_layers到1。操作顺序是优先调Dropout其次是调小维度不要一上来就加数据。5.4 多分类任务里外圈故障的“一点钟方向”类别总混淆现象外圈故障按损伤位置分为3点钟、6点钟、12点钟方向模型总把其中某一类判成另一类尤其集中在相近方向。原因不同方向的外圈故障在振动信号上的冲击形态相似负载变化后传播路径改变部分类别特征重叠严重。这不是融合模型特有的问题纯CNN和纯Transformer都会遇到。解决一种做法是把三个外圈方向合并成一类“外圈故障”类别数从10降到8工程上够用另一种做法是保持细分类但训练时给混淆严重的类别加大权重或者用focal loss替代交叉熵损失。5.5 训练速度比纯CNN慢很多现象同样的CWRU数据纯CNN训练一轮只要1分钟融合模型要5分钟以上。原因Transformer分支的自注意力计算复杂度是序列长度的平方1024个时间步意味着每层约100万次注意力计算计算开销远大于CNN卷积。解决不要把整段1024点都送进Transformer。可以先用CNN分支的中间层输出做降采样比如把1024点压缩到256个特征向量再送入Transformer分支。这个做法本质是前文说的串行变体速度提升明显但需要重新实验确定最佳压缩比例。6. 验证融合模型的真实能力跨工况测试与特征可视化技巧融合模型到底比单分支好多少不能只看CWRU随机划分的准确率更可靠的做法是跨工况验证。我常用的方案是用0负载数据训练分别用1、2、3负载数据做测试。这个设置模拟的是“现场工况与标定工况不一致”的真实场景。在这种验证下纯CNN的准确率通常会下降3到8个百分点纯Transformer下降更多融合模型的下降幅度往往最小。原因就是CNN分支保住了局部冲击特征Transformer分支补偿了转频变化带来的时序间隔变化。如果融合模型在跨工况测试中没有表现出优势先回去检查两个分支的特征是否真的都被用上了——把融合层的权重打印出来看两个分支的贡献占比经常发现某一分支权重接近0这时需要调整分支输出维度的比例或者改用加权融合与可学习的融合系数。进阶验证手法是用t-SNE可视化融合前两个分支的特征分布。如果两个分支在t-SNE图上各自形成了不同的聚类结构说明它们的特征确实存在互补性如果两个分支的特征分布几乎重合融合意义就不大模型等价于一个参数量翻倍的CNN。最后说一个我自己的习惯每次训练完都保存最佳模型权重和对应的数据划分索引方便复现。CWRU上的结果再漂亮也只能说明公开数据上有效真正要投入现场必须在自采数据上重新做跨工况验证。这个融合思路的落地价值在于提供了两个互补的视角而不是一个万能的模型结构。希望这些实操步骤和踩坑记录能帮你在自己的诊断任务上少走几步弯路。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。