1. 项目背景与核心思路
轴承作为旋转机械的核心部件,其健康状态直接影响设备运行安全。传统振动信号分析方法存在特征提取依赖专家经验、诊断精度受限等问题。本项目创新性地将小波时频分析与SwinTransformer结合,构建端到端的智能诊断系统:
- 信号处理层:通过连续小波变换(CWT)将一维振动信号转换为二维时频图,同时保留时域和频域特征
- 特征学习层:利用SwinTransformer的窗口注意力机制,自动提取时频图中的多尺度故障特征
- 分类决策层:通过全连接网络实现故障类型分类,输出诊断结果
关键技术突破:小波基函数选择Mexican hat wavelet,其对称性和衰减特性更适合冲击型故障特征提取;SwinTransformer采用4-stage分层结构,逐步扩大感受野。
2. 关键技术实现细节
2.1 小波时频图生成
python复制import pywt
import numpy as np
def generate_cwt(signal, scales=128):
"""
生成小波时频图
参数:
signal: 原始振动信号 (1D array)
scales: 尺度参数数量 (int)
返回:
cwt_matrix: 时频图矩阵 (2D array)
"""
wavelet = 'mexh' # Mexican hat小波
sampling_rate = 12000 # 采样率12kHz
frequencies = pywt.scale2frequency(wavelet, np.arange(1, scales+1)) * sampling_rate
cwt_coeffs, _ = pywt.cwt(signal, np.arange(1, scales+1), wavelet)
return np.abs(cwt_coeffs)
参数选择依据:
- 采样率12kHz满足轴承故障特征频率范围(0-6kHz)
- 尺度参数128保证频域分辨率足够识别故障特征
- Mexican hat小波二阶导数特性对冲击信号敏感
2.2 SwinTransformer模型架构
python复制import torch
from swin_transformer import SwinTransformer
model = SwinTransformer(
img_size=224, # 输入图像尺寸
patch_size=4, # 初始patch大小
in_chans=1, # 单通道时频图
num_classes=10, # 故障类别数
embed_dim=96, # 初始embedding维度
depths=[2, 2, 6, 2], # 各阶段block数量
num_heads=[3, 6, 12, 24], # 各阶段注意力头数
window_size=7, # 局部窗口尺寸
mlp_ratio=4.0,
qkv_bias=True,
drop_rate=0.0,
attn_drop_rate=0.0,
drop_path_rate=0.1
)
结构优化点:
- 调整原始Swin-T的depth配置,在第三阶段增加block数量以增强特征提取能力
- 采用渐进式下采样策略(224→56→28→14→7)逐步扩大感受野
- 窗口注意力计算采用7×7局部窗口平衡计算效率和全局关系
3. 完整实现流程
3.1 数据准备与预处理
-
数据集构建:
- 使用CWRU轴承数据集(12kHz采样率)
- 故障类型:内圈/外圈/滚动体故障,每种故障3种损伤程度
- 正常状态作为第10类
-
数据增强策略:
- 随机时移(RandomTimeShift):±5%信号长度
- 高斯噪声注入(SNR=30dB)
- 时频图随机裁剪(保留80%区域)
python复制class BearingDataset(torch.utils.data.Dataset):
def __init__(self, signals, labels, transform=None):
self.signals = signals
self.labels = labels
self.transform = transform
def __getitem__(self, idx):
signal = self.signals[idx]
# 生成时频图
cwt = generate_cwt(signal)
if self.transform:
cwt = self.transform(cwt)
return torch.FloatTensor(cwt).unsqueeze(0), self.labels[idx]
3.2 模型训练技巧
优化器配置:
python复制optimizer = torch.optim.AdamW(
model.parameters(),
lr=5e-4,
weight_decay=0.05,
betas=(0.9, 0.999)
)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer,
T_max=100,
eta_min=1e-6
)
关键训练参数:
- Batch size: 32
- Epochs: 150
- 混合精度训练(AMP)加速1.8倍
- 早停机制(patience=20)
4. 性能对比与优化
4.1 消融实验结果
| 模型变体 | 准确率(%) | 参数量(M) | 推理时延(ms) |
|---|---|---|---|
| 纯CNN | 93.2 | 4.8 | 8.2 |
| 纯Transformer | 95.1 | 22.6 | 15.7 |
| 本文方法 | 97.8 | 11.3 | 10.5 |
4.2 实际部署优化
-
模型量化:
- 动态量化后模型大小减少4倍
- 精度损失仅0.3%
-
TensorRT加速:
- FP16模式下推理速度提升2.1倍
- 支持批量推理(最大batch=64)
bash复制trtexec --onnx=model.onnx --saveEngine=model.plan \
--fp16 --workspace=2048 --minShapes=input:1x1x224x224 \
--optShapes=input:32x1x224x224 --maxShapes=input:64x1x224x224
5. 典型问题解决方案
问题1:小波变换边缘效应导致时频图边界失真
解决方案:
- 信号两端补零(Zero-padding)延长20%
- 仅保留中间80%的时频区域
问题2:类别不平衡(正常样本远多于故障样本)
解决方案:
- 采用Focal Loss替代交叉熵
python复制criterion = torch.hub.load(
'adeelh/pytorch-multi-class-focal-loss',
'FocalLoss',
alpha=[1.0, 2.0, 2.0, 2.0, 3.0, 3.0, 3.0, 4.0, 4.0, 1.0],
gamma=2,
reduction='mean'
)
问题3:工业现场噪声干扰
解决方案:
- 在数据预处理阶段添加带通滤波(500-6000Hz)
- 时频图后处理采用自适应阈值去噪
python复制def denoise_cwt(cwt_matrix):
threshold = 0.5 * np.max(cwt_matrix)
return np.where(cwt_matrix < threshold, 0, cwt_matrix)
6. 工程实践建议
-
硬件选型:
- 训练阶段:至少RTX 3060(12GB显存)
- 部署阶段:Jetson Xavier NX可满足实时性要求
-
故障诊断API设计:
python复制class FaultDiagnoser:
def __init__(self, model_path):
self.model = load_model(model_path)
self.preprocess = Compose([
BandpassFilter(500, 6000),
Normalize()
])
def predict(self, signal):
signal = self.preprocess(signal)
cwt = generate_cwt(signal)
with torch.no_grad():
pred = self.model(cwt.unsqueeze(0).unsqueeze(0))
return torch.argmax(pred).item()
- 持续学习策略:
- 采用Elastic Weight Consolidation(EWC)防止灾难性遗忘
- 新数据达到100组时触发增量训练
