ARTICLE DETAIL

资讯详情

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

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

如何用 pykan 快速完成你的第一次 KAN 函数拟合:完整指南 如何用 pykan 快速完成你的第一次 KAN 函数拟合完整指南【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykanpykan 是 Kolmogorov-Arnold NetworksKAN的 PyTorch 实现。它把可学习的基础函数 样条放在网络的每一条边上让 KAN 函数拟合既准又可读训练完你能直接看到每条边上学到了什么形状甚至抽出符号公式。这篇文章带你走完一个完整任务装环境、备数据、训练、评估、剪枝、抽公式每一步的代码都能直接复制运行。全流程一览上面八个节点对应下文八个小节。只想要最短路径的直接跳到第四节的代码块。上图是一个典型的 pykan KAN 网络每条边挂一个可学习函数线越粗代表这条边的权重越强。 安装 pykan一条命令加一个虚拟环境环境干净与否决定了后面所有我的机器有问题式排查的多少。推荐方式PyPI 一键安装python -m venv .venv source .venv/bin/activate pip install pykanpykan 会连带装好 torch 2.2.2、numpy、matplotlib、scikit-learn 等依赖。GPU 机器上如果 CUDA 报冲突先单独装好对应版本的 PyTorch再执行上面的安装即可。什么时候才用源码安装只有两种情况需要克隆源码后执行pip install -e .你要改源码或者需要 PyPI 尚未发布的版本。普通函数拟合任务用不到。装完后运行from kan import *没有报错就算准备就绪。动手前先看懂 KAN 的两个旋钮装好包别急着跑代码花两分钟理解下面两点能省掉后面大量试错。网格像刻度尺的疏密KAN 与 MLP 的差异就一处激活函数挂在边上而且这个函数本身是参数。每条边的函数 基础函数 × 可学尺度 样条 × 可学尺度。样条铺在网格上grid决定网格被切成几段。类比尺子的刻度刻度稀函数只能学出比较直的形状刻度密才能拟合细弯的曲线。grid3是稳妥起点拟合不够再往上加。值得知道的几个初始化参数KAN构造函数的其余参数基本可以不动完整列表在 kan/MultKAN.pyk3表示三次样条绝大多数任务够用noise_scale0.3控制样条初始注入的噪声是平衡值base_fun选初始基础函数默认silusparse_initTrue会把大部分尺度初始化为零适合做特征选择。这里值得记住作者的一条经验从小模型起步。小任务先试width[输入,5,1]、grid3、不加正则跑不通再逐步加宽最后才考虑加深。KAN 里默认一百宽的 MLP 习惯往往帮倒忙。数据准备create_dataset 的三种用法概念清楚了接下来给模型喂真数据。从公式生成数据最直接的用法是把目标函数直接交给create_datasetfrom kan.utils import create_dataset, create_dataset_from_data # 公式生成每个变量独立范围 标签归一化 dataset create_dataset(f, n_var3, ranges[[-1, 1], [0, 10], [-3, 3]], train_num2000, normalize_labelTrue) # 已有数据自动按 8/2 切分训练测试集 dataset create_dataset_from_data(inputs, labels, train_ratio0.8)几个容易踩的点ranges不传时默认[-1,1]套用到所有变量传(n_var, 2)的列表则每个变量各用各的范围。函数里访问变量默认用列模式x[:,[i]]f_modecol习惯写一维下标的可以改成row。输入量纲差异大比如一个 0~1一个 0~100时建议normalize_inputTrue把输入拉进[-1,1]附近恰好落在样条的初始网格范围内训练更稳。两个函数返回的都是同一个结构的字典train_input、train_label、test_input、test_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_var2, devicedevice) model KAN(width[2,5,1], grid3, k3, seed42, devicedevice) model(dataset[train_input]) # 前向一次顺便初始化网格 model.fit(dataset, steps50, lamb0.001) model.evaluate(dataset)官方 notebook 默认用 float64样条拟合在 float32 下误差偏大建议照做。fit里有三个开关值得先懂stepsLBFGS 迭代次数默认 100先跑 50 看趋势即可。lambL1 惩罚默认 0 是纯拟合给 0.001 开始把网络往稀疏方向推是可解释性的主阀门。update_grid默认 True训练中会周期性按数据分布重算网格点数据分位数与均匀网格之间由grid_eps插值。想手动控制网格时关掉。最常动的五个参数参数什么时候动动了之后grid欠拟合形状学不出来样条控制点变多拟合能力增强过大则过拟合并拖慢训练k需要更光滑的边样条阶数升高3 阶之外收益递减lamb想要稀疏可读的网络弱边被压到零结构变稀疏update_grid网格和数据范围对不上False 时冻结初始网格交给你自己调width调完 grid 仍欠拟合先加宽、后加深宽度优先级更高evaluate返回train_loss和test_loss。两者差距明显拉大就是过拟合信号先降grid比降width更见效再考虑加数据或加大lamb。 读出模型学到了什么plot 与符号公式损失数字只是半个答案KAN 的真正卖点是看得见model.plot(in_vars[x,y], out_vars[f]) # 每条边的函数形状 model.symbolic_formula(varx) # 抽取紧凑符号公式 model.speed() # 关掉符号分支提速plot展示每条边的激活函数形状和边权线宽。哪条边在某区间接近直线就说明它在那个区域基本学到了线性关系哪条边几乎平了说明它没用。符号分支默认开启symbolic_enabledTrue但不开并行化不用它的场景下会拖慢前向。README 专门提醒自己写训练循环且不做符号回归时训练前先调一次model.speed()。 精简网络剪枝与状态回滚刚训完的网络常常多到不需要。剪枝的标准流程是先稀疏、再剪、再补训model.fit(dataset, steps50, lamb0.01) # 更强的 L1 推稀疏 model.prune(node_th1e-2, edge_th3e-2) # 剪掉弱节点和弱边 model.fit(dataset, steps20) # 剪枝后补训一轮阈值越小剪得越狠从 1e-2 量级起步比较安全。prune_input是另一类操作按输入重要性把贡献极小的输入变量整个移除相当于让模型自己告诉你哪些特征没用。想查某个输入到底贡献多少用model.attribute()。剪坏了也不需要从头再来auto_save默认开启训练过程中会自动把检查点存进./model目录model.rewind(0)就能回到任何一版状态对比不同稀疏度也很方便。️ 常见坑速查表跑完全流程多半会撞上下面几类问题症状常见原因对策训练明显偏慢符号分支开着但没用到fit前调用model.speed()训练/测试损失差距大过拟合网格过大先降grid再降width或加数据、加大lamb输出出现 nan目标函数含 1/x、log 等奇异点fit和forward传singularity_avoidingTrue两次运行结果不一致种子或精度没固定固定KAN(seed...)与create_dataset(seed...)统一 float64下一步去哪看上面这条流水线是最短上手路径仓库里还有成体系的材料hellokan.ipynb 与本文代码一一对应带逐步输出适合逐行对照。tutorials/ 下分四类Example 是函数拟合与 PDE 的经典案例深公式发现、奇点、相位转变API_demo 逐条讲 APIInterp 是可解释性技巧换边、Hessian、稀疏初始化Physics 是科学应用拉格朗日量、黑洞、本构方程。分类任务先看 Example_4数值精度看 Example_7剪枝与符号回归的 API 想系统掌握Interp_3_KAN_Compiler.ipynb 更集中。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表