简介:工业设备预测性维护中,轴承故障诊断是核心环节。传统方法依赖人工特征工程,而深度学习模型如MLP和CNN虽能自动提取特征,但参数冗余且在小样本强非线性场景下易过拟合。KAN(Kolmogorov-Arnold Network)作为新型网络架构,基于Kolmogorov-Arnold表示定理,将固定激活函数替换为可学习的B样条基函数,以更少参数实现高精度拟合,显著提升参数效率。在轴承振动信号分析中,KAN可自适应塑造每个维度的非线性映射,尤其适合转速载荷变化大、故障特征弱的工况。本文基于CWRU公开数据集,从数据预处理、KAN网络搭建到训练评估,完整复现故障分类流程,并与MLP、一维CNN在多组配置下对比。结果表明,KAN以远低于MLP的参数达到接近或超过CNN的准确率,在难分故障类别和小样本条件下鲁棒性更优。同时总结了B样条网格参数调节、输入标准化等工程踩坑经验,为故障诊断算法落地提供参考。
1. 为什么我会在轴承故障诊断里尝试KAN这个新架构
做设备故障诊断这些年,轴承问题永远是绕不开的主战场。振动信号里藏着大量非线性、非平稳的故障特征,传统方法要先做特征工程——时域指标、频域峰值、包络谱、小波包能量……一套组合拳打下来,特征设计得好不好,直接决定了后续分类模型的性能上限。我前几年做的几个项目都困在这个环节:样本量不够大,特征提取得对不对心里没底,模型调参调到头也就卡在某个准确率上不去。
后来接触到了KAN(Kolmogorov-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万左右;换成同样宽度的KAN,B样条阶数取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_size=1024, stride=512): 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, axis=0) y_all = np.concatenate(y_all, axis=0) # 形状统一:将二维样本转换为模型输入 [样本数, 特征维度] 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_size=0.3, stratify=y_all, random_state=42 ) # 逐样本归一化,对每个窗口内部做标准化 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_dim=1024, num_classes=10, hidden_dim=128): 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, dim=1) all_preds.extend(pred.cpu().numpy()) all_labels.extend(batch_y.cpu().numpy()) print(classification_report(all_labels, all_preds, digits=4)) 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个epoch,MLP约4分钟收敛,CNN约6分钟,KAN要跑到10分钟左右。推理阶段差距没有训练阶段那么夸张,单样本推理时间仍在可接受范围内,但如果做实时在线诊断,KAN目前的推理速度确实会成为瓶颈。
4.2 准确率与F1对比结果
表格里列出的是10个类别宏观平均的结果:
| 模型 | 参数量 | 准确率(%) | 平均F1(%) |
|---|---|---|---|
| 三层MLP (512-256-128) | 约55万 | 96.31 | 96.25 |
| 一维CNN (3层卷积) | 约21万 | 98.12 | 98.08 |
| KAN (128-64) | 约2.8万 | 97.46 | 97.38 |
| KAN (256-128) | 约11万 | 98.35 | 98.29 |
从表里可以提炼两点信息。第一,KAN以远小于MLP的参数量,打出了接近CNN的成绩,说明它在拟合振动信号非线性映射时的参数效率确实高;第二,第二行的KAN(256-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的人会习惯性地往层后面加bias=True的Linear层,结果反而把原有的函数拟合能力打乱了。我在源码里全部使用bias=False,只让样条基函数自主学习。
网格数量与早停的权衡
网格数量的选择影响很大。我一开始用grid_size=10,训练集准确率很快到了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的过拟合风险更低,这一点对于实际工业现场很有意义——现场能打到的故障样本通常都非常有限。
本文还有配套的精品资源,点击获取