
简介本资源是一份面向深度学习初学者与工业智能诊断研究者的实践型学习材料聚焦轴承故障智能识别这一典型工业AI应用场景。方案融合CNN局部特征提取能力与Transformer长程依赖建模优势基于CWRU标准化振动数据集构建端到端诊断模型适用于设备预测性维护、状态评估与异常早期预警等实际运维需求。压缩包共27个文件含10个MATLAB格式原始振动信号覆盖多种故障类型、3个核心Python模型脚本cnn_transformer.py及.ipynb、1个Keras训练权重、2个README说明文档、1个CSV数据集、1个run.bat启动脚本及配套工具与环境配置文件整体大小为16.65MB结构清晰便于复现与二次开发。已有84人学习下载读者可直接运行代码完成数据加载、模型训练与故障分类全流程获取完整可执行的融合架构实现、CWRU数据预处理逻辑、注意力机制在时序信号中的应用范式及典型排错提示。1. 为什么单用CNN或Transformer做轴承故障诊断总差一口气——CWRU数据集上融合模型的真实瓶颈与破局点你在CWRUCase Western Reserve University轴承数据集上跑过ResNet、ViT或者LSTM吗大概率见过这种现象CNN在时频图上识别内圈故障准确率92%但一换到外圈故障就掉到78%Transformer对多工况迁移稍好可训练100轮后验证loss突然震荡GPU显存还爆得莫名其妙。这不是你调参不行而是单一架构天然存在感知盲区——CNN擅长局部纹理但抓不住长程振动相位关系Transformer能建模全局依赖却对微弱冲击脉冲极度不敏感。本方案不是简单拼接两个模型而是把CNN提取的时频局部特征图作为Transformer Encoder的结构化token输入再用位置编码注入采样时间戳物理意义让模型既“看得清”冲击包络形态又“想得透”故障演化时序逻辑。适合已跑通CWRU基础分类如0-9类故障、正卡在跨工况泛化或小样本识别瓶颈的工程师尤其当你手头只有1000条标注样本却要覆盖4种负载3种转速组合时这套融合路径比堆数据或换更大模型更可控、更可解释。2. 从原始振动信号到结构化TokenCWRU数据预处理与特征嵌入的三道硬关CWRU官网下载的.mat文件看似规整但直接喂给CNN-Transformer混合模型会翻车。关键不在模型结构而在信号→图像→token这条链路上的三个隐性陷阱采样率不一致导致时频图失真、标签分布偏斜引发梯度淹没、以及Transformer对token序列长度极度敏感。下面拆解每一步的实操要点附可直接运行的代码块和参数依据。2.1 振动信号重采样与分段为什么必须用2048点/段而非默认512点CWRU原始数据采样率有12k、48k等多档直接FFT会导致频谱分辨率差异巨大。我们统一重采样至12kHz理由覆盖轴承故障特征频率fBPFO≈160Hz的5倍以上满足奈奎斯特准则再按2048点/段切分非常见512点。原因在于Transformer的token序列长度直接影响计算复杂度O(n²)2048点经STFT生成128×128时频图再经CNN下采样得到16×16特征图最终展平为256个token——这个长度在显存占用8GB和建模能力间取得平衡。若用512点切分token数仅64个长程依赖建模能力断崖式下降。import scipy.io as sio import numpy as np from scipy.signal import resample def load_and_resample(mat_path, target_fs12000): data sio.loadmat(mat_path) raw_signal data[X097_DE_time].flatten() # 以DE端为例 # 计算原始采样率CWRU各文件不同需动态获取 original_fs int(1 / (data[time][0,1] - data[time][0,0])) # 重采样至target_fs n_samples_new int(len(raw_signal) * target_fs / original_fs) resampled resample(raw_signal, n_samples_new) return resampled # 示例加载97号工况DE端数据 signal_97 load_and_resample(X097_DE_time.mat) print(f原始长度: {len(signal_97)}, 重采样后: {len(signal_97)}) # 验证是否成功提示resample()函数会自动补零或截断务必检查输出长度是否符合预期。若原始信号含大量静默段如停机间隙需先用能量阈值法裁剪否则STFT生成的时频图会出现大面积黑色噪声带CNN误学为“正常状态”。2.2 STFT参数精调窗长、重叠率与归一化的物理意义CWRU振动信号信噪比低典型SNR≈6dBSTFT参数选错会让故障冲击被淹没。我们采用汉宁窗长128点、重叠率75%、FFT点数256——这组参数经实验验证窗长128对应约10ms时间窗12kHz下足以包裹单次冲击衰减过程75%重叠保证相邻帧间有强相关性避免冲击脉冲被切碎256点FFT提供足够频域分辨率Δf12000/256≈47Hz能清晰分离BPFI/BPFO频带。import matplotlib.pyplot as plt from scipy.signal import stft def generate_spectrogram(signal, fs12000, nperseg128, noverlap96, nfft256): f, t, Zxx stft(signal, fsfs, npersegnperseg, noverlapnoverlap, nfftnfft, windowhann, scalingspectrum) # 取幅度谱并归一化到[0,1] mag_spec np.abs(Zxx) mag_spec (mag_spec - mag_spec.min()) / (mag_spec.max() - mag_spec.min() 1e-8) return mag_spec, f, t # 生成一张时频图示例 spec_97, freqs, times generate_spectrogram(signal_97[:2048*10]) # 取前10段 plt.imshow(spec_97, aspectauto, cmapjet, originlower) plt.title(CWRU X097 DE端 STFT时频图128点汉宁窗) plt.ylabel(Frequency (Hz)) plt.xlabel(Time frame) plt.show()参数说明noverlap96即75%重叠128×0.75scalingspectrum确保能量守恒归一化用(x-min)/(max-min)而非z-score因故障冲击在时频图中表现为局部高亮斑点min-max更能保留对比度。2.3 标签映射与平衡策略CWRU的9类故障如何避免“正常样本霸权”CWRU标准划分含10类0正常、1-3内圈故障不同尺寸、4-6滚动体故障、7-9外圈故障。但原始数据中正常样本占比超40%而外圈故障样本不足8%。若直接按原始比例划分训练集模型会倾向预测“正常”。我们采用分层欠采样SMOTE插值组合对正常类随机丢弃至与其他类均等每类取1200样本再对少样本类如外圈故障用SMOTE在时频图特征空间生成新样本非原始信号空间。from imblearn.over_sampling import SMOTE from sklearn.model_selection import train_test_split # 假设X_all是所有时频图堆叠成的(N, 128, 128)数组y_all是对应标签 X_train, X_test, y_train, y_test train_test_split( X_all, y_all, test_size0.2, stratifyy_all, random_state42 ) # 对训练集进行平衡 smote SMOTE(random_state42, k_neighbors3) # k_neighbors3避免过拟合 X_train_balanced, y_train_balanced smote.fit_resample( X_train.reshape(X_train.shape[0], -1), y_train ) X_train_balanced X_train_balanced.reshape(-1, 128, 128) # 恢复形状 print(f平衡前训练集标签分布: {np.bincount(y_train)}) print(f平衡后训练集标签分布: {np.bincount(y_train_balanced)})注意SMOTE必须作用于展平后的时频图特征N×16384维而非原始信号或未归一化的STFT结果。k_neighbors3是经验值——过大如5会在噪声区域生成虚假冲击斑点过小如1导致插值样本过于相似。3. CNN-Transformer融合架构设计不是简单拼接而是特征流的物理对齐很多论文把CNN输出直接Flatten送进Transformer结果精度还不如单CNN。问题出在特征语义断裂CNN最后一层输出的是高维抽象特征如512维而Transformer期望的是具有空间/时序结构的token序列。我们的方案强制让CNN输出与STFT网格对齐并注入物理位置信息使Transformer真正理解“左上角token对应高频早期冲击”。3.1 CNN分支用轻量级ResNet18替代VGG为何去掉全连接层CWRU时频图尺寸为128×128VGG16在此尺寸上参数量超130M且深层卷积易丢失冲击细节。我们改用ResNet18精简版移除最后的Global Average Pooling和全连接层保留4个stage的输出特征图尺寸依次为64×64、32×32、16×16、8×8。关键改动是将Stage3输出16×16×256作为主token源——该尺度既能保留足够空间细节16×16对应STFT的16个时间帧×16个频带又避免Stage48×8过度压缩导致冲击定位模糊。import torch import torch.nn as nn from torchvision.models import resnet18 class CNNBackbone(nn.Module): def __init__(self, pretrainedFalse): super().__init__() self.resnet resnet18(pretrainedpretrained) # 移除最后的avgpool和fc层 self.resnet.avgpool nn.Identity() self.resnet.fc nn.Identity() def forward(self, x): # x shape: (B, 1, 128, 128) —— 注意加通道维度 x x.unsqueeze(1) # 转为(B, 1, 128, 128) # 经过ResNet前3个stage x self.resnet.conv1(x) x self.resnet.bn1(x) x self.resnet.relu(x) x self.resnet.maxpool(x) x self.resnet.layer1(x) # 64x64x64 x self.resnet.layer2(x) # 32x32x128 x self.resnet.layer3(x) # 16x16x256 ← 关键输出 return x # 实例化并测试输出形状 cnn CNNBackbone() dummy_input torch.randn(4, 128, 128) # batch4 output cnn(dummy_input) print(fCNN输出形状: {output.shape}) # 应为 torch.Size([4, 256, 16, 16])逻辑说明unsqueeze(1)将单通道时频图转为(B, C1, H, W)适配ResNet输入layer3输出256通道×16×16后续直接展平为256个token每个token含256维特征——这比直接Flatten整个特征图16×16×25665536维更符合Transformer的token设计哲学。3.2 Token化与位置编码给每个token打上“时空坐标戳”Transformer的Positional Encoding通常用正弦函数但CWRU时频图的横纵轴有明确物理意义横轴是时间帧索引0~15纵轴是频带索引0~15。我们设计二维位置编码将每个token的位置(i,j)映射为可学习向量而非固定正弦波class PositionEmbedding2D(nn.Module): def __init__(self, num_embeddings256, embedding_dim256): super().__init__() # 创建16x16网格的位置索引 pos_h torch.arange(16).unsqueeze(1) # (16, 1) pos_w torch.arange(16).unsqueeze(0) # (1, 16) pos_grid torch.stack(torch.meshgrid(pos_h.squeeze(), pos_w.squeeze()), dim-1) # (16,16,2) # 将二维位置展平为256个token的位置索引 self.pos_embedding nn.Embedding(256, embedding_dim) # 初始化用正弦编码预热但允许梯度更新 self.pos_embedding.weight.data self._get_sin_pos_embed(embedding_dim, 16, 16) def _get_sin_pos_embed(self, dim, h, w): pos torch.zeros(h*w, dim) for i in range(h): for j in range(w): idx i * w j for k in range(0, dim, 2): pos[idx, k] np.sin(i / (10000 ** ((k)/dim))) if k1 dim: pos[idx, k1] np.cos(j / (10000 ** ((k)/dim))) return pos def forward(self, x): # x shape: (B, 256, 256) —— B个样本256个token每个256维 pos_idx torch.arange(256, devicex.device) pos_emb self.pos_embedding(pos_idx) # (256, 256) return x pos_emb.unsqueeze(0) # (1, 256, 256) → 广播加到x上 # 测试位置编码 pos_enc PositionEmbedding2D() dummy_tokens torch.randn(2, 256, 256) # B2, 256 tokens, dim256 encoded pos_enc(dummy_tokens) print(f位置编码后形状: {encoded.shape}) # (2, 256, 256)参数说明embedding_dim256与CNN输出通道数一致确保可直接相加_get_sin_pos_embed生成初始正弦编码但nn.Embedding允许反向传播更新使模型能自适应CWRU数据的特定时空模式。3.3 Transformer Encoder配置层数、头数与FFN隐藏层的实测阈值CWRU样本量有限平衡后约1.2万张图Transformer层数过多会过拟合。经消融实验4层Encoder 4个注意力头 FFN隐藏层512维为最优组合层数4时长程建模不足4则验证loss波动剧烈头数4在256维token下能充分分割特征子空间头数8导致每头仅32维注意力权重分散FFN隐藏层设为5122×token dim是经典比例设为1024时显存暴涨且精度无提升。from torch.nn import TransformerEncoder, TransformerEncoderLayer def build_transformer_encoder(d_model256, nhead4, num_layers4, dim_feedforward512): encoder_layer TransformerEncoderLayer( d_modeld_model, nheadnhead, dim_feedforwarddim_feedforward, dropout0.1, # CWRU噪声大dropout需略高于NLP任务 activationgelu, # GELU比ReLU更适配振动信号非线性 batch_firstTrue ) transformer_encoder TransformerEncoder(encoder_layer, num_layersnum_layers) return transformer_encoder transformer build_transformer_encoder() # 输入(B, 256, 256) —— B个样本256个token每个256维 dummy_tokens torch.randn(2, 256, 256) output transformer(dummy_tokens) print(fTransformer输出形状: {output.shape}) # (2, 256, 256)关键设置dropout0.1防止过拟合activationgelu因振动信号含丰富谐波分量GELU的平滑非线性比ReLU更适配batch_firstTrue避免维度转换开销。4. 训练与推理全流程从损失函数选择到跨工况验证的落地细节模型结构搭好只是开始CWRU实战中最耗时的环节在训练策略和验证设计。我们不用交叉验证太慢也不用ImageNet预训练领域偏差大而是构建工况感知的训练-验证-测试闭环确保模型真正学会故障物理本质而非记忆数据集ID。4.1 损失函数与优化器Label Smoothing为何比CrossEntropy更稳CWRU故障类别间存在混淆如内圈小故障与正常状态频谱相似直接CrossEntropy会使模型对错误标签过度自信。我们采用Label Smoothingε0.1将真实标签概率从1降为0.9其余类均分0.1迫使模型输出更平滑的概率分布。import torch.nn as nn class LabelSmoothingLoss(nn.Module): def __init__(self, classes10, smoothing0.1, dim-1): super().__init__() self.confidence 1.0 - smoothing self.smoothing smoothing self.cls classes self.dim dim def forward(self, pred, target): pred pred.log_softmax(dimself.dim) with torch.no_grad(): true_dist torch.zeros_like(pred) true_dist.fill_(self.smoothing / (self.cls - 1)) true_dist.scatter_(1, target.data.unsqueeze(1), self.confidence) return torch.mean(torch.sum(-true_dist * pred, dimself.dim)) # 使用示例 criterion LabelSmoothingLoss(classes10, smoothing0.1) optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100)参数依据smoothing0.1经网格搜索确定——0.05时模型仍易过拟合0.15时收敛变慢且最终精度下降weight_decay0.05高于常规0.01因CWRU特征维度高256需更强正则化。4.2 工况划分验证法为什么不能只用随机8:2划分CWRU包含4种负载0、1、2、3HP和3种转速1797、1772、1750、1730rpm共12种工况组合。若随机划分训练集可能含1797rpm而测试集全是1730rpm模型实际是做转速识别而非故障诊断。我们严格按工况隔离选定1797rpm-0HP、1772rpm-1HP、1750rpm-2HP为训练工况其余9种为测试工况。这样验证的是模型真正的泛化能力。# 假设df_cwru包含列file_name,fault_type,load_hp,rpm train_conditions [(0, 1797), (1, 1772), (2, 1750)] test_conditions [(hp, rpm) for hp in [0,1,2,3] for rpm in [1797,1772,1750,1730] if (hp, rpm) not in train_conditions] train_mask df_cwru.apply(lambda row: (row[load_hp], row[rpm]) in train_conditions, axis1) test_mask df_cwru.apply(lambda row: (row[load_hp], row[rpm]) in test_conditions, axis1) X_train, y_train X_all[train_mask], y_all[train_mask] X_test, y_test X_all[test_mask], y_all[test_mask] print(f训练工况数: {len(train_conditions)}, 测试工况数: {len(test_conditions)}) print(f训练样本数: {len(X_train)}, 测试样本数: {len(X_test)})注意此划分下测试集样本量必然少于训练集需在评估时用宏平均F1-score而非accuracy因各类别样本数不均衡。4.3 推理加速技巧ONNX导出与TensorRT部署的关键参数训练好的PyTorch模型直接部署到边缘设备如Jetson AGX会很慢。我们通过ONNX导出TensorRT优化将单图推理时间从120ms降至18ms。核心是静态shape与FP16精度# 导出ONNXPyTorch端 model.eval() dummy_input torch.randn(1, 128, 128) # 单样本 torch.onnx.export( model, dummy_input, cnn_transformer_cwru.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}}, opset_version13, verboseFalse ) # TensorRT优化需在TRT环境执行 import tensorrt as trt def build_engine(onnx_path, engine_path, fp16_modeTrue): logger trt.Logger(trt.Logger.WARNING) builder trt.Builder(logger) network builder.create_network(1 int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser trt.OnnxParser(network, logger) with open(onnx_path, rb) as f: parser.parse(f.read()) config builder.create_builder_config() config.max_workspace_size 1 30 # 1GB if fp16_mode: config.set_flag(trt.BuilderFlag.FP16) engine builder.build_engine(network, config) with open(engine_path, wb) as f: f.write(engine.serialize()) return engine # 执行构建 build_engine(cnn_transformer_cwru.onnx, cnn_transformer_cwru.engine, fp16_modeTrue)避坑提示opset_version13必须匹配PyTorch版本1.10否则ONNX解析失败FP16开启后需验证精度损失0.5%否则回退到FP32。5. 避坑指南CWRU融合模型训练中踩过的5个真实血泪坑这些坑不是理论推演而是我在3台不同配置GPURTX3090/4090/A100上反复调试、记录日志、逐行对比输出后确认的硬伤。跳过任一个都可能导致训练loss不降、验证acc震荡或部署后结果错乱。5.1 现象训练初期loss下降极慢10轮后才开始明显收敛原因CNN分支的BatchNorm层在小batch16下统计量不准导致特征分布偏移Transformer输入不稳定。CWRU平衡后每类1200样本若batch8则每个epoch仅150步BN统计量严重失真。解决改用SyncBatchNorm多卡时或GroupNorm单卡后者将通道分组归一化不受batch size影响。代码中替换nn.BatchNorm2d为nn.GroupNorm(num_groups32, num_channels256)。5.2 现象验证集F1-score在第40轮突降15%之后持续波动原因学习率调度器CosineAnnealingLR的T_max设为100但模型在40轮已收敛余下60轮在低lr下微调反而破坏已学特征。CWRU数据噪声大过长微调易过拟合。解决改用ReduceLROnPlateau监控验证F1连续5轮不升则lr×0.5最多降2次。同时设置早停patience10避免无效训练。5.3 现象跨工况测试时外圈故障识别率仅62%远低于内圈的89%原因STFT参数对高频分量敏感而外圈故障特征频率BPFO低于内圈BPFI原STFT窗长128点10ms对低频周期捕捉不足。解决对外圈故障样本单独使用**窗长256点、重叠率50%**的STFT保持其他参数不变。在数据加载时根据标签动态切换STFT参数而非统一处理。5.4 现象TensorRT推理结果与PyTorch差异达12%且类别混乱原因ONNX导出时未冻结BN层model.eval()后需model.train(False)确保BN参数固定导致TRT读取的running_mean/var与训练时不一致。解决导出前显式调用model.apply(lambda m: setattr(m, training, False) if hasattr(m, training) else None)并验证ONNX输出与PyTorch一致。5.5 现象模型在CWRU上效果好但换到自家产线数据同轴承型号准确率暴跌至55%原因CWRU数据采集环境理想实验室振动台而产线存在电机电磁干扰、安装螺栓松动等未建模因素导致时频图出现非故障相关条纹。解决在STFT后增加**自适应谱减法Adaptive Spectral Subtraction**预处理用滑动窗口估计噪声谱再从时频图中减去——此步骤不改变标签但显著提升跨域鲁棒性。6. 故障可解释性增强用Attention Map反向定位冲击源不只是“黑匣子判别”模型给出92%准确率还不够产线工程师需要知道“为什么判为外圈故障”。我们不满足于Grad-CAM这类通用可视化而是设计故障导向的Attention Rollout将Transformer最后一层的注意力权重沿token空间反向传播至原始时频图精准标出模型决策依据的时空区域。6.1 Attention Rollout原理为什么比Grad-CAM更适合振动信号Grad-CAM基于CNN梯度但CNN-Transformer融合模型中最终决策由Transformer的全局注意力聚合CNN梯度无法反映跨token关联。Attention Rollout则利用Transformer各层注意力矩阵A(l)∈ℝn×nn256递归计算累积注意力$$ R^{(L)} A^{(L)}, \quad R^{(l)} A^{(l)} \cdot R^{(l1)} $$其中R(0)即为最终token重要性展平后映射回16×16网格再双线性插值回128×128时频图。def rollout_attention(attentions): # attentions: list of (B, n_heads, n_tokens, n_tokens), len4 # 取最后一层平均注意力 last_att attentions[-1].mean(dim1) # (B, n_tokens, n_tokens) # 初始化累积注意力 residual_att torch.eye(last_att.shape[1], devicelast_att.device) for att in reversed(attentions): att_mean att.mean(dim1) # (B, n_tokens, n_tokens) residual_att torch.matmul(att_mean, residual_att) # 取cls token我们用mean pooling故取所有token平均 rollout residual_att.mean(dim0) # (n_tokens, n_tokens) # 取每行平均得每个token重要性 token_importance rollout.mean(dim1) # (n_tokens,) return token_importance # 在推理时hook注意力 attentions [] def hook_fn(module, input, output): attentions.append(output[1]) # output[1] is attention weights # 注册hook到Transformer Encoder Layer for layer in transformer.layers: layer.self_attn.register_forward_hook(hook_fn) # 推理后获取rollout token_imp rollout_attention(attentions) # 映射回16x16网格 imp_grid token_imp.reshape(16, 16) # 插值回128x128 import torch.nn.functional as F imp_map F.interpolate(imp_grid.unsqueeze(0).unsqueeze(0), size(128,128), modebilinear)[0,0]效果验证在CWRU外圈故障样本上Attention Rollout高亮区域与理论BPFO频带~160Hz及冲击周期~12ms完全吻合误差2个像素——这意味着工程师能直接看到模型关注的是不是真实的故障特征。6.2 故障模式聚类用Transformer最后一层token聚类发现未知故障CWRU只定义10类但产线可能出现未标注的复合故障如内圈润滑不良。我们提取Transformer Encoder输出的**[CLS] token**实际用mean pooling在1000个测试样本上做UMAP降维HDBSCAN聚类聚类ID样本数主要标签物理特征0320正常低频平稳无冲击1210内圈故障高频冲击周期稳定2185外圈故障中频调制包络周期长342未知高频冲击宽频噪声疑似润滑失效操作步骤提取所有测试样本的output.mean(dim1)256维umap.UMAP(n_components2, n_neighbors15).fit_transform(features)hdbscan.HDBSCAN(min_cluster_size30).fit(embedding)发现聚类3中87%样本来自同一台电机且振动加速度RMS值异常高——现场检查确认为轴承缺油。我坚持在每次模型上线前跑一遍Attention Rollout和聚类分析不是为了发论文而是当产线老师傅指着屏幕问“这红点为啥在这”时我能指着时频图说“看这里每12ms一次冲击频率162Hz正是外圈故障特征频率BPFO。”——技术的价值不在指标数字而在让不可见的故障变得可见、可对话、可行动。希望帮到你。本文还有配套的精品资源点击获取