如何用 pykan 快速完成你的第一次 KAN 函数拟合:完整指南
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
pykan 是 Kolmogorov-Arnold Networks(KAN)的 PyTorch 实现。它把可学习的"基础函数 + 样条"放在网络的每一条边上,让 KAN 函数拟合既准又"可读":训练完你能直接看到每条边上学到了什么形状,甚至抽出符号公式。这篇文章带你走完一个完整任务:装环境、备数据、训练、评估、剪枝、抽公式,每一步的代码都能直接复制运行。
全流程一览
上面八个节点对应下文八个小节。只想要最短路径的,直接跳到第四节的代码块。
上图是一个典型的 pykan KAN 网络:每条边挂一个可学习函数,线越粗代表这条边的权重越强。
📦 安装 pykan:一条命令加一个虚拟环境
环境干净与否,决定了后面所有"我的机器有问题"式排查的多少。
推荐方式:PyPI 一键安装
python -m venv .venv source .venv/bin/activate pip install pykanpykan 会连带装好 torch 2.2.2、numpy、matplotlib、scikit-learn 等依赖。GPU 机器上如果 CUDA 报冲突,先单独装好对应版本的 PyTorch,再执行上面的安装即可。
什么时候才用源码安装
只有两种情况需要克隆源码后执行pip install -e .:你要改源码,或者需要 PyPI 尚未发布的版本。普通函数拟合任务用不到。装完后运行from kan import *没有报错,就算准备就绪。
动手前,先看懂 KAN 的两个旋钮
装好包别急着跑代码,花两分钟理解下面两点,能省掉后面大量试错。
网格:像刻度尺的疏密
KAN 与 MLP 的差异就一处:激活函数挂在边上,而且这个函数本身是参数。每条边的函数 = 基础函数 × 可学尺度 + 样条 × 可学尺度。样条铺在网格上,grid决定网格被切成几段。类比尺子的刻度:刻度稀,函数只能学出比较"直"的形状;刻度密,才能拟合细弯的曲线。grid=3是稳妥起点,拟合不够再往上加。
值得知道的几个初始化参数
KAN构造函数的其余参数基本可以不动,完整列表在 kan/MultKAN.py:k=3表示三次样条,绝大多数任务够用;noise_scale=0.3控制样条初始注入的噪声,是平衡值;base_fun选初始基础函数,默认silu;sparse_init=True会把大部分尺度初始化为零,适合做特征选择。
这里值得记住作者的一条经验:从小模型起步。小任务先试width=[输入,5,1]、grid=3、不加正则,跑不通再逐步加宽,最后才考虑加深。KAN 里"默认一百宽"的 MLP 习惯往往帮倒忙。
数据准备:create_dataset 的三种用法
概念清楚了,接下来给模型喂真数据。
从公式生成数据
最直接的用法是把目标函数直接交给create_dataset:
from kan.utils import create_dataset, create_dataset_from_data # 公式生成:每个变量独立范围 + 标签归一化 dataset = create_dataset(f, n_var=3, ranges=[[-1, 1], [0, 10], [-3, 3]], train_num=2000, normalize_label=True) # 已有数据:自动按 8/2 切分训练测试集 dataset = create_dataset_from_data(inputs, labels, train_ratio=0.8)几个容易踩的点:
ranges不传时默认[-1,1]套用到所有变量;传(n_var, 2)的列表则每个变量各用各的范围。- 函数里访问变量默认用列模式
x[:,[i]](f_mode='col');习惯写一维下标的可以改成'row'。 - 输入量纲差异大(比如一个 0~1,一个 0~100)时建议
normalize_input=True,把输入拉进[-1,1]附近,恰好落在样条的初始网格范围内,训练更稳。 - 两个函数返回的都是同一个结构的字典:
train_input、train_label、test_input、test_label,注意 tensor 要和模型在同一个 device 上。
跑通第一次训练:fit 的三个开关
数据在手,KAN 函数拟合只需要一次fit调用:
import torch from kan import * torch.set_default_dtype(torch.float64) device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') f = lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) + x[:,[1]]**2) dataset = create_dataset(f, n_var=2, device=device) model = KAN(width=[2,5,1], grid=3, k=3, seed=42, device=device) model(dataset['train_input']) # 前向一次,顺便初始化网格 model.fit(dataset, steps=50, lamb=0.001) model.evaluate(dataset)官方 notebook 默认用 float64,样条拟合在 float32 下误差偏大,建议照做。fit里有三个开关值得先懂:
steps:LBFGS 迭代次数,默认 100,先跑 50 看趋势即可。lamb:L1 惩罚,默认 0 是纯拟合;给 0.001 开始把网络往稀疏方向推,是"可解释性"的主阀门。update_grid:默认 True,训练中会周期性按数据分布重算网格点(数据分位数与均匀网格之间由grid_eps插值)。想手动控制网格时关掉。
最常动的五个参数
| 参数 | 什么时候动 | 动了之后 |
|---|---|---|
grid | 欠拟合,形状学不出来 | 样条控制点变多,拟合能力增强;过大则过拟合并拖慢训练 |
k | 需要更光滑的边 | 样条阶数升高;3 阶之外收益递减 |
lamb | 想要稀疏可读的网络 | 弱边被压到零,结构变稀疏 |
update_grid | 网格和数据范围对不上 | False 时冻结初始网格,交给你自己调 |
width | 调完 grid 仍欠拟合 | 先加宽、后加深,宽度优先级更高 |
evaluate返回train_loss和test_loss。两者差距明显拉大就是过拟合信号:先降grid(比降width更见效),再考虑加数据或加大lamb。
📊 读出模型学到了什么:plot 与符号公式
损失数字只是半个答案,KAN 的真正卖点是"看得见":
model.plot(in_vars=['x','y'], out_vars=['f']) # 每条边的函数形状 model.symbolic_formula(var='x') # 抽取紧凑符号公式 model.speed() # 关掉符号分支,提速plot展示每条边的激活函数形状和边权(线宽)。哪条边在某区间接近直线,就说明它在那个区域基本学到了线性关系;哪条边几乎平了,说明它没用。
符号分支默认开启(symbolic_enabled=True),但不开并行化,不用它的场景下会拖慢前向。README 专门提醒:自己写训练循环且不做符号回归时,训练前先调一次model.speed()。
🌿 精简网络:剪枝与状态回滚
刚训完的网络常常"多到不需要"。剪枝的标准流程是先稀疏、再剪、再补训:
model.fit(dataset, steps=50, lamb=0.01) # 更强的 L1 推稀疏 model.prune(node_th=1e-2, edge_th=3e-2) # 剪掉弱节点和弱边 model.fit(dataset, steps=20) # 剪枝后补训一轮阈值越小剪得越狠,从 1e-2 量级起步比较安全。prune_input是另一类操作:按输入重要性把贡献极小的输入变量整个移除,相当于让模型自己告诉你哪些特征没用。想查某个输入到底贡献多少,用model.attribute()。
剪坏了也不需要从头再来:auto_save默认开启,训练过程中会自动把检查点存进./model目录,model.rewind(0)就能回到任何一版状态,对比不同稀疏度也很方便。
🛠️ 常见坑速查表
跑完全流程,多半会撞上下面几类问题:
| 症状 | 常见原因 | 对策 |
|---|---|---|
| 训练明显偏慢 | 符号分支开着但没用到 | fit前调用model.speed() |
| 训练/测试损失差距大 | 过拟合,网格过大 | 先降grid再降width,或加数据、加大lamb |
| 输出出现 nan | 目标函数含 1/x、log 等奇异点 | fit和forward传singularity_avoiding=True |
| 两次运行结果不一致 | 种子或精度没固定 | 固定KAN(seed=...)与create_dataset(seed=...),统一 float64 |
下一步去哪看
上面这条流水线是最短上手路径,仓库里还有成体系的材料:
- hellokan.ipynb 与本文代码一一对应,带逐步输出,适合逐行对照。
- tutorials/ 下分四类:Example 是函数拟合与 PDE 的经典案例(深公式发现、奇点、相位转变),API_demo 逐条讲 API,Interp 是可解释性技巧(换边、Hessian、稀疏初始化),Physics 是科学应用(拉格朗日量、黑洞、本构方程)。
- 分类任务先看 Example_4,数值精度看 Example_7;剪枝与符号回归的 API 想系统掌握,Interp_3_KAN_Compiler.ipynb 更集中。
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考