)
pykan 实战用线性基函数与样条系数惩罚鼓励 KAN 激活函数线性化Example 11 解析【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan导读当不确定一个任务需要多深的 KAN 时除了从最小模型逐步加宽加深更实用的策略是先建一个大模型再修剪掉冗余结构。本教程对应仓库 docs/Example/Example_11_encouraing_linear.rst 及同源 Notebook tutorials/Example/Example_11_encouraing_linear.ipynb演示了如何在宽度稀疏之外进一步让激活函数沿深度方向线性化——即让不需要的层退化为恒等映射shortcut。读完本文你将掌握两个可直接套用的技巧将基函数base_fun设为线性以及用lamb_coef惩罚样条系数并理解它们在 pykan 源码中的底层实现机制。背景为什么要鼓励线性在设计 KAN 时我们通常不知道到底需要多少层、多宽的模型。有两种常见策略自底向上从小模型开始逐步加宽、加深直到找到能较好完成任务的最小模型自顶向下先初始化一个足够大的模型再通过剪枝pruning等手段把它缩小。本示例走的是第二种路线。除了在宽度维度上追求稀疏让无用的节点/边被剪掉我们还想在深度维度上短路希望不必要的激活函数变成线性函数。这样深层网络中冗余的层会退化为恒等映射从而让模型结构在功能上等价于一个更浅的网络为后续的深度剪枝铺路。要让激活函数线性化文档给出了两个关键技巧把基函数base_fun设为线性惩罚样条系数——当样条系数全部为零时激活函数就只剩线性基函数项退化为线性函数。实验设定一个杀鸡用牛刀的任务为了直观展示效果示例选用了最简单的单变量函数$$f(x)\sin(\pi x)$$事实上这个函数用一个[1,1]单层、单输入单输出的 KAN 就足以精确拟合。但我们假设自己不知道这一点故意用一个[1,1,1,1]的三层 KAN 去拟合它——三层结构是冗余的这正是测试线性化技巧的理想场景。数据集的构造使用 pykan 提供的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):其中n_var指定输入变量个数ranges默认[-1,1]训练/测试样本数默认各 1000。不使用技巧默认训练下的非线性激活首先按默认配置训练模型基函数默认是 SiLUsilufrom 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);训练过程中的典型输出如下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/s] saving model version 0.1随后用默认参数绘制模型model.plot()可以看到在没有任何技巧干预的情况下虽然最终 train/test loss 已经降到 1e-4 量级、拟合精度足够高但每一层的激活函数都呈现出明显的非线性曲线上凸、下凸、类抛物线等形状冗余层并没有退化成线性映射。也就是说模型虽然精度达标但结构上仍然又深又密不利于后续剪枝。使用技巧线性基函数 样条系数惩罚接下来在同样的模型结构和数据上应用两个技巧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) # penalize spline coefficients model.fit(dataset, optLBFGS, steps20, lamb1e-4, lamb_coef10.0);与无技巧版本相比只改动了两个地方参数位置取值作用base_funidentityKAN(...)构造参数identity把残差基函数从默认的 SiLU 换成恒等映射lamb1e-4model.fit(...)1e-4开启正则化整体惩罚强度lamb_coef10.0model.fit(...)10.0惩罚样条系数的 L1 范数把样条系数推向零训练输出为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/s] saving model version 0.1绘制模型时把beta调大突出低重要度边的透明度对比model.plot(beta10)对比两张图可以直观看到使用技巧后第一层与第三层的激活函数变成了严格的直线对应恒等映射只有中间层保留了非线性的波动——这正是把正弦函数的非线性集中在某一层表达的理想结果。虽然 train/test loss约 8e-3比无技巧版本约 4e-4略高但换来的是结构上的极大简化为后续剪枝提供了干净的骨架。注意plot(beta...)中的beta控制边/激活的显示透明度由节点/边分数经score2alpha映射数值越大低分数部分越透明便于看清真正重要的激活函数详见 kan/MultKAN.py 中plot方法。源码级解析这两个技巧为什么有效要理解技巧为什么奏效需要看 pykan 底层KANLayer的前向计算。在 kan/KANLayer.py 中每个激活函数被建模为基函数项与样条项之和base self.base_fun(x) # 基函数项 b(x) y coef2curve(x_evalx, gridself.grid, coefself.coef, kself.k) # 样条项 spline(x) y self.scale_base[None,:,:] * base[:,:,None] self.scale_sp[None,:,:] * y即 $\phi(x) s_b \cdot b(x) s_{sp} \cdot \text{spline}(x)$。两个技巧分别作用于这两项技巧一base_funidentity让基函数项本身是线性的MultKAN.__init__对字符串形式的base_fun做了映射见 kan/MultKAN.pyself.base_fun_name base_fun if base_fun silu: base_fun torch.nn.SiLU() elif base_fun identity: base_fun torch.nn.Identity() elif base_fun zero: base_fun lambda x: x*0.默认值base_funsilu使用 SiLU 激活b(x)$ 本身是非线性的因此即使样条项被压为零激活函数依然是非线性的而base_funidentity 时 $b(x)x$一旦样条项消失激活函数就退化为线性甚至恒等映射。这正是线性基函数技巧的源码依据。技巧二lamb_coef把样条系数推向零fit()的训练目标是loss reg其中正则项由reg()方法计算见 kan/MultKAN.py。除默认的 L1/熵稀疏惩罚外正则项专门包含了对样条系数的惩罚# regularize coefficient to encourage spline to be zero for i in range(len(self.act_fun)): coeff_l1 torch.sum(torch.mean(torch.abs(self.act_fun[i].coef), dim1)) coeff_diff_l1 torch.sum(torch.mean(torch.abs(torch.diff(self.act_fun[i].coef)), dim1)) reg_ lamb_coef * coeff_l1 lamb_coefdiff * coeff_diff_l1代码注释写得很直白regularize coefficient to encourage spline to be zero。其中lamb_coef * coeff_l1对样条系数取绝对值求和强制系数整体趋近于零lamb_coefdiff * coeff_diff_l1对相邻系数差取绝对值求和强制系数在网格上平滑系数相等则曲线更平直。当样条系数全部为零时spline(x)项消失激活函数只剩线性基函数项——这正是样条系数归零 ⇒ 激活函数线性化的机制。fit()的完整签名kan/MultKAN.py中lamb整体惩罚强度、lamb_l1、lamb_entropy、lamb_coef、lamb_coefdiff的默认值分别为0.、1.、2.、0.、0.注意lamb默认关闭正则化、lamb_coef默认不惩罚系数——本示例显式开启它们才产生效果。训练输出里的reg是什么两个版本的训练日志中都出现了reg:字段无技巧 8.88e00有技巧 1.83e01。它就是在 kan/MultKAN.py 处调用self.get_reg(...)计算出的正则项数值。有技巧版本reg明显更高正是lamb_coef10.0对样条系数施加强惩罚的直接体现——惩罚力度越大正则项越大代价是训练 loss 略升收益是结构线性化。结合源码看参数取值范围与默认值把示例用到的所有参数与仓库源码中的默认值/取值对应起来参数示例取值源码默认值说明源码位置width[1,1,1,1]None各层神经元数含输入输出层kan/MultKAN.pygrid53网格区间数kan/MultKAN.pyk33B 样条阶数kan/MultKAN.pyseed01随机种子初始化时对 torch/np/random 统一设置kan/MultKAN.pynoise_scale0.10.3初始化注入到样条上的噪声幅度kan/MultKAN.pybase_funidentitysilu残差基函数支持silu/identity/zerokan/MultKAN.pylamb1e-40.整体正则化强度0表示关闭kan/MultKAN.pylamb_coef10.00.样条系数 L1 惩罚强度kan/MultKAN.pylamb_coefdiff未设0.样条系数平滑性惩罚强度kan/MultKAN.pyoptLBFGSLBFGS优化器可选LBFGS或Adamkan/MultKAN.pysteps20100训练步数kan/MultKAN.py实操提示与适用场景何时该用这两个技巧当你先建大模型再修剪时希望冗余层退化为恒等映射以便安全剪枝或者在可解释性分析中希望激活函数尽可能简单直线比复杂曲线更容易解释。两个技巧建议搭配使用只设base_funidentity而不用lamb_coef样条项仍可能保持非零非线性只用lamb_coef而基函数仍是 SiLU即使样条归零激活函数也非线性。二者配合才能让激活函数真正退化为线性。注意正则强度带来的 trade-off如示例所示lamb_coef10.0使 loss 从约 4e-4 升到约 9e-3属于用一点精度换结构简洁。实际使用时可按需调节lamb、lamb_coef、lamb_coefdiff的数值后两者默认都是0.在精度与稀疏/线性化之间取平衡。训练日志会触发自动保存示例运行会在当前目录创建./model检查点日志中的checkpoint directory created: ./model与saving model version 0.x可结合 API_12_checkpoint_save_load_model 中的saveckpt/loadckpt继续后续实验。小结本示例给出了在 pykan 中鼓励激活函数线性化的一套完整、可复现的做法把KAN构造参数base_fun设为identity线性基函数同时在fit()中开启lamb与lamb_coef样条系数惩罚。从源码看前者让激活函数的残差项本身线性后者通过reg()中的lamb_coef * |coef|把样条系数推向零二者叠加即可让冗余层退化为线性短路为从大模型出发的剪枝与可解释性分析铺平道路。当你不确定网络深度、又想避免盲目试错时这条先大后剪 鼓励线性的路线值得优先尝试。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考