十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

pykan 实战指南:Kolmogorov-Arnold Networks(KAN)安装、训练与超参数调优全解析

pykan 实战指南:Kolmogorov-Arnold Networks(KAN)安装、训练与超参数调优全解析 pykan 实战指南Kolmogorov-Arnold NetworksKAN安装、训练与超参数调优全解析【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan导读本文以 pykan 仓库的 README 为主线系统讲解 Kolmogorov-Arnold NetworksKAN的数学动机、安装方式、训练流程含model.fit()与model.speed()等核心 API、可解释性与稀疏化技巧并给出作者基于论文实验经验总结的超参数调优方法论。读完本文你将能够从零搭建并训练一个 KAN 模型掌握网格扩展、剪枝、符号回归等进阶能力并理解 KAN 与 MLP 在架构设计上的本质差异。一、KAN 是什么与 MLP 对偶的新网络架构Kolmogorov-Arnold NetworksKAN被定位为多层感知机MLP的替代方案。二者的数学根基一一对应MLP 建立在通用逼近定理之上而 KAN 建立在 Kolmogorov-Arnold 表示定理之上。两者互为对偶——MLP 把激活函数放在节点神经元上KAN 则把激活函数放在边连接上。这一看似微小的改变带来了模型**精度accuracy与可解释性interpretability**两方面的收益。在当前仓库中KAN 的核心实现位于 kan/MultKAN.py 的MultKAN类文档示例中通常以KAN的名义使用例如from kan import *后直接model KAN(width[2,5,1], grid5, k3, seed0)。每个 KAN 层由可学习的一维样条激活函数组成激活函数形如phi(x) sb_scale * b(x) sp_scale * spline(x)其中b(x)是残差基函数默认 SiLU样条部分则承载了主要的拟合能力。二、2024-07-14 重大更新API 变更一览README 顶部特别标注了 2024 年 7 月 14 日的破坏性更新这是阅读旧教程时必须注意的版本分水岭model.train()更名为model.fit()训练入口统一为 fit 方法源码签名中的默认参数如optLBFGS、steps100、lamb0.均可在调用时覆盖create_dataset迁移到kan.utils合成数据集工具函数位于 kan/utils.py不再从模型对象上直接调用教程已全部验证可在 CPU 上运行./tutorials下的 Notebook 与更新后的 API 保持同步而文档docs尚未完全同步更新。此外README 提到 PyPI 最新版本为 0.2.1以当前仓库为准setup.py 中pykan的版本号已推进到0.2.8python_requires3.6。新增功能一览新功能说明仓库中的对应教程/源码乘法节点MultKANKAN 中引入乘法运算宽度列表可写成[[n0,m0],[n1,m1],...]形式tutorials/Interp/Interp_1_Hello, MultKAN.ipynb速度模式不使用符号分支时调用model.speed()大幅加速kan/MultKAN.py符号公式编译将符号公式编译进 KANtutorials/Interp/Interp_3_KAN_Compiler.ipynb特征归因与输入剪枝计算节点/边/子节点的归因分数剪掉无用的输入tutorials/Interp/Interp_4_feature_attribution.ipynb三、安装指南PyPI / GitHub / 开发模式 / Conda前置要求Python 3.9.7 或更高版本 pip方式一PyPI 安装推荐大多数用户pip install pykan方式二直接从 GitHub 安装pip install githttps://github.com/KindXiaoming/pykan.git方式三开发模式安装适合贡献代码或修改源码git clone https://github.com/KindXiaoming/pykan.git cd pykan pip install -e .方式四Conda 环境可选conda create --name pykan-env python3.9.7 conda activate pykan-env pip install githttps://github.com/KindXiaoming/pykan.git # GitHub 安装 # 或 pip install pykan # PyPI 安装依赖清单requirements.txt 给出了精确的依赖版本matplotlib3.6.2 numpy1.24.4 scikit_learn1.1.3 setuptools65.5.0 sympy1.11.1 torch2.2.2 tqdm4.66.2 pandas2.0.1 seaborn pyyaml激活虚拟环境后执行pip install -r requirements.txt注意README 原始版本清单以python3.9.7为基准当前仓库的 requirements.txt 已补充pandas、seaborn、pyyamlMultKAN的模型存档机制使用 yaml 保存配置绘图与数据处理用到 pandas/seaborn。四、计算资源需求单 CPU 即可运行仓库明确说明见 README.mdtutorials中的示例在单 CPU 上通常 10 分钟以内即可跑完论文中的所有示例在单 CPU 上一天以内可复现训练 KAN 求解 PDE 是开销最大的场景单 CPU 上可能需要数小时到数天作者之所以用 CPU 而非 GPU是因为需要做 MLP 与 KAN 的大规模参数扫描Pareto 前沿涉及数千个小模型如果你的任务规模较大超出科学计算类任务的典型规模建议改用 GPU。五、快速上手从 hellokan 到第一个模型5.1 起步 Notebook官方快速入门入口是 hellokan.ipynb更多示例分布在 tutorials 目录按主题划分为API_demo、Example、Interp、Community、Physics等子目录。5.2 用 kan.utils 合成数据集create_dataset已经从模型迁移到工具层其完整签名kan/utils.pydef create_dataset(f, n_var2, f_modecol, ranges[-1,1], train_num1000, test_num1000, normalize_inputFalse, normalize_labelFalse, devicecpu, seed0)关键参数说明f用于生成合成数据的符号函数n_var输入变量个数ranges每个输入变量的采样区间默认[-1,1]也支持形状(n_var, 2)的逐变量区间train_num/test_num训练/测试样本数默认各 1000normalize_input/normalize_label是否对输入/标签做标准化返回字典包含train_input、train_label、test_input、test_label四个键。典型用法from kan import * f lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) x[:,[1]]**2) dataset create_dataset(f, n_var2, train_num100) model KAN(width[2,5,1], grid5, k3, seed0)5.3 训练model.fit()当前版本统一使用model.fit()而非旧的model.train()核心签名与参数kan/MultKAN.pymodel.fit(dataset, optLBFGS, steps100, log1, lamb0., lamb_l11., lamb_entropy2., lamb_coef0., lamb_coefdiff0., update_gridTrue, grid_update_num10, loss_fnNone, lr1., start_grid_update_step-1, stop_grid_update_step50, batch-1, metricsNone, save_figFalse, ...)opt优化器LBFGS默认内部使用带strong_wolfe线搜索的 L-BFGS或Adamlamb整体正则强度默认为 0lamb_l1/lamb_entropy/lamb_coef/lamb_coefdiff分别控制 L1 稀疏、熵、系数幅值、相邻系数平滑度惩罚update_grid训练中定期更新网格默认开启grid_update_num次更新分布在stop_grid_update_step步内这是 KAN 网格扩展技术的实现基础batch批量大小-1表示全量返回的results字典包含train_loss、test_loss、regRMSE 形式以及自定义metrics的序列。5.4 精度提示如果追求高精度README 建议设置默认浮点类型为 float64torch.set_default_dtype(torch.float64)六、精度与可解释性KAN 的两大卖点6.1 精度AccuracyREADME 给出的结论性描述为KAN 的缩放scaling速度快于 MLP用更少的参数获得比 MLP 更好的精度。论文中的四类典型精度场景包括拟合符号公式symbolic formulas拟合特殊函数special functionsPDE 求解避免灾难性遗忘continual learning对应的可复现教程见 tutorials/Example/Example_1_function_fitting.ipynb、tutorials/Example/Example_5_special_functions.ipynb、tutorials/Example/Example_6_PDE_interpretation.ipynb、tutorials/Example/Example_8_continual_learning.ipynb。6.2 可解释性InterpretabilityKAN 可以被直观地可视化model.plot()并提供 MLP 无法给出的可交互性理论上可用于辅助发现新的科学定律。README 展示了四类场景符号公式提取训练后通过model.suggest_symbolic(l, i, j)对每条边拟合符号函数发现纽结knots的数学规律见 tutorials/Example/Example_14_knot_supervised.ipynb发现 Anderson 局域化的物理规律三层 KAN 的训练过程可视化。从源码结构看符号能力由 kan/utils.py 中的SYMBOLIC_LIB支撑内置了x、x^2、1/x、sqrt、exp、log、sin、cos、tanh、arcsin、gaussian等约 30 个候选符号函数每个条目同时提供 torch 实现、sympy 实现、复杂度权重与奇点保护实现suggest_symbolic与auto_symbolic正是在该库上做参数扫描与拟合。七、新特性源码级解读7.1 乘法节点MultKANMultKAN.__init__kan/MultKAN.py支持两种宽度描述无乘法节点width[2,5,5,3]表示 2 维输入、3 维输出、两层各 5 个隐藏神经元有乘法节点width[[2,0],[5,3],[5,1],3]表示在对应层额外引入 3 个 / 1 个乘法节点。mult_arity控制每个乘法节点相乘的输入个数可以是整数此时mult_homoTrue可并行或整数列表的列表如[[2,3],[4]]逐节点指定需要 for 循环。7.2 速度模式model.speed()当满足两个条件时务必在训练前调用model.speed()kan/MultKAN.py需要自己手写训练循环而不是使用model.fit()完全不使用符号分支symbolic branch。原因符号分支的计算没有被并行化开着会非常慢。speed()的实现即把symbolic_enabledFalse、save_actFalse、auto_saveFalse若传compileTrue则额外返回torch.compile(self)的编译版本。7.3 符号公式编译KAN Compiler可以把已知的符号公式编译进 KAN 结构作为先验知识注入教程见 tutorials/Interp/Interp_3_KAN_Compiler.ipynb。这与auto_symbolic自动把训练好的边替换为符号函数形成双向闭环先验符号 → 编译进网络 → 训练 → 提取符号。7.4 特征归因与输入剪枝kan/MultKAN.py 提供prune_input(threshold1e-2)配合归因分数可以识别并剪掉无用的输入特征教程见 tutorials/Interp/Interp_4_feature_attribution.ipynb。八、超参数调优建议作者经验谈README 明确提醒关于 MLP 的很多直觉不能直接迁移到 KAN。以下是作者基于论文中规模较小、面向科学计算问题调参的经验总结8.1 从最简单的设置起步从小规模开始小 KAN 形状、小网格、小数据、无正则lamb0。这与 MLP 文献默认宽度 10^2 量级的做法截然不同例如 5 输入 1 输出的任务先尝试KAN(width[5,1,1], grid3, k3)不行则先加宽width再加深depth在调试阶段选小数据集是为了获得更快的反馈其隐含假设是小数据与大数据的定性行为相似在小规模问题上通常成立但并非普遍真理。8.2 达到可接受精度后逐步精炼追求精度尝试网格扩展grid extension技术fit(update_gridTrue)默认开启示例见 tutorials/Example/Example_1_function_fitting.ipynb但注意过拟合风险追求可解释性用稀疏化训练例如model.fit(dataset, lamb0.01)并建议逐步增大lamb训练后调用model.plot()观察若发现明显无用的神经元可执行pruned_model model.prune()得到剪枝模型之后继续训练鼓励精度或稀疏或做符号回归剪枝相关方法的默认阈值可从源码确认prune_node(threshold1e-2)、prune_edge(threshold3e-2)、prune(node_th1e-2, edge_th3e-2)。8.3 精度与可解释性并非必然矛盾作者强调精度、可解释性以及参数效率不一定互相冲突论文 Figure 2.3 展示了二者正相关或存在权衡的不同情形建议一次只专注一个目标。但如果你有充分理由相信剪枝可解释性也能提升精度可以提前规划先推可解释性再推精度。8.4 收尾得到相当好的结果后尝试增大数据量做最终一轮训练通常会得到更好的结果。九、欠拟合 vs 过拟合如何判断调参的核心是时刻确认模型处于哪种状态若训练损失与测试损失差距很大说明过拟合优先增大数据量或减小模型减小模型时注意grid比width更重要所以先减小grid再减小width作者建议从简单模型开始先确保模型处于欠拟合区再逐步扩张到金发姑娘区Goldilocks zone即恰到好处的拟合区间。十、引文与联系方式论文引用格式BibTeXarticle{liu2024kan, title{KAN: Kolmogorov-Arnold Networks}, author{Liu, Ziming and Wang, Yixuan and Vaidya, Sachin and Ruehle, Fabian and Halverson, James and Solja{\v{c}}i{\c}, Marin and Hou, Thomas Y and Tegmark, Max}, journal{arXiv preprint arXiv:2404.19756}, year{2024} }如有疑问可联系作者邮箱 zmliumit.edu。十一、作者的话KAN 的定位与边界README 末尾的作者附言值得所有使用者阅读它划定了 KAN 的适用边界仓库最初面向科学发现与科学计算用户代码没有过多考虑效率与复用性优化官方也明确推荐了 efficientkan、FourierKAN 等致力于效率提升的社区实现本文不展开外部链接对于机器学习方向的用户作者坦诚 KAN目前还不是开箱即用的即插即用方案超参数需要调试且不同应用有各自的专属技巧例如在强化学习中固定部分可训练参数以提升稳定性、在图任务中于潜空间使用 KAN 等关于KAN 是否会成为下一代大语言模型作者表示没有可靠直觉。KAN 面向的是高精度和/或高可解释性的场景而精度可解释性在 LLM 与科学计算中的含义差别很大不能直接套用论文结论作者欢迎对 KAN 的批评也提醒对批评本身保持批判实践是检验真理的唯一标准。KAN 与 MLP 在作者看来互相不可替代各有优势场景与局限。一句话总结在动手训练之前请先确认自己的场景是否需要 KAN 所擅长的高精度与可解释性如果是就从KAN(width[5,1,1], grid3, k3)这样的简单模型开始遵循先欠拟合、再逐步扩张的调参路线并在最终阶段结合model.speed()、网格扩展、稀疏化与剪枝把模型打磨到最佳状态。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表