资讯动态

PyKAN 怎么通过 base_fun=‘identity‘ 与 lamb_coef 惩罚让深层 KAN 的激活函数趋向线性

发布时间:2026/9/15 21:51:59 来源:尧图企业网站定制
PyKAN 怎么通过 base_funidentity 与 lamb_coef 惩罚让深层 KAN 的激活函数趋向线性【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan当不确定 KAN 该设多深时pykan 官方示例Example 11: Encouraging linearity给出的策略之一是先建一个足够深的模型训练后再把冗余部分剪掉。但沿宽度方向的稀疏化还不够——如果某些深度方向上其实不需要非线性希望对应的激活函数退化成线性即“shortcut along depth”。这个示例演示了两个配合使用的技巧把 base functionbase_fun设为线性惩罚 spline 系数。当 spline 系数为零时激活函数就是线性的。本文按官方 Example_11_encouraing_linear.ipynb 的步骤完整走一遍先跑一个不带技巧的对照模型再开启两个技巧最后用训练日志和model.plot()对比激活函数的变化。base_fun 与 spline 系数各管哪一部分在 pykan 中每条边上的激活函数由残差函数和样条两部分组成见 MultKAN.py 的参数说明phi(x) sb_scale * b(x) sp_scale * spline(x)其中b(x)就是base_fun。MultKAN.__init__支持三个取值默认是silusilu→torch.nn.SiLU()默认identity→torch.nn.Identity()zero→lambda x: x*0.所以把base_fun设为identity相当于把公式里的b(x)直接换成恒等函数激活的非线性只剩 spline 一项。再配合训练时对 spline 系数的惩罚系数被压到接近零时整个激活函数就趋向线性。lamb_coef正是fit里控制这一项惩罚的参数源码 reg() 中的实现是对各层样条系数取 L1 后按lamb_coef * coeff_l1加进正则项。准备环境并定义任务示例代码依赖kan包仓库内kan/目录from kan import *导入和 torch。官方示例用单变量任务f(x) sin(πx)文档明确说明一个[1,1]KAN 就足够完成它但为了模拟“不知道该用多深的网络”的情况示例故意使用[1,1,1,1]的过深网络。对照模型不带技巧的完整代码与示例 Notebook 保持一致from kan import * device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) # create dataset f(x,y) sin(pi*x). This task can be achieved by a [1,1] KAN f lambda x: torch.sin(torch.pi*x[:,[0]]) dataset create_dataset(f, n_var1, devicedevice) model KAN(width[1,1,1,1], grid5, k3, seed0, noise_scale0.1, devicedevice) model.fit(dataset, optLBFGS, steps20);说明device按环境自动选择 CUDA 或 CPU模型auto_save默认为真训练会在当前目录创建./model并保存检查点示例输出中的checkpoint directory created: ./model、saving model version 0.0即来自此。打开两个技巧训练深层 KAN带技巧的版本只改两处构造KAN时传base_funidentity调用fit时传lamb1e-4, lamb_coef10.0。示例 Notebook 的原文如下f、dataset、device沿用上一节已定义的对象from kan import * # create dataset f(x,y) sin(pi*x). This task can be achieved by a [1,1] KAN f lambda x: torch.sin(torch.pi*x[:,[0]]) dataset create_dataset(f, n_var1, devicedevice) # set base_fun to be linear model KAN(width[1,1,1,1], grid5, k3, seed0, base_funidentity, noise_scale0.1, devicedevice) # penality spline coefficients model.fit(dataset, optLBFGS, steps20, lamb1e-4, lamb_coef10.0);参数用途依据 MultKAN.py 的fit文档串lamboverall penalty strength整体惩罚强度lamb_coefcoefficient magnitude penalty strength系数幅度惩罚强度本例用 10.0另有lamb_coefdiff惩罚相邻系数差值smoothness本例保持默认 0。注意这两个技巧是配套使用的只设base_funidentity不惩罚系数spline 部分仍可以学出强非线性只加lamb_coef惩罚而base_fun保持默认的silub(x)那一项也仍是非线性的。两者一起才让激活函数整体趋向线性。验证训练日志与激活函数图文档给出的判断方式是观察训练日志和激活函数图而不是比较固定数值。以下是文档记录的示例输出不同硬件/环境下数值会不同不要当作固定预期不带技巧cuda checkpoint directory created: ./model saving model version 0.0 | train_loss: 3.74e-04 | test_loss: 3.84e-04 | reg: 8.88e00 | : 100%|█| 20/20 [00:0500:00, 3.79it saving model version 0.1带技巧checkpoint directory created: ./model saving model version 0.0 | train_loss: 8.89e-03 | test_loss: 8.40e-03 | reg: 1.83e01 | : 100%|█| 20/20 [00:0400:00, 4.20it saving model version 0.1图中示例的reg项明显更大1.83e01 对 8.88e00这是lamb_coef惩罚生效在日志上的直观痕迹train_loss略高则体现了为线性化付出的拟合代价。激活函数是否真的趋向线性用model.plot()看。默认调用即可beta控制每条激活的透明度transparency tanh(beta*l1)示例中为了把线性化效果画得更清晰使用了beta10model.plot() # 示例文档中使用: model.plot(beta10)plot会把 PNG 存到默认目录./figures。官方文档中两组结果对比均为文档示例图判读方法对比两张图中各层边上的曲线带技巧后曲线更接近直线说明冗余深度的激活函数已经“短路”成线性为后续沿深度方向的剪枝示例开头提到的从大模型剪下来的策略提供了依据。边界与限制lamb1e-4与lamb_coef10.0是官方示例[1,1,1,1]拟合sin(πx)、LBFGS 20 步中使用的取值文档没有给出其他任务的推荐值换任务时需要按拟合损失与线性化程度自行权衡。技巧的代价是拟合损失变差文档示例中train_loss从 3.74e-04 变为 8.89e-03文档示例数值线性化程度与拟合精度是此策略要权衡的一对指标。该示例任务只有 1 个输入变量n_var1技巧本身与变量数无关但上述数值结论不能直接外推到多变量任务。完整可执行路径见 Example_11_encouraing_linear.ipynb 及其渲染版 Example_11_encouraing_linear.rstbase_fun的分支实现位于 MultKAN.pylamb_coef参与的正则计算位于 reg()。若确认深层激活已线性化下一步可结合官方示例中提到的剪枝策略把多余的深度剪掉。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价