pykan 文档导览与快速上手:Kolmogorov-Arnold Networks 的安装、入门与教程体系
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
Kolmogorov-Arnold Networks(KAN)是一种以 Kolmogorov-Arnold 表示定理为数学基础、把可学习激活函数放在边(edge)上的新型神经网络架构,本仓库(pykan)是该论文的开源实现。本文以docs/index.rst(Sphinx 文档主入口)为骨架,完整梳理 pykan 的安装方式、依赖环境、入门路径与文档导航,并结合仓库源码与intro.rst快速上手示例,帮助你从零开始跑通一个 KAN 模型的完整生命周期:初始化 → 训练 → 剪枝 → 符号化 → 提取公式。
项目概览:KAN 与 MLP 的本质差异
pykan 是论文 "KAN: Kolmogorov-Arnold Networks" 的开源仓库(对应论文 DOI 信息可查README.md)。文档开篇即点明其核心定位:
Kolmogorov-Arnold Networks, inspired by the Kolmogorov-Arnold representation theorem, are promising alternatives of Multi-Layer Preceptrons (MLPs). KANs have activation functions on edges, whereas MLPs have activation functions on nodes.
KAN 与 MLP 是"对偶"关系:
- MLP在线性层之间用非线性激活函数
σ作用在节点上:MLP(x) = W_{L-1} ∘ σ ∘ ⋯ ∘ W_1 ∘ σ ∘ W_0 ∘ x; - KAN把可学习的 1D 函数放置在边上,网络是 KAN 层的堆叠:
KAN(x) = Φ_{L-1} ∘ ⋯ ∘ Φ_1 ∘ Φ_0 ∘ x,其中每个Φ都是一个"函数矩阵"(Kolmogorov-Arnold 层)。
这一简单改变使 KAN 在精度与可解释性两方面都具备优势:KAN 层可以可视化为全连接图,每条边上是一条可学习的样条曲线,训练后可以直接"读"出函数形态,并通过符号回归自动还原成数学公式。这也是本仓库从文档到源码始终强调的两大关键词——accuracy 与 interpretability。
文档主入口的首页配图(见 docs/kan_plot.png)直观展示了 KAN 训练后各边激活函数的可视化形态,这正是"边上的激活函数"这一核心设计的直观体现。
安装方式:GitHub 源码安装与 PyPI 安装
docs/index.rst提供了两种官方安装途径,与README.md保持一致。
方式一:从 GitHub 源码安装(推荐开发场景)
git clone https://github.com/KindXiaoming/pykan.git cd pykan pip install -e . # pip install -r requirements.txt # install requirements使用pip install -e .(可编辑安装)时,setup.py会将包注册为pykan(版本号0.2.8,python_requires='>=3.6'),packages=setuptools.find_packages()会自动收集kan/下的源码模块,include_package_data=True会把assets/img/sum_symbol.png、assets/img/mult_symbol.png等绘图所需的符号资源一并打包。
方式二:从 PyPI 安装
pip install pykan此外README.md还提供了面向开发者的快速通道:pip install git+https://github.com/KindXiaoming/pykan.git,以及可选的 Conda 环境方案(conda create --name pykan-env python=3.9.7后激活再执行上述安装命令)。
依赖环境要求
docs/index.rst中列出的核心依赖(与仓库根目录 requirements.txt 完全对应)如下:
| 依赖包 | 版本 | 用途 |
|---|---|---|
| python | 3.9.7(建议) | 解释器版本,README 注明 "Python 3.9.7 or higher" |
| matplotlib | 3.6.2 | 网络可视化与激活函数绘图 |
| numpy | 1.24.4 | 数值计算 |
| scikit_learn | 1.1.3 | 回归指标(R²)等辅助计算 |
| setuptools | 65.5.0 | 包安装 |
| sympy | 1.11.1 | 符号公式提取与 LaTeX 输出 |
| torch | 2.2.2 | 深度学习后端 |
| tqdm | 4.66.2 | 训练进度条 |
| pandas / seaborn / pyyaml | 2.0.1 等 | 数据整理、统计绘图与 checkpoint 配置读写(README 补充项) |
安装时可在激活虚拟环境后执行pip install -r requirements.txt精确锁定上述版本。
入门路径:从 Hello KAN 到专题教程的文档地图
docs/index.rst的 "Get started" 一节为读者规划了清晰的学习路径:
- Quickstart(快速上手):
docs/intro.rst(即 "Hello, KAN!"); - KANs in Action(动手实战):
docs/demos.rst(API 演示)与docs/examples.rst(完整示例); - API(高级接口):
docs/modules.rst。
文档目录(toctree)包含七个部分,对应仓库中的实际文档文件:
| 文档模块 | 相对路径 | 内容定位 |
|---|---|---|
| 入门 | docs/intro.rst | KAN 数学原理与 Hello KAN 完整演练 |
| API | docs/modules.rst | 源码级 API 文档(kan包) |
| API 演示 | docs/demos.rst | 12 个 API 专题(索引、绘图、激活提取、初始化、网格、训练超参、剪枝、正则化、视频、设备、数据集、checkpoint) |
| 示例 | docs/examples.rst | 15 个示例(函数拟合、深层公式、分类、特殊函数、PDE、持续学习、奇点、相对论速度叠加、无监督学习、相变、结分类等) |
| 可解释性 | docs/interp.rst | 12 个可解释性专题(MultKAN、KAN Compiler、特征归因、对称性检验、Hessian、稀疏初始化等) |
| 物理 | docs/physics.rst | 7 个物理应用(拉格朗日量、守恒律、黑洞、本构律) |
| 社区 | docs/community.rst | 社区贡献教程(物理信息 KAN、蛋白质序列分类) |
这些.rst文档与 tutorials 目录下的.ipynbNotebook 一一对应,实际动手时可直接打开 Notebook 逐格运行。
快速上手:跑通一个 KAN 的完整生命周期
以 docs/intro.rst 中的 Quickstart 为例,完整演示从初始化到提取符号公式的全流程。这也是仓库根目录 hellokan.ipynb 的核心内容。
1. 初始化 KAN
from kan import * # create a KAN: 2D inputs, 1D output, and 5 hidden neurons. # cubic spline (k=3), 5 grid intervals (grid=5). model = KAN(width=[2,5,1], grid=5, k=3, seed=0)kan/__init__.py通过from .MultKAN import *导出核心类MultKAN(KAN 的别名)。从 kan/MultKAN.py 的构造签名可以看到完整初始化参数:width(各层神经元数)、grid(网格区间数,默认 3)、k(样条阶数,默认 3)、noise_scale(样条初始噪声,默认 0.3)、base_fun(残差基函数,默认'silu')、grid_eps(均匀网格与自适应网格的插值系数,默认 0.02)、grid_range(网格范围,默认 [-1,1])、seed(随机种子,默认 1)、device(默认 'cpu')等。每条边上的激活函数φ(x) = sb_scale * b(x) + sp_scale * spline(x),即"基函数 + 样条"的组合。
2. 创建合成数据集
# create dataset 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,默认train_num=1000、test_num=1000、输入范围[-1,1];支持f_mode='col'/'row'两种标签计算方式,以及normalize_input/normalize_label标准化选项。返回的字典含train_input、train_label、test_input、test_label四个键,后续model.train(dataset)直接消费。
3. 初始化可视化与训练
# plot KAN at initialization model(dataset['train_input']); model.plot(beta=100) # train the model 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初始化状态下各边近似一条直线(见 docs/intro_files/intro_15_0.png);训练时默认使用仓库自带的 L-BFGS 优化器(kan/LBFGS.py),lamb控制 L1 稀疏正则强度,lamb_entropy控制熵正则以鼓励激活函数稀疏化。
训练 20 步后,稀疏正则使得大量边被"压平"(接近常函数),网络结构趋于稀疏可解释(见 docs/intro_files/intro_19_0.png)。
4. 剪枝:从稀疏结构到更小模型
model.prune() # 原地剪枝(保留原形状) model.plot(mask=True) # 用 mask 隐藏被剪掉的边 model = model.prune() # 返回更小形状的剪枝模型 model(dataset['train_input']) model.plot()prune实现在 kan/MultKAN.py,默认阈值node_th=1e-2、edge_th=3e-2:低于阈值的节点与边会被移除,网络自动收缩为更紧凑的形状,这正是 KAN "可解释性即模型压缩"的体现。
5. 符号回归:自动还原数学公式
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)auto_symbolic(kan/MultKAN.py)会为每条边遍历符号库、以 R² 为指标拟合最优符号函数(输出如fixing (0,0,0) with sin, r2=0.999987252534279)。符号库SYMBOLIC_LIB定义在 kan/utils.py,包含x、x^2、sqrt、exp、log、sin、cos、tan、tanh、abs、gaussian、倒数族等 20 余个函数,每个条目都提供数值实现、sympy 符号实现、拟合参数量与数值稳定性裁剪函数。
继续训练至近机器精度后,用model.symbolic_formula()输出最终公式:
model.train(dataset, opt="LBFGS", steps=50) # train loss: 2.02e-10 | test loss: 1.13e-10 model.symbolic_formula()[0][0] # 1.0 * exp(1.0 * x_2^2 + 1.0 * sin(3.14 * x_1))网络自主还原出了f(x,y) = exp(sin(πx) + y²),验证了 KAN "先学结构、再符号化、最终得到闭式解"的完整可解释性工作流。
进阶技巧:Efficiency 模式与超参调优
关闭符号分支提升训练速度
README.md特别强调:当你(1)需要自己写训练循环而非使用model.fit(),(2)不使用符号分支时,务必在训练前调用model.speed()(实现见 kan/MultKAN.py)。否则符号分支默认开启,而符号计算未做并行化,会严重拖慢训练。
面向科学问题的调参建议(摘自 README)
- 从极简配置起步:小 width、小 grid、小数据、无正则(
lamb=0)。例如 5 输入 1 输出任务可先试KAN(width=[5,1,1], grid=3, k=3),不收敛时先加宽、再加深——这与 MLP 文献默认O(10²)宽度的习惯截然不同。 - 追求精度:优先使用网格扩展(grid extension)技术;但要警惕过拟合。
- 追求可解释性:用
model.train(lamb=0.01)稀疏化,逐步增大lamb;剪枝出明显无用的神经元后调用model.prune()获得紧凑模型。 - 判断拟合状态:train/test loss 差距大说明过拟合,优先减小
grid(比减小width更有效),再考虑减小width。 - 收尾:模型表现良好后,增大数据量做一次最终训练,通常还能进一步提升效果。
计算资源需求
README.md说明:tutorials 中的示例在单 CPU 上通常 10 分钟内可跑完;论文中全部示例单 CPU 一天内可完成。KAN 训练 PDE 是最昂贵的场景,单 CPU 可能需要数小时到数天。仓库示例规模偏向科研任务,任务规模较大时建议使用 GPU(模型支持device参数,可参考 docs/API_demo/API_10_device.rst)。
结语
从docs/index.rst这张文档地图出发,你可以按"安装 → Hello KAN → API 演示 → 专题示例 → 可解释性 → 物理应用"的路径系统掌握 pykan。KAN 的核心心智模型是:把可学习函数放在边上,先以样条拟合任意函数,再通过稀疏正则、剪枝与符号回归将网络"翻译"成人类可读的数学公式——这正是其在精度与可解释性上优于 MLP 的根本原因。建议从 hellokan.ipynb 开始逐格运行,再按需查阅 docs/demos.rst 与 docs/examples.rst 中的专题教程。
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考