Kolmogorov-Arnold 网络(KAN)实战入门:基于 pykan 的安装配置、模型训练与可解释性建模全指南
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
本篇技术指南以 pykan 项目官方文档(docs/index.rst 与 docs/intro.rst)为主体,系统讲解 Kolmogorov-Arnold 网络的数学原理、环境安装与依赖配置,并给出从初始化、训练、剪枝到符号公式提取的完整可运行代码流程。读完本文,你将掌握KAN模型的核心 API 用法(train/prune/fix_symbolic/auto_symbolic/symbolic_formula),并理解其底层实现如何在 kan/MultKAN.py 与 kan/KANLayer.py 中落地。
一、KAN 是什么:从 Kolmogorov-Arnold 表示定理说起
Kolmogorov-Arnold 网络(Kolmogorov-Arnold Network,KAN)的灵感来源于Kolmogorov-Arnold 表示定理。该定理指出:有界域上的多元连续函数 $f$,可以写成有限个一元连续函数与加法二元运算的复合。具体地,对于光滑函数 $f:[0,1]^n\to\mathbb{R}$,有:
$$f(x) = f(x_1,...,x_n)=\sum_{q=1}^{2n+1}\Phi_q\left(\sum_{p=1}^n \phi_{q,p}(x_p)\right)$$
其中 $\phi_{q,p}:[0,1]\to\mathbb{R}$,$\Phi_q:\mathbb{R}\to\mathbb{R}$。从某种意义上说,该定理表明"唯一真正的多元函数是加法"——因为任何其他函数都可以用一元函数与求和来构造。然而,这种 2 层、宽度为 $(2n+1)$ 的 Kolmogorov-Arnold 表示由于表达能力有限,可能不够光滑。pykan 通过将其推广到任意深度和宽度来增强表达能力,从而得到实用的可学习网络。
1.1 矩阵形式的 KAN 层
Kolmogorov-Arnold 表示可以写成矩阵形式:
$$f(x)={\bf \Phi}{\rm out}\circ{\bf \Phi}{\rm in}\circ {\bf x}$$
其中 $\Phi_{\rm in}$ 是 $(2n+1)\times n$ 的函数矩阵,$\Phi_{\rm out}$ 是 $1\times(2n+1)$ 的函数矩阵。二者都是如下"函数矩阵"(Kolmogorov-Arnold 层)的特例——该层有 $n_{\rm in}$ 个输入、$n_{\rm out}$ 个输出:
$${\bf \Phi}= \begin{pmatrix} \phi_{1,1}(\cdot) & \cdots & \phi_{1,n_{\rm in}}(\cdot) \ \vdots & & \vdots \ \phi_{n_{\rm out},1}(\cdot) & \cdots & \phi_{n_{\rm out},n_{\rm in}}(\cdot) \end{pmatrix}$$
定义好层之后,只需堆叠层即可构造 KAN 网络:设有 $L$ 层,第 $l$ 层 $\Phi_l$ 的形状为 $(n_{l+1}, n_l)$,则整个网络为:
$${\rm KAN}({\bf x})={\bf \Phi}_{L-1}\circ\cdots \circ{\bf \Phi}_1\circ{\bf \Phi}_0\circ {\bf x}$$
1.2 KAN 与 MLP 的本质差异
作为对比,多层感知机(MLP)是线性层 $\bf W_l$ 与非线性激活 $\sigma$ 的交替复合:
$${\rm MLP}({\bf x})={\bf W}_{L-1}\circ\sigma\circ\cdots\circ {\bf W}_1\circ\sigma\circ {\bf W}_0\circ {\bf x}$$
两者最核心的区别(官方文档原文要点)是:KAN 把激活函数放在边上(edges),而 MLP 把激活函数放在节点上(nodes)。正是这一简单改变,让 KAN 在精度与可解释性两个维度上相比 MLP 更具优势——这一论断在官方文档中被明确表述,并在 docs/index.rst 中作为项目定位提出。
从源码结构看,这一设计体现在 kan/KANLayer.py 中:每个KANLayer在初始化时为每一条输入-输出边维护一组 B 样条系数(coef)与对应的网格(grid),前向传播时通过 kan/spline.py 中的B_batch、coef2curve完成样条求值,从而实现"边上的一维函数"。
二、安装与依赖配置
pykan 支持两种安装方式,官方文档给出了源码安装与 PyPI 安装两条路径。
2.1 从源码安装(开发模式)
git clone https://gitcode.com/GitHub_Trending/pyk/pykan.git cd pykan pip install -e . # pip install -r requirements.txt # 如需同时安装全部依赖pip install -e .以可编辑模式安装当前仓库,安装元数据定义在 setup.py 中:包名为pykan,版本号0.2.8,并声明python_requires='>=3.6'。仓库顶层init.py 与 kan/init.py 会导出MultKAN与utils中的全部符号,因此安装后可直接from kan import *使用。
2.2 从 PyPI 安装
pip install pykan2.3 依赖版本要求
官方文档给出的核心依赖清单(按 requirements.txt 校验)如下:
| 依赖包 | 版本 | 用途 |
|---|---|---|
| matplotlib | 3.6.2 | KAN 可视化绘图(model.plot()) |
| numpy | 1.24.4 | 数值计算 |
| scikit_learn | 1.1.3 | 指标计算、实验工具 |
| setuptools | 65.5.0 | 打包安装 |
| sympy | 1.11.1 | 符号公式推导(symbolic_formula) |
| torch | 2.2.2 | 深度学习框架与自动求导 |
| tqdm | 4.66.2 | 训练进度条 |
| pandas | 2.0.1 | 数据处理(requirements.txt 补充) |
| seaborn | 最新 | 统计绘图(requirements.txt 补充) |
| pyyaml | 最新 | 配置解析(requirements.txt 补充) |
文档标注的参考 Python 版本为python==3.9.7。需要说明的是:这些版本是官方文档与 requirements.txt 记录的验证环境;使用其他 Python 版本时,建议以实际环境兼容性为准适当调整版本号。
三、快速上手:完整实战流程(Hello, KAN!)
下面的流程完整复现 docs/intro.rst 的官方入门教程:用一个 2 输入、1 输出的 KAN 去拟合函数 $f(x,y)=\exp(\sin(\pi x)+y^2)$,并最终自动提取出符号公式。
3.1 初始化 KAN 模型
from kan import * # 创建一个 KAN:2 维输入、1 维输出、5 个隐藏神经元;三次样条(k=3)、5 个网格区间(grid=5) model = KAN(width=[2,5,1], grid=5, k=3, seed=0)对照 kan/MultKAN.py 中KAN.__init__的完整签名,KAN(实际类名为KAN,实现于MultKAN.py,同时支持乘法边即 MultKAN)的核心初始化参数包括:
width:各层宽度列表,例如[2,5,1]表示输入 2 维、隐藏 5 个神经元、输出 1 维;grid=3:网格区间数(默认 3,本例设为 5);k=3:B 样条阶数(cubic spline);mult_arity=2:乘法边(乘法节点)的元数;noise_scale=0.3:初始化时的噪声幅度;scale_base_mu=0.0、scale_base_sigma=1.0:基础函数(base function)的初始化分布参数;base_fun='silu':基础激活函数(可传入字符串或可调用对象);grid_eps=0.02、grid_range=[-1,1]:网格更新时的边界缓冲与网格取值范围;seed=1:随机种子(本例显式指定seed=0保证可复现);device='cpu':运行设备(详见 API_10_device.ipynb);- 以及
sp_trainable、sb_trainable、auto_save、ckpt_path='./model'等训练与断点存档相关开关。
3.2 创建数据集
# 创建数据集 f(x,y) = exp(sin(pi*x)+y^2) f = lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) + x[:,[1]]**2) dataset = create_dataset(f, n_var=2) dataset['train_input'].shape, dataset['train_label'].shape输出为:
(torch.Size([1000, 2]), torch.Size([1000, 1]))create_dataset定义在 kan/utils.py,其完整参数为create_dataset(f, n_var=2, f_mode='col', ranges=[-1,1], train_num=1000, test_num=1000, normalize_input=False, normalize_label=False, device='cpu', seed=0)。默认在 $[-1,1]^n$ 上均匀采样 1000 个训练样本和 1000 个测试样本,返回包含train_input、train_label、test_input、test_label四个键的字典。
3.3 训练前可视化
# 绘制初始化时的 KAN model(dataset['train_input']) model.plot(beta=100)先调用一次前向传播以确定激活范围,再调用plot绘图(kan/MultKAN.py 中plot(folder="./figures", beta=3, metric='backward', scale=0.5, ...),beta控制边的透明度/线宽,metric控制边的重要性度量方式)。初始化时各边粗细接近、函数曲线接近线性样条:
3.4 带稀疏正则化训练
# 训练模型:LBFGS 优化器、20 步、L1 稀疏正则 lamb=0.01、熵正则 lamb_entropy=10.0 model.train(dataset, opt="LBFGS", steps=20, lamb=0.01, lamb_entropy=10.)官方示例输出:
train loss: 1.57e-01 | test loss: 1.31e-01 | reg: 2.05e+01 : 100%|██| 20/20 [00:18<00:00, 1.06it/s]train即fit方法(kan/MultKAN.py 中fit),其完整签名包含opt="LBFGS", steps=100, log=1, lamb=0., lamb_l1=1., lamb_entropy=2., lamb_coef=0., lamb_coefdiff=0., update_grid=True, grid_update_num=10, loss_fn=None, lr=1., start_grid_update_step=-1, stop_grid_update_step=50, batch=-1, metrics=None, ...。要点:
opt支持"LBFGS"(默认,使用仓库自带的 kan/LBFGS.py 实现)与"Adam"等 PyTorch 优化器;lamb是 L1 稀疏正则强度,lamb_entropy是熵正则强度(对应reg_metric='edge_forward_spline_n'下的正则计算,见get_reg);- 训练过程中
update_grid=True会基于输入样本自适应更新样条网格(update_grid_from_samples),grid_update_num控制网格更新次数。
3.5 训练后的可视化
model.plot()训练后各边粗细出现明显差异,部分边趋于消失,曲线形态开始显现函数结构:
3.6 剪枝 KAN
剪枝是 KAN 可解释性工作流的关键步骤。源码中剪枝由 kan/MultKAN.py 的prune(node_th=1e-2, edge_th=3e-2)实现,内部调用prune_node与prune_edge:将重要性低于阈值的节点/边移除(默认节点阈值1e-2、边阈值3e-2)。
方式一:保持原始形状,仅展示掩码:
model.prune() model.plot(mask=True)plot(mask=True)会将剪掉的边以虚线形式保留,便于对比原结构:
方式二:真正缩小网络形状:
model = model.prune() model(dataset['train_input']) model.plot()prune()返回裁剪后的新模型,重新前向传播后绘制,此时网络宽度显著变小,仅保留对输出有贡献的边。
3.7 继续训练并再次可视化
model.train(dataset, opt="LBFGS", steps=50)train loss: 4.74e-03 | test loss: 4.80e-03 | reg: 2.98e+00 : 100%|██| 50/50 [00:07<00:00, 7.03it/s]继续训练后损失进一步下降,正则项reg也因剪枝后结构更稀疏而大幅降低。
3.8 自动或手动符号化激活函数
pykan 支持把学习到的样条函数替换为已知的符号函数(如sin、x^2、exp),这是其可解释性能力的核心体现。
mode = "auto" # "manual" if mode == "manual": # 手动模式:逐边指定符号函数 model.fix_symbolic(0,0,0,'sin') model.fix_symbolic(0,1,0,'x^2') model.fix_symbolic(1,0,0,'exp') elif mode == "auto": # 自动模式:从候选函数库中按拟合优度自动挑选 lib = ['x','x^2','x^3','x^4','exp','log','sqrt','tanh','sin','abs'] model.auto_symbolic(lib=lib)官方示例中自动模式输出:
fixing (0,0,0) with sin, r2=0.999987252534279 fixing (0,1,0) with x^2, r2=0.9999996536741071 fixing (1,0,0) with exp, r2=0.9999988529417926底层实现说明:
fix_symbolic(l, i, j, fun_name, fit_params_bool=True, a_range=(-10,10), b_range=(-10,10), verbose=True, random=False, log_history=True)(kan/MultKAN.py)将第l层、第i行、第j列的边替换为名为fun_name的符号函数,并调用 kan/utils.py 的fit_params拟合符号函数的仿射参数(a,b的可搜索范围由a_range/b_range控制);auto_symbolic(a_range=(-10,10), b_range=(-10,10), lib=None, verbose=1, weight_simple=0.8, r2_threshold=0.0)(kan/MultKAN.py)内部逐边调用suggest_symbolic,在候选库中按 R² 拟合优度挑选最佳符号函数;用户也可通过 kan/utils.py 的add_symbolic注册自定义符号函数扩充候选库;- 每条被替换边的底层由 kan/Symbolic_KANLayer.py 承接,该层以前向传播方式求值符号函数及其仿射参数。
3.9 继续训练至接近机器精度
model.train(dataset, opt="LBFGS", steps=50)train loss: 2.02e-10 | test loss: 1.13e-10 | reg: 2.98e+00 : 100%|██| 50/50 [00:02<00:00, 22.59it/s]符号化之后继续微调,训练损失从1e-1量级骤降至1e-10量级(官方示例输出),验证了"先学样条、再符号化、后精修"这一范式的高效性。
3.10 提取符号公式
model.symbolic_formula()[0][0]输出:
$$\displaystyle 1.0 e^{1.0 x_{2}^{2} + 1.0 \sin{\left(3.14 x_{1} \right)}}$$
这正是目标函数 $\exp(\sin(\pi x)+y^2)$ 的精确恢复!symbolic_formula(var=None, normalizer=None, output_normalizer=None)(kan/MultKAN.py)利用 sympy 将网络逐层复合为单个符号表达式,从而把黑盒网络转化为人类可读的公式。
四、核心 API 速查:从入门到进阶
除上述入门流程外,官方文档目录(docs/demos.rst、docs/modules.rst)覆盖了更完整的 API 面:
| 能力域 | 对应教程 | 涉及核心 API |
|---|---|---|
| 索引与切片 | API_1_indexing.ipynb | model[l][i][j]、get_fun、get_act |
| 绘图 | API_2_plotting.ipynb | plot、attribute |
| 激活提取 | API_3_extract_activations.ipynb | get_act |
| 初始化 | API_4_initialization.ipynb | initialize_from_another_model |
| 网格细化 | API_5_grid.ipynb | refine(new_grid)(kan/MultKAN.py) |
| 训练超参数 | API_6_training_hyperparameter.ipynb | fit的lamb/lamb_entropy/opt等 |
| 剪枝 | API_7_pruning.ipynb | prune、prune_node、prune_edge、prune_input |
| 正则化 | API_8_regularization.ipynb | reg、get_reg、disable_symbolic_in_fit |
| 训练视频 | API_9_video.ipynb | fit(..., save_fig=True) |
| 设备 | API_10_device.ipynb | device='cuda'、model.to(device) |
| 数据集 | API_11_create_dataset.ipynb | create_dataset、create_dataset_from_data |
| 断点保存/加载 | API_12_checkpoint_save_load_model.ipynb | saveckpt、loadckpt、rewind、checkout |
符号化训练完成后,还可以再次绘制以观察最终结构——此时每条保留的边都对应一个已确认的符号函数,网络从"样条黑盒"彻底转变为透明公式:
五、文档生态导航:从入门到研究前沿
pykan 官方文档(根目录为 docs/index.rst)是一个分层体系,便于按需深入:
- 快速入门:docs/intro.rst(本文第三部分即其完整走查);
- API 演示(advanced):docs/demos.rst 汇总 12 个 API 专题 Notebook,覆盖索引、绘图、剪枝、正则化、设备、数据集、断点等;
- 实例应用(Examples):docs/examples.rst 收录 15 个端到端案例,包括函数拟合(Example_1)、公式深度解析(Example_3)、分类(Example_4)、特殊函数(Example_5)、PDE 求解与解释(Example_6/7)、持续学习(Example_8)、奇点(Example_9)、相对论速度叠加(Example_10)、无监督学习(Example_12)、相变(Example_13)与拓扑不变量(Example_14/15)等;
- 可解释性(Interp):docs/interp.rst 聚焦模型解释,涵盖特征归因、对称性检验、辅助变量、稀疏初始化、Hessian 分析等;
- 物理应用(Physics):docs/physics.rst 展示拉格朗日量学习、守恒律、黑洞度规、本构定律等物理场景;
- 社区贡献:docs/community.rst 收录物理信息 KAN、蛋白质序列分类等社区案例;
- 模块 API 参考:docs/kan.rst 以 Sphinx automodule 方式自动生成
kan.MultKAN、kan.KANLayer、kan.LBFGS、kan.Symbolic_KANLayer、kan.spline、kan.utils、kan.compiler、kan.hypothesis等模块的完整成员文档。
从源码结构看,kan/包还包含 kan/compiler.py(符号表达式到 KAN 的编译器)、kan/hypothesis.py(对称性/可分性检验与树形可视化)、kan/experiment.py(网格搜索实验 runner)、kan/MLP.py(对照实验用 MLP)等模块,为进阶研究与复现实验提供了完整工具链。
六、结语
本文围绕 pykan 官方文档完成了从理论到实战的闭环:Kolmogorov-Arnold 表示定理给出了"多元函数 = 一元函数 + 加法"的构造性证明,pykan 将其落地为边上的 B 样条层(kan/KANLayer.py + kan/spline.py),并通过"初始化 → 稀疏正则训练 → 剪枝 → 符号化 → 精修 → 公式提取"的标准化流程,把黑盒网络还原为人类可读的数学表达式。对于希望快速验证 KAN 效果的研究者与工程师,推荐从本文第三部分的完整代码入手,再结合第四节列出的 API 专题 Notebook 逐项深入;涉及原理级问题的读者,可对照 docs/kan.rst 的模块文档与对应源码文件进一步钻研。
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考