news 2026/9/11 10:32:05

如何用 pykan 快速完成你的第一次 KAN 函数拟合:完整指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
如何用 pykan 快速完成你的第一次 KAN 函数拟合:完整指南

如何用 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 pykan

pykan 会连带装好 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选初始基础函数,默认silusparse_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_inputtrain_labeltest_inputtest_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_losstest_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 等奇异点fitforwardsingularity_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 10:32:02

让AMD显卡跑CUDA程序:ZLUDA从源码构建到预编译缓存的4步完整指南

让AMD显卡跑CUDA程序:ZLUDA从源码构建到预编译缓存的4步完整指南 【免费下载链接】ZLUDA CUDA on non-NVIDIA GPUs 项目地址: https://gitcode.com/GitHub_Trending/zl/ZLUDA 给NVIDIA显卡编译的程序,换到AMD显卡上就完全跑不起来,这是…

作者头像 李华
网站建设 2026/9/11 10:27:56

Transformers与LoRA微调实战:从迁移学习到Qwen模型落地

1. 迁移学习不是玄学:先把“落地”这件事想清楚 1.1 从预训练到微调:迁移学习的工程化表达 做 NLP 这行超过十年的人,应该都经历过那个“每个任务都要从零开始训练模型”的年代。那时候做个文本分类,你得准备几百万条标注数据&am…

作者头像 李华
网站建设 2026/9/11 10:27:46

NR2047/A-47停产背后的语音芯片代际升级指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 10:26:07

gvim复制粘贴详解:寄存器、系统剪贴板与vimrc配置实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 10:22:53

图书管理系统开发实战:从SpringBoot表设计到部署上线全流程

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华