简介:基于KAN(Kolmogorov–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 数据集选型
这个项目用的是CWRU(Case 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个采样点,步长512(50%重叠),这样做既能保证每个样本包含足够多的振动周期信息(12kHz采样率下1024点约为85毫秒信号),又能通过重叠切分实现数据增强,把样本数量扩增到足够训练的水平。切完之后每种工况大约能拿到几百到上千个样本,总样本量在7000到8000左右。
第二步,统一标签编码。将10种工况(正常+9种故障)做one-hot编码。正常状态标签为0,内圈轻度故障为1,内圈中度故障为2,以此类推。这里建议把标签映射关系单独存成一个JSON文件,方便后续做混淆矩阵和分类报告时反查。
第三步,特征增强。直接用原始时域波形也能训练出不错的效果,但为了进一步提升模型鲁棒性,我额外计算了三组频域特征作为辅助输入:FFT幅值谱的前256个频点、功率谱密度(PSD)的主要峰值频率,以及小波包分解后各子带的能量占比。这样每个样本的特征维度就从1024扩展到了1024+256+64+16,相当于在送入KAN之前先做了一轮显式的特征提炼。实验证明,特征增强后模型收敛速度明显加快,最终准确率也有约1到2个百分点的提升。
第四步,划分训练集和测试集。这里有个关键细节:不能随机打乱后划分,而是要在连续时间序列的维度上划分,确保同一段原始信号的邻近窗口不会同时出现在训练集和测试集里,否则会造成信息泄漏(leakage),测试精度虚高。我是按7:3的比例划分,每个工况类别内部单独切分,保证各类别在训练集和测试集中的比例一致。
3. 模型搭建与代码实现
3.1 环境依赖与版本说明
这个项目基于Python 3.9实现,核心依赖如下:
- torch >= 1.13.0(KAN的实现需要自动求导框架)
- pykan == 0.0.6(官方KAN实现库)
- scipy == 1.10.1
- numpy == 1.23.5
- matplotlib == 3.7.1
- scikit-learn == 1.2.2
- h5py(用于读取部分.mat文件)
重点提示一下,pykan这个库对PyTorch版本有一定要求,我用的是PyTorch 2.0.1。如果你环境里装了更高版本的PyTorch(比如2.1+),部分旧版pykan可能因为API变动报错,建议先装官方要求的版本组合,稳定运行之后再考虑升级。
安装命令很简单:
pip install pykan==0.0.6 pip install torch==2.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( width=layers, grid=5, # B样条网格数量 k=3, # B样条阶数 seed=42, # 固定随机种子,保证实验可复现 device='cuda' # GPU训练 ) return model这里几个超参数值得展开说:
grid参数控制每个B样条函数内部的控制点数量,可以理解为“每条边上函数的复杂度”。取值太小(比如grid=3),函数的拟合能力不够,欠拟合;取值太大(比如grid=10),参数数量暴涨,容易过拟合。我实验下来grid=5是最合适的,准确率和训练速度的平衡点最好。
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, epochs=100): optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs, eta_min=1e-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, dim=1) 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(f"Epoch {epoch+1}/{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的训练对学习率并不算敏感(我试过固定lr=1e-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_names=class_names, digits=4))从分类报告可以观察到:正常状态、内圈故障、外圈故障这几类的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.84 | 42s |
| CNN (2层卷积+全连接) | 约23.1万 | 98.72 | 68s |
| KAN (1024-128-64-10) | 约9.4万 | 99.31 | 55s |
结论很清楚: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结合,用注意力机制处理跨传感器的长程依赖。这类方向在工业场景中的价值会更大,期待你跑出更好的结果。
本文还有配套的精品资源,点击获取