Kolmogorov-Arnold 网络(KAN)实战入门:基于 pykan 的安装配置、模型训练与可解释性建模全指南
2026/9/14 10:39:51 网站建设 项目流程

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_batchcoef2curve完成样条求值,从而实现"边上的一维函数"。

二、安装与依赖配置

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 会导出MultKANutils中的全部符号,因此安装后可直接from kan import *使用。

2.2 从 PyPI 安装

pip install pykan

2.3 依赖版本要求

官方文档给出的核心依赖清单(按 requirements.txt 校验)如下:

依赖包版本用途
matplotlib3.6.2KAN 可视化绘图(model.plot()
numpy1.24.4数值计算
scikit_learn1.1.3指标计算、实验工具
setuptools65.5.0打包安装
sympy1.11.1符号公式推导(symbolic_formula
torch2.2.2深度学习框架与自动求导
tqdm4.66.2训练进度条
pandas2.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.0scale_base_sigma=1.0:基础函数(base function)的初始化分布参数;
  • base_fun='silu':基础激活函数(可传入字符串或可调用对象);
  • grid_eps=0.02grid_range=[-1,1]:网格更新时的边界缓冲与网格取值范围;
  • seed=1:随机种子(本例显式指定seed=0保证可复现);
  • device='cpu':运行设备(详见 API_10_device.ipynb);
  • 以及sp_trainablesb_trainableauto_saveckpt_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_inputtrain_labeltest_inputtest_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]

trainfit方法(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_nodeprune_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 支持把学习到的样条函数替换为已知的符号函数(如sinx^2exp),这是其可解释性能力的核心体现。

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.ipynbmodel[l][i][j]get_funget_act
绘图API_2_plotting.ipynbplotattribute
激活提取API_3_extract_activations.ipynbget_act
初始化API_4_initialization.ipynbinitialize_from_another_model
网格细化API_5_grid.ipynbrefine(new_grid)(kan/MultKAN.py)
训练超参数API_6_training_hyperparameter.ipynbfitlamb/lamb_entropy/opt
剪枝API_7_pruning.ipynbpruneprune_nodeprune_edgeprune_input
正则化API_8_regularization.ipynbregget_regdisable_symbolic_in_fit
训练视频API_9_video.ipynbfit(..., save_fig=True)
设备API_10_device.ipynbdevice='cuda'model.to(device)
数据集API_11_create_dataset.ipynbcreate_datasetcreate_dataset_from_data
断点保存/加载API_12_checkpoint_save_load_model.ipynbsaveckptloadckptrewindcheckout

符号化训练完成后,还可以再次绘制以观察最终结构——此时每条保留的边都对应一个已确认的符号函数,网络从"样条黑盒"彻底转变为透明公式:

五、文档生态导航:从入门到研究前沿

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.MultKANkan.KANLayerkan.LBFGSkan.Symbolic_KANLayerkan.splinekan.utilskan.compilerkan.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),仅供参考

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

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

立即咨询