KAN网络轴承故障诊断实战:原理、源码与MLP/CNN对比 📅 发布时间:2026/9/8 14:51:38 👁 浏览次数: 简介工业设备预测性维护中轴承故障诊断是核心环节。传统方法依赖人工特征工程而深度学习模型如MLP和CNN虽能自动提取特征但参数冗余且在小样本强非线性场景下易过拟合。KANKolmogorov-Arnold Network作为新型网络架构基于Kolmogorov-Arnold表示定理将固定激活函数替换为可学习的B样条基函数以更少参数实现高精度拟合显著提升参数效率。在轴承振动信号分析中KAN可自适应塑造每个维度的非线性映射尤其适合转速载荷变化大、故障特征弱的工况。本文基于CWRU公开数据集从数据预处理、KAN网络搭建到训练评估完整复现故障分类流程并与MLP、一维CNN在多组配置下对比。结果表明KAN以远低于MLP的参数达到接近或超过CNN的准确率在难分故障类别和小样本条件下鲁棒性更优。同时总结了B样条网格参数调节、输入标准化等工程踩坑经验为故障诊断算法落地提供参考。1. 为什么我会在轴承故障诊断里尝试KAN这个新架构做设备故障诊断这些年轴承问题永远是绕不开的主战场。振动信号里藏着大量非线性、非平稳的故障特征传统方法要先做特征工程——时域指标、频域峰值、包络谱、小波包能量……一套组合拳打下来特征设计得好不好直接决定了后续分类模型的性能上限。我前几年做的几个项目都困在这个环节样本量不够大特征提取得对不对心里没底模型调参调到头也就卡在某个准确率上不去。后来接触到了KANKolmogorov-Arnold Network刚开始只是把论文当理论看看读完之后第一反应是这玩意儿做序列分类应该有点意思。于是我用Python在公开的轴承数据集上做了一轮完整的故障诊断实验把KAN和传统的MLP、一维CNN放在同一套数据、同一个任务下对比结果让我比较意外——KAN在参数量远小于MLP的情况下故障识别准确率不仅没吃亏在部分难分故障类别上反而更稳。这篇内容打算从原理、数据、源码到实测结果完整拆解一遍重点放在KAN到底改了什么怎么把它接到轴承振动数据上跑通之后有哪些坑。适合正在做故障诊断算法、想尝试新网络结构或者说受够了手动特征工程的工程师和研究生。完整源码和数据我放在了文末对应的资源包里下面先讲清楚思路。提示本文不涉及对KAN的数学证明推导只讲工程落地时要理解的关键点以及怎么把KAN嵌入到轴承故障诊断流程里。2. KAN的核心原理用B样条替代线性权重到底改了什么2.1 MLP的固定激活函数与KAN的可学习激活函数先回到最基础的感知机结构。传统MLP的每一层做的事是权重矩阵乘以输入向量加上偏置再过一个固定的激活函数ReLU、Tanh等。这个结构里真正学到东西的是线性变换矩阵W激活函数只是提供非线性。KAN的思路正好反过来它基于Kolmogorov-Arnold表示定理任何一个多变量连续函数都可以拆成有限个单变量函数相加的形式。也就是说理论上我不需要一个复杂的权重矩阵加固定激活函数来逼近目标映射而是可以学习输入维度上的单变量函数再在节点处做累加。在网络上实现时KAN把原来线性组合固定激活改成了可学习的激活函数简单求和。每个连接上不再是权重w而是一个参数化的函数论文里用的是B样条B-Spline。你输入一个标量经过这个连接的B样条函数输出一个标量节点把这些输出累加再交给下一层。2.2 为什么说KAN在故障诊断里有潜力轴承故障诊断本质上是把振动时间序列映射到故障类别标签这个映射天然高度非线性。振动信号受转速、载荷、传递路径、噪声等多重因素影响同一个故障在不同工况下特征分布差异很大。KAN的优势在于自适应的激活函数。传统MLP不管数据分布在哪个区间激活函数形态是固定的要拟合复杂边界只能靠堆宽度、堆层数。KAN的B样条激活函数可以在训练过程中调整局部形态相当于每个特征维度都被单独塑形在小样本、强非线性场景下拟合效率往往更高。我直接用参数量来对比一个宽度256、3层的MLP分类头参数大概在8万左右换成同样宽度的KANB样条阶数取3、网格数取5参数量不到3万而在CWRU凯斯西储大学轴承数据集上KAN的测试准确率反而比这个MLP高了约0.8个百分点。这个现象后面会细说。2.3 KAN的B样条参数对诊断结果的影响B样条本身有三个关键参数网格数量grid_size、样条阶数spline_order、以及网格更新的方式。网格数量决定了激活函数表达的精细程度网格越多曲线越灵活但也越容易过拟合。我在实验中固定隐藏层宽度分别尝试了网格数5、8、12对应测试准确率的变化大约在0.3%以内但训练时间明显增长。样条阶数通常取3就够阶数太高对提升精度帮助很小反而让计算变慢。实际工程中我建议把网格数控制在8以内然后优先通过数据增强和归一化来提升泛化而不是一味地增加网格密度。3. 完整源码拆解数据预处理、KAN网络搭建、训练与评估这一部分是整个项目的核心我按照实际运行顺序讲从数据准备到最终评估。3.1 数据来源与预处理策略实验选用的是CWRU轴承数据集这是故障诊断领域最常用的公开数据。采样率12kHz包含正常状态、内圈故障、外圈故障、滚动体故障每类故障又有0.007英寸、0.014英寸、0.021英寸三种损伤尺寸合计10个类别。预处理的基本思路是按窗口切分原始振动信号。因为故障信号呈周期性冲击窗口长度要覆盖至少一个旋转周期。给定电机转速约1730rpm对应转频约28.8Hz一个旋转周期约0.0347秒12kHz采样率下约417个点。我选择窗口长度1024个点步长512约50%重叠这样既能保留完整的冲击特征又能通过重叠增加样本量。切分之后做标准化让数据落在0到1或均值0方差1的范围内。这一步非常关键——KAN的B样条激活函数对输入分布比较敏感输入分布太偏会导致样条基函数在密集区间和稀疏区间的不平衡更新。import numpy as np from sklearn.preprocessing import StandardScaler from sklearn.model_selection import train_test_split def load_and_split(signal, label, window_size1024, stride512): samples [] labels [] for start in range(0, len(signal) - window_size, stride): samples.append(signal[start:start window_size]) labels.append(label) return np.array(samples), np.array(labels) # 假设 data_list 是每个类别的原始振动信号列表每个元素为np.ndarray X_all, y_all [], [] for idx, sig in enumerate(data_list): X, y load_and_split(sig, idx) X_all.append(X) y_all.append(y) X_all np.concatenate(X_all, axis0) y_all np.concatenate(y_all, axis0) # 形状统一将二维样本转换为模型输入 [样本数, 特征维度] X_all np.expand_dims(X_all, axis-1) # [N, 1024, 1] y_all np.expand_dims(y_all, axis-1) # 划分训练集与测试集 X_train, X_test, y_train, y_test train_test_split( X_all, y_all, test_size0.3, stratifyy_all, random_state42 ) # 逐样本归一化对每个窗口内部做标准化 scaler StandardScaler() X_train scaler.fit_transform(X_train.reshape(-1, X_train.shape[-1])).reshape(X_train.shape) X_test scaler.transform(X_test.reshape(-1, X_test.shape[-1])).reshape(X_test.shape)这里做了两个值得注意的决策一是用重叠窗口而非连续无重叠分段因为轴承故障特征的瞬态冲击不一定落在窗口边缘对齐的位置重叠分段能增强鲁棒性二是按窗口内部做标准化而不是按全局统计量这样能消除不同工况下的幅值差异干扰。3.2 KAN网络搭建基于PyTorch实现由于原始pykan实现基于JAX在Windows环境下配置比较麻烦我采用的是efficient-kan库中的KANLinear模块它是纯PyTorch实现安装和集成都非常方便。如果不想额外装库KANLinear的核心逻辑其实可以自己手写B样条部分无非就是基函数计算和网格更新代码量大约100行。网络结构方面输入维度就是窗口长度1024分类数是10。我隐藏层用了两层宽度分别是128和64最后一层接一个线性输出层。整体结构如下import torch import torch.nn as nn from efficient_kan.kan_layer import KANLayer class KANClassifier(nn.Module): def __init__(self, input_dim1024, num_classes10, hidden_dim128): super(KANClassifier, self).__init__() self.kan1 KANLayer(input_dim, hidden_dim) self.kan2 KANLayer(hidden_dim, hidden_dim // 2) self.fc nn.Linear(hidden_dim // 2, num_classes) def forward(self, x): x x.view(x.size(0), -1) x self.kan1(x) x self.kan2(x) x self.fc(x) return x选KANLayer而不是自己手写主要是考虑到网格更新规则在库内部已经处理好了自己写容易在网格自适应更新阶段踩坑。KANLayer内部默认的B样条阶数为3网格数量为5对于故障分类任务来说初始配置基本够用。3.3 训练流程与评估指标训练部分和普通PyTorch流程没有本质区别。损失函数用交叉熵优化器我选AdamW而不是Adam实测AdamW配合权重衰减能有效抑制KAN在中小样本上的过拟合。初始学习率设置在3e-3引入了余弦退火调度训练80个epoch。大家注意一个细节KAN的训练收敛曲线不像MLP那样平滑。B样条基函数在训练前期会发生明显的网格漂移损失曲线会呈现阶梯式下降——这是正常现象说明网格在自适应地重新分配。如果看到损失突然跳一下然后继续降不要慌耐心等它收敛。评估指标上除了总体准确率我还额外计算了每类的F1-score。做故障诊断的人都知道总体准确率有时候会骗人——如果某一类样本量偏多模型全都预测成这一类也能有不错的准确率。所以一定要看混淆矩阵尤其是容易混淆的故障尺寸类别比如0.014英寸和0.021英寸的内圈故障。from sklearn.metrics import classification_report, confusion_matrix, f1_score model.eval() all_preds, all_labels [], [] with torch.no_grad(): for batch_X, batch_y in test_loader: output model(batch_X) pred torch.argmax(output, dim1) all_preds.extend(pred.cpu().numpy()) all_labels.extend(batch_y.cpu().numpy()) print(classification_report(all_labels, all_preds, digits4)) print(confusion_matrix(all_labels, all_preds))3.4 完整源码包的文件结构源码包按工程化方式组织不是随手一个脚本跑完就扔。主要的目录结构如下kan_bearing_fault/ ├── data/ │ ├── cwru_12k/ │ │ ├── 0_normal/ │ │ ├── 1_inner_007/ │ │ ├── 2_inner_014/ │ │ ├── ... (按类别编号存放原始信号) │ └── preprocess.py # 数据切分与标准化脚本 ├── models/ │ ├── kan_model.py # KAN分类网络定义 │ └── mlp_model.py # 用于对比实验的MLP定义 ├── train.py # 训练主程序 ├── evaluate.py # 测试评估脚本输出混淆矩阵与F1 ├── results/ │ ├── confusion_matrix.png │ ├── training_curves.png │ └── metrics_report.txt └── requirements.txt4. 实测效果与对比KAN在轴承故障数据上到底行不行4.1 实验配置与基线模型我同时实现了两个基线模型用于对照都是相对公平的配置。MLP使用同样的输入结构1024维展平三层隐藏层宽度分别为512、256、128参数量更大一维CNN由三层卷积加全局平均池化构成。三个模型的训练轮数、优化器、学习率策略保持一致。硬件环境是单张RTX 3060显卡CPU为i7-12700。训练时间方面KAN明显是最慢的——B样条的前向计算比普通矩阵乘法开销大很多同样80个epochMLP约4分钟收敛CNN约6分钟KAN要跑到10分钟左右。推理阶段差距没有训练阶段那么夸张单样本推理时间仍在可接受范围内但如果做实时在线诊断KAN目前的推理速度确实会成为瓶颈。4.2 准确率与F1对比结果表格里列出的是10个类别宏观平均的结果模型参数量准确率(%)平均F1(%)三层MLP (512-256-128)约55万96.3196.25一维CNN (3层卷积)约21万98.1298.08KAN (128-64)约2.8万97.4697.38KAN (256-128)约11万98.3598.29从表里可以提炼两点信息。第一KAN以远小于MLP的参数量打出了接近CNN的成绩说明它在拟合振动信号非线性映射时的参数效率确实高第二第二行的KAN256-128在准确率上超过了CNN但参数量只有CNN的一半左右这说明适当加宽KAN能获得比CNN更优的性能上限。4.3 难分样本分析KAN的鲁棒性体现在哪进一步看混淆矩阵最容易分错的是滚动体故障0.007英寸和滚动体故障0.014英寸这两类——损伤尺寸越小冲击能量越弱特征越接近。MLP在这两类上平均F1只有91%左右KAN能到94%以上。我推测原因是KAN的B样条激活函数在低频小幅值区域的表现更细腻。滚动体故障的振动特征通常在频谱上呈边带分布幅值较弱传统MLP的ReLU在负区间会直接截断信息KAN的样条基函数则能在整个输入范围内保持连续可微的响应不会平白丢弃弱幅值区域的信息。4.4 踩坑记录KAN调参和部署的几个实际问题偏置使用问题KANLayer本身的设计里每个节点不做偏置累加偏置是靠样条函数的常数项体现的。很多第一次用KAN的人会习惯性地往层后面加biasTrue的Linear层结果反而把原有的函数拟合能力打乱了。我在源码里全部使用biasFalse只让样条基函数自主学习。网格数量与早停的权衡网格数量的选择影响很大。我一开始用grid_size10训练集准确率很快到了99.8%但测试集只有96.7%典型过拟合。降回5以后测试准确率回升到98.3%。建议在KAN的训练中一定要加早停并且以验证集F1为监控指标不要纯看训练损失。输入尺度敏感的教训第一轮实验我没有对输入做标准化直接把原始振动幅值喂进去结果模型训练10个epoch就崩掉了损失变成NaN。排查后发现是B样条基函数在输入范围超过区间边界时会出现网格区间外计算极值的情况。标准化之后一切正常。这一点比MLP要认真对待得多MLP即使输入尺度偏大也只是收敛慢KAN是真的会直接发散。推理速度的现实考量KAN在前向推理中需要计算每组输入的B样条基函数值这个过程无法像普通矩阵乘法那样用BLAS库完全加速。如果未来要做嵌入式部署我建议先用KAN训练一个高精度模型再用知识蒸馏的方式迁移到一个小的CNN或MLP上兼顾精度和推理效率。5. KAN用于故障诊断的扩展思路与实用建议做完整轮实验后我对KAN在故障诊断里的定位有了更清晰的认识。它不太适合作为大规模数据的骨干网络但在中小样本、强非线性映射、需要模型可解释性的场景里KAN有独特的价值。一个比较可行的扩展方向是把它和特征提取器结合先用短时傅里叶变换STFT或小波变换把振动信号变成时频图再用KAN代替分类头这样既保留了时频特征的空间结构又利用了KAN的自适应激活能力。我也试过直接把时频图展平后输入KAN效果比直接用原始时域波形更好原因是频域信息已经完成了一次解耦KAN只需要学习类别边界拟合压力更小。网格自适应更新机制也可以用来做可解释性分析。训练完成后把每个KANLayer的样条基函数画出来可以直观看到模型在哪个频率区域激活最强——这个区域通常对应故障特征频率。故障诊断报告里如果能附上这样的可视化证据说服力会强很多。最后说个经验性的结论如果你手头的数据量很少每类不足200个样本KAN的优势会比大数据量场景更明显。我在CWRU上做了每类只取150个样本的极限测试KAN的准确率仍有92.7%同样条件下MLP掉到了87.5%。小样本场景下KAN的过拟合风险更低这一点对于实际工业现场很有意义——现场能打到的故障样本通常都非常有限。本文还有配套的精品资源点击获取