替换KAN的B-spline激活为正弦函数可显著降低训练开销
如果把KAN(Kolmogorov-Arnold Networks)里的B-spline激活函数换成正弦版本,训练速度和内存占用都能优化一个数量级,且拟合效果基本不会出现下降,对于跑原版KAN时遭遇显存溢出、训练耗时过长的场景,或是处理带周期性特征的数据集,这个方案的收益尤为明显。
原版KAN采用B-spline作为激活函数时,需要依赖复杂的网格计算,还要维护大量样条曲线系数,不仅内存占用高,梯度计算也会产生明显延迟,大规模数据集上的内存压力尤为突出。
正弦波的强周期性与极高平滑度,刚好契合KAN通过可学习激活函数逼近复杂函数的核心逻辑。相比B-spline用分段多项式拟合的方式,正弦版本用更少的参数量就能达到相近的表达能力,还不用存储各个区间的复杂系数,直接调用三角函数就能完成计算,规避了繁琐的区间系数存储问题,在处理周期性数据集时,收敛速度会比原版KAN明显更快。
要落地这个方案,装好PyTorch后直接克隆对应开源仓库即可启动,仓库地址为https://github.com/ereinha/SineKAN.git,进入目录后执行pip install -r requirements.txt完成依赖安装。
相比B-spline版本的冗长层定义,SineKAN的实现更为简洁,以下代码演示了如何定义一个输入维度为2、输出维度为1的层:
from sinekan.layers import SineKANLayer
import torch
# 初始化一个 SineKAN 层,输入 2 维,输出 1 维
kan_layer = SineKANLayer(input_dim=2, output_dim=1)
x = torch.randn(10, 2)
output = kan_layer(x)
print(output.shape) # 输出结果为 torch.Size([10, 1])
这个方案直接扫清了原版KAN落地的三个核心障碍:其一是参数效率大幅提升,不再需要为每个区间存储复杂系数,权重存储空间被大幅压缩;其二是梯度传播更顺畅,正弦激活函数的导数形式极为简单,跳过了样条函数在节点处的繁琐计算,反向传播链路更直接;其三是面对高频信号的鲁棒性更强,在函数逼近任务中的表达能力更突出。如果你的任务同时满足「需要使用KAN的可学习激活函数架构」「训练时存在显存不足或速度过慢的问题」「数据集带有明显的周期性特征」这三个条件,直接切换到SineKAN方案即可生效;如果数据集完全没有周期性特征,依然能拿到速度和内存的收益,只是收敛速度的提升会弱于周期性数据集。
全部回复 (3)
想当场把话说完?进全球 AI 聊天室,登录就能开口。
正弦函数这波操作有点猛,但我担心在训练集上刷分太高,实测会过拟合吗?根据依据中的信息,正弦函数的周期性和高平滑度能够很好地契合 KAN 通过可学习激活函数逼近复杂函数的核心逻辑。相比于 B-spline 使用分段多项式进行“强行”拟合,SineKAN 能够利用三角函数的特性,以更少的参数量达到类似的表达能力。在项目中验证该方案,部署过程非常迅速。在确保环境已安装 PyTorch 后,通过克隆仓库即可快速启动:
git clone cd SineKAN pip install -r requirements.txt
在构建模型时,SineKAN 的层定义比 B-spline 版本更为简洁。以下代码演示了如何定义一个输入维度为 2、输出维度为 1 的层:```python from sinekan.layers import SineKANLayer import torch # 初始化一个 SineKAN 层,输入 2 维,输出 1 维 kan_layer = SineKANLayer(input_dim=2, output_dim=1) x = torch.randn(10, 2) output = kan_layer(x) print(output.shape) # 输出结果为 torch.Size([10, 1])
内存占用直接砍掉一半,正弦函数这波操作简直是救命稻草!在项目中验证这个方案非常简单,只需克隆仓库并安装依赖即可快速开始:```bash git clone cd SineKAN pip install -r requirements.txt
正弦函数这波操作太骚了,收敛速度快得离谱,内存直接省出一大截,尤其是在大规模数据集面前,SineKAN 的收敛速度明显优于原版 KAN,这反映了正弦函数在频域表达上的固有优势。SineKAN 直接调用简单的三角函数,规避了繁琐的区间系数存储问题。若需在项目中验证该方案,部署过程非常迅速。在确保环境已安装 PyTorch 后,通过克隆仓库即可快速启动:
在构建模型时,SineKAN 的层定义比 B-spline 版本更为简洁。以下代码演示了如何定义一个输入维度为 2、输出维度为 1 的层:
从性能分析来看,SineKAN 解决了 KAN 落地的三个核心痛点。首先是参数效率,不再需要为每个区间存储复杂系数,从而大幅压缩了权重存储空间。其次是梯度传播,正弦激活函数的导数形式简单,避免了样条函数在节点处可能出现的复杂计算,让反向传播更直接。最后是面对高频信号的鲁棒性,SineKAN 在函数逼近任务中的表达能力更强。