☰
KAN 实战指南:5 步把 hellokan 跑到输出解析公式
2026/9/27 2:49:30 网站建设 项目流程

KAN 实战指南:5 步把 hellokan 跑到输出解析公式

【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan

本文面向有 PyTorch 基础的工程师,解决用 KAN(pykan 实现)时的三个真实问题:为什么训练比 MLP 慢一个量级、loss 停在 1e-2 不动怎么办、以及如何把训好的样条激活函数变成可解释的符号公式。

整个过程在仓库根目录的 hellokan.ipynb 一个 notebook 里就能走完:任务是拟合 f(x,y) = exp(sin(πx)+y²),终点是拿到公式 exp(1.0·x₂² + 1.0·sin(3.1416·x₁))。下面 5 步就是通往这个结果的路径,对应 KAN 论文(arXiv 2404.19756,2024)给出的核心工作流。

一、先跑通 hellokan:起步真正的坑在哪

安装就一行pip install pykan(当前版本 0.2.8),依赖里 torch 2.2.2、numpy 1.24.4 都是 requirements.txt 钉死的版本。如果要在源码上开发,从 https://gitcode.com/GitHub_Trending/pyk/pykan 克隆后执行pip install -e .即可。

跑起来之后第一眼是"慢"。hellokan 默认 50 步 LBFGS 要跑十秒以上,比同规模 MLP 慢得多。原因 README.md 写得很直白:symbolic_enabled默认 True,每次前向都会多算一条 symbolic 分支,而符号计算没有并行化。如果你自己写训练循环、用不到符号功能,训练前调model.speed()把这条分支砍掉,速度立刻正常。

model = KAN(width=[2,5,1], grid=3, k=3, seed=42) model.fit(dataset, opt="LBFGS", steps=50, lamb=0.001) model = model.prune(); model = model.refine(10) model.auto_symbolic(lib=['x','x^2','exp','log','sqrt','tanh','sin','abs']) model.symbolic_formula() # -> exp(1.0*x_2**2 + 1.0*sin(3.1416*x_1))

上面 5 行就是全文骨架:训练 → 剪枝 → 网格加密 → 符号回归 → 出公式。另一个细节:hellokan 里torch.set_default_dtype(torch.float64)不是随手写的,样条系数对数值精度敏感,float32 下会看到莫名其妙的震荡。

二、KAN 真正的三个超参:width、grid、k

KAN 和 MLP 是双生的:MLP 的激活函数在节点上,KAN 的激活函数在边上——每条边放一个可学的一维函数。看 kan/KANLayer.py 的 forward 会发现它由两部分相加:一条 B-spline 曲线,加一个 base 函数(默认 SiLU)再乘可训练系数。所谓"基于 Kolmogorov-Arnold 表示定理"落到代码里就是这个结构。

grid 和 k 到底控制什么

grid是 B-spline 的网格区间数,k是分段多项式阶数(默认三次,k=3)。kan/spline.py 里的curve2coef/coef2curve负责"曲线值"与"控制系数"两种表示的互转,训练时优化的是后者。grid 越大,单条边能刻画的函数越精细,参数量也随之上去——所以 grid 对过拟合的影响往往比 width 还大。

由此引出的调参直觉和 MLP 文献是反的。README 的建议:从小模型起步。5 输入 1 输出的任务,先试KAN(width=[5,1,1], grid=3, k=3),而不是照搬 MLP 习惯上 O(10²) 的宽度;不行先加宽,再加深。小模型反馈快,而且小数据定性上通常能代表大数据——这是作者"物理学家思维"的注脚。

三、判断过拟合看 grid,提精度靠 refine

训几轮之后最常见的状态是 train/test loss 拉开差距。此时 README 的处方是先降grid再降width,别急着加数据。

反过来,想要精度就上 KAN 独有的"网格加密":model.refine(new_grid)把现有样条在更密的网格上重新近似,再继续训练。hellokan 里 refine(10) 之后接 50 步 LBFGS,test loss 从 1.7e-2 压到 4.7e-4,三个数量级。这也是为什么"refine 之后要警惕过拟合"——网格越密,单条边的表达能力越接近任意曲线。

另外fit的update_grid=True默认开启:训练中会根据样本分布自适应更新网格,grid_eps(默认 0.02)在均匀网格和分位数网格之间插值。数据分布不均时,这个机制比 MLP 的归一化技巧省心得多。

四、把 KAN 变"稀":lamb 与 prune 的配合

可解释性是 KAN 相对 MLP 的卖点,工程上来自稀疏正则。fit(lamb=0.01)会对边前向激活加 L1 类惩罚(reg_metric 默认edge_forward_spline_n),把没用的边往零上压。

lamb 幅度怎么定

从 0.001 起步,能收敛就往上调,直到 plot 里出现大片接近零的边;调过头(loss 明显恶化)就回退。训完调model.prune(),按edge_th=3e-2、node_th=1e-2的默认阈值把弱边、弱节点物理删掉(实现见 kan/MultKAN.py),然后继续训练让剩下的边补回功能。hellokan 里 plot 的子图从 22 个变成 6 个,就是这一步。

⚠️ 精度和稀疏不是天然矛盾的。论文里两者可以正相关,换个任务又变成 trade-off,所以别贪心:一个阶段只追一个目标。先稀疏剪枝拿到可解释骨架,需要精度时再用 refine 和加数据收尾。

五、从曲线到公式:符号回归的最后一步

样条终究是"曲线",要"公式"得把边锁死到显式函数。两条路:

  • model.fix_symbolic(l, i, j, 'sin')手动锁定某条边为 sin,之后只学幅度系数;
  • model.auto_symbolic(lib=lib)对每条边按 R² 自动匹配 lib 里的候选。

lib 的选择是关键手感。hellokan 用['x','x^2','x^3','x^4','exp','log','sqrt','tanh','sin','abs']起步,缺函数就用 kan/utils.py 的add_symbolic往函数库里加。锁完之后模型自由度骤降,50 步 LBFGS 直接打到 test loss 7e-11 的机器精度量级;symbolic_formula()输出 sympy 表达式,ex_round(..., 4)四舍五入系数就是最终公式。

📌 同一套 API 平移到 PDE、拉格朗日量等科学问题上,tutorials/Physics/ 的 notebook 从守恒律、黑洞到本构方程都有走通案例。注意 PDE 训练是最贵的场景,CPU 上小时到天级别;其余 tutorial 单 CPU 十分钟以内能跑完。

六、KAN 的适用与不适用:作者自己的冷判断

问得最多的问题是"KAN 能不能替掉 MLP"。作者在 README 里的态度很坦率:对纯 ML 任务,KAN 目前不是开箱即用的插件,超参要调,场景更偏向关心高精度或可解释性的小中规模问题——科学发现、函数拟合、PDE。社区里两个方向的改造值得留意:GraphKAN 把 KAN 放进隐空间、前后加线性嵌入层;KANRL 在强化学习里固定部分参数换训练稳定性。边界还在被试出来。

最后一个长程开发的实用细节:MultKAN 的auto_save默认开启,任何状态改动都会往 ./model 目录存带版本的 checkpoint,checkout(model_id)、rewind(model_id)可以回到任意历史状态。KAN 的"训 → 剪 → refine → 锁符号"是非线性流程,有回滚,敢大胆试。

下一步动作:克隆仓库把 hellokan.ipynb 完整跑一遍,重点看prune()之后model.plot()留下的边;能读懂剪枝后的网络为什么还能拟合,再开始调自己数据的超参。

【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询