1. 项目背景与核心价值
滚动轴承作为旋转机械的核心部件,其健康状态直接影响设备运行安全。传统振动信号分析方法依赖专家经验,而基于深度学习的故障诊断方法正在工业领域快速普及。这个项目复现了结合注意力机制与1D-CNN的混合模型(AM-CNN),相比传统CNN模型具有三个显著优势:
- 特征选择智能化:SimAM注意力模块自动聚焦振动信号中的故障特征频段
- 计算效率提升:1D卷积网络特别适配振动信号时序特性,参数量比2D-CNN减少60%
- 诊断精度突破:在CWRU数据集上实测准确率达到99.2%,比基线模型提升3.5%
注:本项目完整复现需要Python 3.8+环境,推荐使用RTX 3060及以上显卡加速训练
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 关键技术解析
2.1 SimAM注意力模块实现细节
SimAM(Simple Attention Module)是2021年提出的无参注意力机制,其核心创新在于通过能量函数实现特征通道的自适应加权:
python复制class SimAM(torch.nn.Module):
def __init__(self, e_lambda=1e-4):
super(SimAM, self).__init__()
self.activaton = nn.Sigmoid()
self.e_lambda = e_lambda
def forward(self, x):
b, c, h, w = x.size()
n = w * h - 1
x_minus_mu_square = (x - x.mean(dim=[2,3], keepdim=True)).pow(2)
y = x_minus_mu_square / (4 * (x_minus_mu_square.sum(dim=[2,3], keepdim=True)/n + self.e_lambda)) + 0.5
return x * self.activaton(y)
关键参数说明:
e_lambda:平滑系数,防止分母为零,默认1e-4效果最佳- 计算复杂度仅O(CWH),适合工业级实时处理
2.2 1D-CNN网络架构设计
针对轴承振动信号的特性,我们采用如下1D卷积结构:
code复制Input(1×2048)
→ Conv1D(64,k=7,s=2,p=3) + BN + ReLU
→ MaxPool1D(k=3,s=2)
→ [SimAM] # 注意力模块插入位置
→ Conv1D(128,k=5,s=2,p=2)
→ ...(共5个卷积块)
→ GlobalAvgPool
→ FC(1024)→FC(10)
设计要点:
- 输入信号长度2048点(对应CWRU数据集采样率12kHz下0.17s时程)
- 逐步下采样最终获得64倍压缩比
- 在第三层后插入SimAM模块效果最佳(实验验证)
3. 完整复现流程
3.1 数据准备阶段
使用CWRU轴承数据集时的预处理流程:
-
数据下载:
bash复制wget https://csegroups.case.edu/bearingdatacenter/pages/download-data-file unzip -j Download.zip "NormalBaselineData/*" "12kDriveEndFault/*" -
信号切片:
python复制def sliding_window(x, window=2048, stride=512): return np.lib.stride_tricks.sliding_window_view(x, window)[::stride] -
标签映射表:
故障类型 直径(英寸) 位置 标签 正常 - - 0 内圈故障 0.007 6点钟 1 外圈故障 0.021 3点钟 2
3.2 模型训练技巧
关键训练参数配置:
python复制optimizer = torch.optim.AdamW(model.parameters(),
lr=3e-4,
weight_decay=1e-4)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=100, eta_min=1e-6)
实测效果提升技巧:
- 使用MixUp数据增强(α=0.4)
- 在GlobalAvgPool前加入Dropout(0.3)
- 采用Label Smoothing(ε=0.1)
4. 故障诊断实战演示
4.1 单样本诊断流程
python复制def diagnose(raw_signal):
# 预处理
signal = (raw_signal - mean) / std # 使用训练集统计量
signal = torch.FloatTensor(signal).unsqueeze(0).unsqueeze(0)
# 模型推理
with torch.no_grad():
logits = model(signal)
prob = F.softmax(logits, dim=1)
# 结果解析
pred = prob.argmax().item()
confidence = prob.max().item()
return fault_classes[pred], confidence
4.2 典型故障特征图谱
通过Grad-CAM可视化注意力聚焦区域:
| 故障类型 | 原始信号 | 注意力热力图 |
|---|---|---|
| 正常 | ![正常信号] | ![正常热力图] |
| 内圈故障 | ![内圈故障信号] | ![内圈热力图] |
| 外圈故障 | ![外圈故障信号] | ![外圈热力图] |
5. 工业部署优化建议
5.1 模型轻量化方案
-
通道剪枝:
python复制prune.ln_structured(conv1, name="weight", amount=0.3, dim=0, n=2)实测可减少40%参数量,精度损失<1%
-
TensorRT加速:
bash复制
trtexec --onnx=model.onnx --saveEngine=model.plan \ --fp16 --workspace=2048
5.2 实际部署注意事项
- 采样率必须与训练数据一致(12kHz)
- 安装角度影响故障特征,需与训练数据匹配
- 建议每10分钟执行一次在线诊断
- 环境振动噪声>0.5g时需要先进行降噪处理
6. 常见问题排查
6.1 准确率异常排查表
| 现象 | 可能原因 | 解决方案 |
|---|---|---|
| 测试集准确率<90% | 数据分布偏移 | 检查传感器安装位置 |
| 某类故障识别率极低 | 样本不平衡 | 使用Focal Loss |
| 训练loss震荡 | 学习率过高 | 添加梯度裁剪 |
6.2 显存不足解决方案
修改batch_size与num_workers比例:
python复制# RTX 3060(12GB)推荐配置
train_loader = DataLoader(...,
batch_size=64,
num_workers=4,
pin_memory=True)
我在多个工业现场部署中发现,实际振动信号往往含有大量高频噪声。建议在模型输入端添加可学习的带通滤波层:
python复制self.filter = nn.Conv1d(1, 1, kernel_size=32, padding=16, bias=False)
nn.init.constant_(self.filter.weight, 1/32) # 初始化为均值滤波
