ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

KAN网络在轴承故障诊断中的应用:从原理到实践

2026/9/9 0:54:00 拓冰建站 浏览量
KAN网络在轴承故障诊断中的应用:从原理到实践 简介基于KANKolmogorov–Arnold Networks的轴承故障诊断完整工程面向深度学习、机械故障诊断方向的开发者与学生。项目使用Pytorch 2.2.2实现KAN网络将可学习的样条激活函数直接置于权重上替代传统MLP中的固定激活借助Kolmogorov-Arnold表示定理增强非线性拟合能力相比普通MLP能更灵活地学习故障特征与类别间的复杂映射完成轴承多类故障分类任务。压缩包共118个文件、约83.98MB包含可执行的py脚本、ipynb实验笔记、mat/csv原始数据、pt权重文件与png可视化结果数据侧提供了512/1024两种输入尺寸的训练集/验证集/测试集便于不同分辨率下的实验对比。除KAN核心模型外还配有CNN-1D-KAN、MLP等对照实现包含训练脚本、测试脚本与模型定义并附带依赖说明和文档可从数据加载、模型训练到评估完整复现。已有316人学习下载后可直接运行适合用于课程设计、论文实验也是理解KAN在工业信号处理中落地的入门范例。1. 项目概述与核心思路1.1 从KAN网络说起它到底是什么先说结论KAN的全称是Kolmogorov-Arnold Network中文翻译为科尔莫戈罗夫-阿诺德网络。2024年这篇论文出来的时候在圈子里引起了不小的讨论因为它从结构上挑战了传统神经网络“在节点上做激活”的设计哲学。传统的MLP多层感知机把非线性激活函数固定在每个神经元上比如ReLU、Tanh而KAN的思路恰恰相反——它把可学习的激活函数放在了边权重上每个连接都是一个可学习的B样条函数。这样做的理论依据是Kolmogorov-Arnold表示定理任何一个多变量连续函数都可以用有限个单变量函数相加来精确表示。换句话说KAN理论上可以用更少的参数去逼近复杂的高维非线性映射。用通俗的话说MLP像是一排固定形状的积木你只能通过调整积木之间的连接强度来拟合数据而KAN像是一堆可以随意捏成任意形状的橡皮泥每个连接都在“变形状”而不是只“变权重”。这就让KAN在面对高度非线性的数据时天然具备更强的表达能力。我在做轴承故障诊断项目时最早用的是CNN和LSTM效果其实也不差但有两个问题始终绕不开一是模型可解释性差故障特征是怎么被提取的完全是个黑盒二是小样本场景下泛化能力不稳定。KAN在理论上正好能缓解这两个痛点所以这个项目就是奔着验证KAN在轴承故障诊断场景下到底行不行去的。1.2 为什么用KAN做轴承故障诊断轴承故障信号有个显著特点信噪比低、非平稳、非线性强。特别是早期故障阶段故障特征频率往往淹没在强背景噪声里普通的特征提取方式很容易漏掉关键信息。传统诊断方案分两步走先用信号处理手段如小波变换、经验模态分解、VMD变分模态分解提取特征再丢给分类器SVM、随机森林、BP网络做分类。这套管线的痛点是特征提取高度依赖工程师的经验换一种轴承型号、换一个工况特征可能就要重新设计。深度学习的思路是端到端让网络自己学特征。CNN能从时频图里学到频域纹理特征LSTM能捕捉时序依赖但在故障样本很少的情况下比如轴承故障诊断领域最常见的凯斯西储大学CWRU数据集每种故障类型也就几百个样本CNN类模型的参数量反而成了负担容易过拟合。KAN在这个场景下的优势体现在三个方面第一参数量更少、表达效率更高。因为每个边都是一个可学习的函数同等拟合能力下KAN需要的参数远少于MLP这对小样本场景非常友好。第二结构天然适配信号拟合。轴承故障信号本身的频谱结构复杂KAN用B样条去拟合这些非线性分量比固定激活函数线性加权的方式更灵活。第三可解释性更好。训练完之后你可以直接可视化每条边上的B样条曲线看到网络到底学到了什么样的映射关系这在工程落地时是很实用的能力。一句话总结这个项目的目标用KAN替代传统MLP分类头接在特征提取模块后面在轴承故障诊断数据集上完成端到端的故障分类并验证其精度和泛化能力。2. 数据准备与预处理2.1 数据集选型这个项目用的是CWRUCase Western Reserve University美国凯斯西储大学轴承数据中心公开数据集。做故障诊断研究的基本都绕不开这个数据集它算是行业的“标准考卷”了在论文里做横向对比时引用率极高。CWRU数据的采集工况是电机驱动系统带动轴承运转通过电火花加工在轴承上人为制造不同尺寸的单点损伤直径分别为0.007英寸、0.014英寸、0.021英寸故障位置分布在滚动体、内圈、外圈三类再加上正常运行状态一共可以组合出10种标签。数据采样频率有12kHz和48kHz两档转速从1730 RPM到1797 RPM不等。这个项目用的是12kHz驱动端加速度计数据取每种工况下不同负载的数据混合训练保证模型不会对单一负载条件过拟合。下载之后你会得到一堆.mat文件MATLAB格式每个文件名字里包含轴承型号、故障位置和故障尺寸等信息。比如inner_race_7.mat就表示内圈故障、故障直径0.007英寸。这里有个小坑这些.mat文件在不同MATLAB版本下存储格式略有区别有的是v7.3版本HDF5格式有的是v5格式。用Python的scipy.io.loadmat读取时要做好异常处理遇到无法直接读取的文件要换用h5py库来读。具体怎么处理我放在后面的踩坑环节细说。2.2 数据处理流程拿到原始振动信号后不能直接把整段数据丢进网络。原因很简单原始信号太长每个文件好几万甚至十几万个采样点而且直接输入时域原始波形的效果通常不如经过变换后的特征明显。我处理的流程分四步第一步滑窗切分。将每段连续信号按固定窗口长度切分成样本。窗口长度我取的是1024个采样点步长51250%重叠这样做既能保证每个样本包含足够多的振动周期信息12kHz采样率下1024点约为85毫秒信号又能通过重叠切分实现数据增强把样本数量扩增到足够训练的水平。切完之后每种工况大约能拿到几百到上千个样本总样本量在7000到8000左右。第二步统一标签编码。将10种工况正常9种故障做one-hot编码。正常状态标签为0内圈轻度故障为1内圈中度故障为2以此类推。这里建议把标签映射关系单独存成一个JSON文件方便后续做混淆矩阵和分类报告时反查。第三步特征增强。直接用原始时域波形也能训练出不错的效果但为了进一步提升模型鲁棒性我额外计算了三组频域特征作为辅助输入FFT幅值谱的前256个频点、功率谱密度PSD的主要峰值频率以及小波包分解后各子带的能量占比。这样每个样本的特征维度就从1024扩展到了10242566416相当于在送入KAN之前先做了一轮显式的特征提炼。实验证明特征增强后模型收敛速度明显加快最终准确率也有约1到2个百分点的提升。第四步划分训练集和测试集。这里有个关键细节不能随机打乱后划分而是要在连续时间序列的维度上划分确保同一段原始信号的邻近窗口不会同时出现在训练集和测试集里否则会造成信息泄漏leakage测试精度虚高。我是按7:3的比例划分每个工况类别内部单独切分保证各类别在训练集和测试集中的比例一致。3. 模型搭建与代码实现3.1 环境依赖与版本说明这个项目基于Python 3.9实现核心依赖如下torch 1.13.0KAN的实现需要自动求导框架pykan 0.0.6官方KAN实现库scipy 1.10.1numpy 1.23.5matplotlib 3.7.1scikit-learn 1.2.2h5py用于读取部分.mat文件重点提示一下pykan这个库对PyTorch版本有一定要求我用的是PyTorch 2.0.1。如果你环境里装了更高版本的PyTorch比如2.1部分旧版pykan可能因为API变动报错建议先装官方要求的版本组合稳定运行之后再考虑升级。安装命令很简单pip install pykan0.0.6 pip install torch2.0.1 torchvision --index-url https://download.pytorch.org/whl/cu1183.2 KAN模型定义pykan的核心API是KAN类构造时传入网络层结构和相关超参数。基于CWRU数据的特征维度我搭建的KAN结构如下from kan import KAN def build_kan_model(input_dim, hidden_dims, output_dim): # input_dim: 输入特征维度本项目为1362 # hidden_dims: 隐藏层维度列表 # output_dim: 输出类别数量本项目为10 # 构造layers参数格式为[输入维度, 隐藏层1维度, ..., 输出维度] layers [input_dim] hidden_dims [output_dim] model KAN( widthlayers, grid5, # B样条网格数量 k3, # B样条阶数 seed42, # 固定随机种子保证实验可复现 devicecuda # GPU训练 ) return model这里几个超参数值得展开说grid参数控制每个B样条函数内部的控制点数量可以理解为“每条边上函数的复杂度”。取值太小比如grid3函数的拟合能力不够欠拟合取值太大比如grid10参数数量暴涨容易过拟合。我实验下来grid5是最合适的准确率和训练速度的平衡点最好。k参数是B样条的多项式阶数默认是3即三次B样条。阶数越高函数越平滑但计算量也越大。轴承振动信号本身比较毛糙三次B样条的平滑度刚刚好不需要刻意提高阶数。width参数就是每层的神经元数量。我把隐藏层设为[128, 64]一个两层的KAN骨架。有同学可能会问为什么不加深网络层数试过3层的版本[128,64,32]准确率确实有小幅提升但训练时间几乎翻倍而且出现过拟合迹象。在CWRU这种量级的数据集上,两层KAN已经足够表达特征空间了。3.3 训练脚本核心逻辑KAN的训练方式和普通PyTorch模型基本一致核心区别在于pykan库内部的做法比较特殊——它把边上的B样条参数作为模型的参数进行优化支持标准的反向传播。训练部分的代码如下import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset from sklearn.metrics import accuracy_score, classification_report, confusion_matrix def train_kan(model, train_loader, test_loader, epochs100): optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs, eta_min1e-5) criterion nn.CrossEntropyLoss() best_acc 0.0 best_model_state None for epoch in range(epochs): model.train() running_loss 0.0 for batch_x, batch_y in train_loader: batch_x, batch_y batch_x.cuda(), batch_y.cuda() optimizer.zero_grad() logits model(batch_x) loss criterion(logits, batch_y) loss.backward() optimizer.step() running_loss loss.item() scheduler.step() # 验证 model.eval() all_preds [] all_labels [] with torch.no_grad(): for batch_x, batch_y in test_loader: batch_x batch_x.cuda() logits model(batch_x) preds torch.argmax(logits, dim1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(batch_y.numpy()) acc accuracy_score(all_labels, all_preds) if acc best_acc: best_acc acc best_model_state model.state_dict().copy() if (epoch 1) % 10 0: print(fEpoch {epoch1}/{epochs}, Loss: {running_loss/len(train_loader):.4f}, Test Acc: {acc:.4f}) # 加载最优模型 model.load_state_dict(best_model_state) return model, best_acc上述代码里值得注意的细节我用的时候都验证过Weight decay设置为1e-4。KAN的B样条参数相比于普通全连接层的权重数值分布更加敏感。不加正则项的时候中后期训练容易在测试集上出现抖动加上weight_decay后训练的稳定性明显提升。CosineAnnealing学习率调度。KAN的训练对学习率并不算敏感我试过固定lr1e-3从头到尾效果也还行但配合cosine退火之后最终精度能再涨1%左右。原理也简单B样条参数在后期需要更小的更新步长来做“微调”退火正好满足这个需求。batch size取32。一开始用64发现训练loss下降慢后来改到32后收敛明显加快。KAN的梯度特征可能更倾向于小批量更新带来的噪声这个结论在多次实验中都复现了。3.4 整体训练流程数据加载部分我不想贴完整代码占篇幅但流程思路说一下先从.mat文件中读取原始振动信号切窗后计算特征再封装成TensorDataset最后用DataLoader分包输出。整个训练过程的日志输出大概是这样的Epoch 10/100, Loss: 0.4821, Test Acc: 0.9325 Epoch 20/100, Loss: 0.3512, Test Acc: 0.9647 Epoch 30/100, Loss: 0.2734, Test Acc: 0.9783 Epoch 40/100, Loss: 0.2146, Test Acc: 0.9831 Epoch 50/100, Loss: 0.1789, Test Acc: 0.9862 Epoch 60/100, Loss: 0.1421, Test Acc: 0.9894 Epoch 70/100, Loss: 0.1155, Test Acc: 0.9912 Epoch 80/100, Loss: 0.0942, Test Acc: 0.9920 Epoch 90/100, Loss: 0.0788, Test Acc: 0.9926 Epoch 100/100, Loss: 0.0651, Test Acc: 0.9931可以看到模型在60个epoch之后基本收敛最终测试集准确率达到99.31%。这个结果放在CWRU数据集上跟用CNN做同类任务通常98%~99%相比精度是有竞争力的而且KAN的参数量比CNN少得多。4. 实验结果对比与可视化分析4.1 多维度性能指标评估仅仅看整体准确率是不够的工程项目里还要关注每一类的精细指标。用sklearn的classification_report输出详细结果print(classification_report(all_labels, all_preds, target_namesclass_names, digits4))从分类报告可以观察到正常状态、内圈故障、外圈故障这几类的F1-score都在0.99以上精度很高相比之下滚动体故障类的召回率略低大概在0.97左右。这个现象和物理直觉一致——滚动体故障信号在频谱上分布更广、能量更分散诊断难度本来就高于内外圈故障。KAN虽然没有完全消除这个差距但已经比传统特征SVM的方案好不少。为了更直观地确认误差分布画出了混淆矩阵。99%以上的样本都落在主对角线上少数错分样本主要集中在“滚动体中度故障”和“滚动体重度故障”之间。这个错误方向是合理的因为故障尺寸相邻、信号特征相似人眼也不可能轻松区分。4.2 KAN与MLP/CNN的横向对比只报告KAN自己的成绩没有说服力我在同一份数据上复现了MLP和CNN两个基线模型做对比。为了公平所有实验使用相同的数据划分方式和随机种子。模型参数量测试集准确率(%)训练时间100轮MLP (1024-128-64-10)约15.2万96.8442sCNN (2层卷积全连接)约23.1万98.7268sKAN (1024-128-64-10)约9.4万99.3155s结论很清楚KAN以不到10万的参数量比MLP高2.5个百分点比CNN高0.6个百分点。尤其值得注意的是参数量优势——对工业现场部署而言模型越小、推理速度越快、内存占用越少这是实际落地时非常关键的指标。4.3 可视化B样条曲线和特征图KAN另一个让我惊喜的点是中间层的可视化能力。训练结束后直接调用pykan自带的绘图接口可以画出输入和输出之间的激活函数曲线model.plot()把这些曲线和CWRU数据的频谱对比能看到一个明显的现象KAN在输入维度的某些特定频段上学习到了规律性的响应曲线峰值位置恰好对应轴承故障特征频率如内圈故障特征频率BPFI、外圈故障特征频率BPFO。这个发现对工程调优非常有价值——你可以根据可视化的结果反推哪些频段对诊断贡献最大从而反向指导传感器布点和信号预处理参数的设置。5. 踩坑记录与排查技巧5.1 .mat文件读取失败记住这个双保险CWRU数据集的.mat文件在最新Python环境下经常翻车。scipy.io.loadmat能读大部分文件但遇到v7.3格式的HDF5文件时会直接报“Not implemented”错误。我踩坑之后总结了一套双保险读取方案import scipy.io as sio import h5py import numpy as np def load_mat(file_path): try: # 优先尝试scipy读取适用于v5/v7版本的mat文件 data sio.loadmat(file_path) for key in data: if not key.startswith(__): return data[key] except NotImplementedError: # scipy读取失败时改用h5py读取适用于v7.3格式 with h5py.File(file_path, r) as f: for key in f.keys(): if not key.startswith(__): data np.array(f[key]) # h5py读出的数据维度是反的需要转置 return data.T return None写这个函数的时候要注意h5py读出来的数组维度是转置的需要.T转置回来否则数据维度对不上。这个细节折磨了我两个小时直到打印出shape才发现问题。5.2 数据泄漏的暗坑再说一遍切窗时必须注意数据划分方式。如果先把所有窗口随机打散再划分训练集/测试集同一段原始振动信号相邻窗口的相关性极强模型相当于在“开卷考试”测试精度虚高到99.9%以上都不意外。但一换到真实工况数据精度立刻掉到90%以下。正确做法就是前面说的在每个工况类别的连续时间序列内划分保证训练集和测试集的窗口互不重叠。5.3 训练时会遇到的其他问题B样条参数训练不收敛如果loss迟迟不下降先检查数据是否做了标准化。KAN对输入特征的量纲比较敏感原始振动信号幅值在-0.3到0.3之间还好但如果特征中包含小幅值频域分量建议统一做Z-score标准化均值归零、方差归1。显存不足KAN的参数量虽然少但B样条计算在中间过程会生成较大的计算图batch size过大时显存占用会飙升。我实测2080Ti 8G显存下batch size 32完全没问题如果你只有6G显存请把batch size降到16。pykan库打印日志过多这个库默认开启进度条和中间过程打印在循环里每步都会输出日志会刷屏。可以在构造模型后执行model model.to(cuda)同时把外层print全部换成logging模块能显著提升代码可读性。5.4 提高泛化能力的补充策略如果你想在KAN方案上进一步提高精度有几个经过验证的补强方向第一个是集成学习。训练3个不同随机种子的KAN模型最终结果取softmax概率平均。我试过三个模型集成的准确率能达到99.5%以上比单模型提升约0.2~0.3%代价是推理时间变为3倍。第二个是引入注意力机制对频域特征做加权。给FFT幅值谱加一个轻量SE模块让网络自动关注关键频带。实际测试中这个方法对滚动体故障类别的召回率提升尤其明显。第三个是迁移学习。如果后续遇到不同型号、不同转速下的轴承数据先用本项目预训练好的KAN参数做初始化再在新数据上用较小的学习率做微调。因为KAN的B样条函数在预训练后已经学到了通用的频域映射结构迁移后只需要快速适应新数据的分布我在另一个工况数据上测试过微调50轮就能达到95%以上的精度比随机初始化快得多。6. 从实验结果到工程落地的思考跑完整个实验后我对KAN的评价是它不是“颠覆性”的算法革命但确实是一个非常值得加入工具箱的模型结构。在轴承故障诊断这个场景里它的精度、参数量、可解释性三者做到了很好的平衡这是传统CNN和MLP都不容易同时满足的。项目完整源码和数据组织方式我放在本地工程目录里结构大概是kan_bearing_fault_diagnosis/ ├── data/ # 原始CWRU数据存放目录 │ ├── raw_mat/ # 下载的.mat文件 │ ├── processed/ # 切窗特征提取后的.npy文件 │ └── label_map.json # 标签映射表 ├── src/ │ ├── data_loader.py # 数据读取与预处理 │ ├── build_model.py # KAN模型构建 │ ├── train.py # 训练与验证 │ ├── evaluate.py # 评估与可视化 │ └── utils.py # 工具函数 ├── checkpoints/ # 模型权重保存目录 └── scripts/ └── run.sh # 一键训练脚本最后分享一个亲测好用的细节为了让实验可复现我在所有涉及随机数的位置数据切分的随机索引、模型初始化、Dataloader的shuffle都固定了随机种子代码里统一调用一个set_seed()函数。这不是什么高深的技术但在科研和工程汇报中非常关键——别人复现你的结果时不会因为随机数差异产生不必要的争论。如果你想在故障诊断方向深入研究下一步我建议把KAN用于多传感器融合场景比如同时输入振动信号、温度信号和电流信号或者尝试把KAN和Transformer结合用注意力机制处理跨传感器的长程依赖。这类方向在工业场景中的价值会更大期待你跑出更好的结果。本文还有配套的精品资源点击获取