深度学习做故障诊断,第一层卷积核往往是一堆随机数,学完后也看不出物理意义。但滚动轴承的故障冲击本来就有明确的形态——非对称、陡升缓降,且频率集中在特征频带,能不能把这些已知的物理规律直接“刻”进卷积核里?

从这一动机出发,设计了一种非对称调制高斯差分小波卷积层,将冲击形态先验与可学习参数融合;同时引入基于B样条可学习激活函数的动态非线性降噪机制,在抑制噪声的同时增强故障特征的非线性表达。

01 让卷积核学会“非对称冲击”

传统CNN的卷积核对信号进行无差别滑动卷积,缺乏对故障冲击形态的针对性。然而,旋转机械的冲击响应本就有规律:上升快、下降慢,左右不对称;不同故障的冲击频率也不同。

AMPConv通过3个设计将这个先验嵌入网络:

  1. 非对称高斯差分核:给每个卷积核分配独立的左右高斯尺度参数,使它能够自由调节波形的陡峭程度和拖尾长度,适配冲击响应的非对称包络。

  2. 正弦调制因子:在非对称DoG波形上叠加可学习的正弦载波,将核能量精确聚焦于故障特征频率附近,相当于给每个核指定了关注频段。

  3. 全参数可训练:左右尺度、调制频率、幅值等全部通过反向传播自适应优化,训练后直接可视化出与故障物理频率高度吻合的卷积核。

换句话说,AMPConv的第一层输出不再是一堆抽象的特征图,而是可直接解读为时频原子响应的物理量。

02 用KAN替换静态激活函数

振动信号中的噪声会淹没早期弱故障。传统软阈值降噪的阈值是静态的,或者仅由全局池化生成,缺乏局部自适应能力。

设计了双注意力引导的柔性软阈值机制,并用基于B样条的可学习激活函数替代常规Sigmoid/ReLU生成基础阈值:

全局+局部双分支:全局分支通过全局平均池化捕捉整体噪声水平,局部分支通过卷积和注意力图捕捉各位置的噪声差异,两者融合生成阈值。

KAN激活:B样条基函数赋予激活函数更强的非线性拟合能力,使阈值曲面能够建模复杂的噪声分布,而不是简单地线性映射。

图片

model = AMGDWKCNN().to(device)
criterion = nn.CrossEntropyLoss()                # 无标签平滑
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)

def train_epoch(model, loader, optimizer, criterion):
    model.train()
    total_loss, correct = 0, 0
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad()
        out = model(x)
        loss = criterion(out, y)
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        optimizer.step()
        total_loss += loss.item()*x.size(0)
        correct += (out.argmax(1)==y).sum().item()
    return total_loss/len(loader.dataset), correct/len(loader.dataset)

def eval_model(model, loader, criterion):
    model.eval()
    total_loss, correct = 0, 0
    with torch.no_grad():
        for x, y in loader:
            x, y = x.to(device), y.to(device)
            out = model(x)
            loss = criterion(out, y)
            total_loss += loss.item()*x.size(0)
            correct += (out.argmax(1)==y).sum().item()
    return total_loss/len(loader.dataset), correct/len(loader.dataset)

best_val_loss = float('inf')
patience_counter = 0
train_losses, val_losses = [], []
train_accs, val_accs = [], []

print(f"\n--- 开始训练 (Early Stopping patience={EARLY_STOP_PATIENCE}) ---")
for epoch in range(EPOCHS):
    tr_loss, tr_acc = train_epoch(model, train_loader, optimizer, criterion)
    val_loss, val_acc = eval_model(model, val_loader, criterion)
    scheduler.step()
    train_losses.append(tr_loss); val_losses.append(val_loss)
    train_accs.append(tr_acc);   val_accs.append(val_acc)

    if (epoch+1) % 20 == 0:
        print(f"Epoch {epoch+1:3d}/{EPOCHS} | TrLoss: {tr_loss:.4f}, TrAcc: {tr_acc:.4f} | ValLoss: {val_loss:.4f}, ValAcc: {val_acc:.4f}")

    if val_loss < best_val_loss:
        best_val_loss = val_loss
        patience_counter = 0
        torch.save(model.state_dict(), 'best_model.pth')
    else:
        patience_counter += 1
        if patience_counter >= EARLY_STOP_PATIENCE:
            print(f"Early stopping at epoch {epoch+1}")
            break

model.load_state_dict(torch.load('best_model.pth'))
test_loss, test_acc = eval_model(model, test_loader, criterion)
print(f"\n测试集准确率: {test_acc:.4f}")

# ----------------------------- 5. 可视化(全部 .detach() 安全) -----------------------------
# 5.1 训练曲线
plt.figure(figsize=(10,4))
plt.subplot(1,2,1); plt.plot(train_losses, label='Train'); plt.plot(val_losses, label='Val')
plt.xlabel('Epoch'); plt.ylabel('Loss'); plt.title('Training and Validation Loss'); plt.legend()
plt.subplot(1,2,2); plt.plot(train_accs, label='Train'); plt.plot(val_accs, label='Val')
plt.xlabel('Epoch'); plt.ylabel('Accuracy'); plt.title('Training and Validation Accuracy'); plt.legend()
plt.tight_layout(); plt.savefig('training_curves.png', dpi=150); plt.show()

# 5.2 混淆矩阵
all_preds, all_labels = [], []
model.eval()
with torch.no_grad():
    for x, y in test_loader:
        x = x.to(device)
        all_preds.extend(model(x).argmax(1).cpu().numpy())
        all_labels.extend(y.numpy())
cm = confusion_matrix(all_labels, all_preds)
plt.figure(figsize=(8,6))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=range(10), yticklabels=range(10))
plt.xlabel('Predicted Label'); plt.ylabel('True Label'); plt.title('Confusion Matrix (Test Set)')
plt.tight_layout(); plt.savefig('confusion_matrix.png', dpi=150); plt.show()

# 5.3 t-SNE 特征可视化
feats_list, labs_list = [], []
model.eval()
with torch.no_grad():
    for x, y in test_loader:
        _, f = model(x.to(device), return_features=True)
        feats_list.append(f.cpu().numpy()); labs_list.append(y.numpy())
feats = np.concatenate(feats_list); labs = np.concatenate(labs_list)
tsne = TSNE(n_components=2, random_state=42).fit_transform(feats)
plt.figure(figsize=(8,6))
sc = plt.scatter(tsne[:,0], tsne[:,1], c=labs, cmap='tab10', alpha=0.7)
plt.colorbar(sc, ticks=range(10))
plt.xlabel('t-SNE Component 1'); plt.ylabel('t-SNE Component 2')
plt.title('t-SNE Visualization of Test Features')
plt.tight_layout(); plt.savefig('tsne_features.png', dpi=150); plt.show()

图片

图片

图片

图片

图片

参考论文:

把冲击形态刻进卷积核:融合物理先验和KAN激活的旋转机械故障诊断(Python)

如果你对信号滤波/降噪,机器学习/深度学习,时间序列预分析/预测,设备故障诊断/缺陷检测/异常检测有疑问,或者需要论文思路上的建议,欢迎学术付费咨询

担任《MSSP》《中国电机工程学报》《宇航学报》《控制与决策》等期刊审稿专家,擅长领域:信号滤波/降噪,机器学习/深度学习,时间序列预分析/预测,设备故障诊断/缺陷检测/异常检测

更多推荐