ARTICLE DETAIL

资讯详情

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

pykan 剪枝实战:用 prune_node / prune_edge / prune 让 KAN 模型更稀疏、更可解释

pykan 剪枝实战:用 prune_node / prune_edge / prune 让 KAN 模型更稀疏、更可解释 pykan 剪枝实战用 prune_node / prune_edge / prune 让 KAN 模型更稀疏、更可解释【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan本篇指南以 pykan 官方 API 教程 API_7_pruning.rst对应可运行 NotebookAPI_7_pruning.ipynb为核心系统讲解 Kolmogorov-Arnold NetworkKAN的剪枝能力KAN 提供自动剪枝与手动剪枝两条路径分别作用于节点node与边edge。读完本文你将掌握prune_node()、prune_edge()、prune()三个核心 API 的参数语义、调用方式以及它们背后的归因分数attribution score机制能够独立对训练好的 KAN 模型进行稀疏化与结构压缩。为什么需要剪枝稀疏、高效、可解释剪枝pruning是神经网络工程中的常用手段其目的是去除冗余参数使网络更稀疏从而带来两方面的收益一是推理与存储更高效二是结构更简洁、更容易被人类解读。KAN 也不例外——由于 KAN 使用可学习的样条激活函数spline activation替代固定激活函数其网络中每个节点和每条边都承载着实际的函数逼近信息剪掉低贡献的节点与边可以让谁在计算什么一目了然。pykan 的剪枝能力建立在**归因分数attribution score**之上。从源码看MultKAN即KAN的别名见 kan/MultKAN.py维护了三类分数node_scores节点归因、edge_scores边归因和subnode_scores子节点归因由attribute()方法通过反向传播计算得到kan/MultKAN.py。剪枝时分数低于阈值的组件即被视为死亡组件并被移除或置零。准备工作构建一个待剪枝的 KAN教程统一使用如下实验配置一个宽度为[2, 5, 1]的 KAN2 维输入、5 个隐藏神经元、1 维输出三次样条k3、5 个网格区间grid5并用合成函数f(x,y) exp(sin(πx) y²)生成数据集进行训练from kan import * device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) # create a KAN: 2D inputs, 1D output, and 5 hidden neurons. cubic spline (k3), 5 grid intervals (grid5). model KAN(width[2,5,1], grid5, k3, seed1, devicedevice) # create dataset f(x,y) exp(sin(pi*x)y^2) f lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) x[:,[1]]**2) dataset create_dataset(f, n_var2, devicedevice) dataset[train_input].shape, dataset[train_label].shape # train the model model.fit(dataset, optLBFGS, steps20, lamb0.01); model(dataset[train_input]) model.plot()训练日志与教程运行环境一致可参考 API_7_pruning.rst 中保存的输出cuda checkpoint directory created: ./model saving model version 0.0 | train_loss: 3.46e-02 | test_loss: 3.46e-02 | reg: 4.91e00 | : 100%|█| 20/20 [00:0500:00, 3.36it/s saving model version 0.1几点补充说明create_dataset的完整签名位于 kan/utils.py默认在[-1, 1]范围内各采样 1000 个训练与测试样本返回的字典包含train_input、train_label、test_input、test_label四个键训练使用 LBFGS 优化器并开启lamb0.01的正则项。lamb对应的正则约束L1、熵正则等会在训练中压低低贡献边的幅度为后续剪枝奠定基础由于auto_saveTrue是默认行为kan/MultKAN.py每次模型被修改训练、剪枝等都会在./model目录自动保存 checkpoint 版本saving model version 0.0 / 0.1 / 0.2。若想自定义存储目录可在构造KAN时传入ckpt_path参数。训练完成后model.plot()会画出网络结构。从教程截图可见训练后的网络仍保留了 5 个隐藏节点与大量连接边存在明显的剪枝空间。剪枝节点prune_node 的自动与手动模式节点剪枝对应model.prune_node()它有两种模式**自动auto**按归因分数阈值删除死亡节点**手动manual**则由用户直接指定保留哪些节点。教程代码如下mode auto if mode auto: # automatic model model.prune_node(threshold1e-2) # by default the threshold is 1e-2 model.plot() elif mode manual: # manual model model.prune_node(active_neurons_id[[0]])运行后输出saving model version 0.2剪枝后的网络结构如下参数语义与源码实现prune_node的完整定义位于 kan/MultKAN.py签名如下def prune_node(self, threshold1e-2, modeauto, active_neurons_idNone, log_historyTrue)threshold默认1e-2自动模式下的归因分数阈值。若某个节点的归因分数低于该阈值则视为死亡节点并被删除。调大阈值会剪得更狠调小则保留更多节点mode默认autoauto或manual。自动模式依赖self.attribute()计算出的node_scoreskan/MultKAN.py通过self.node_scores[i1] threshold生成保留掩码active_neurons_id默认None手动模式的保留清单。只要传入该参数mode会被强制置为manualkan/MultKAN.py随后按active_neurons_id[i]逐层标记需要保留的神经元kan/MultKAN.py。教程示例active_neurons_id[[0]]表示第 1 个隐藏层只保留 0 号神经元。从源码结构看剪枝过程还涉及两个关键步骤节点移除对每个非活跃节点调用remove_node(l, i, modeup/down)本质是把该节点所有入边与出边的 mask 置零kan/MultKAN.py模型重建剪枝不会原地修改原模型而是基于copy.deepcopy(self.width)构造一个新的MultKAN实例仅拷贝保留子集对应的权重与偏置node_bias、node_scale、act_fun、symbolic_fun等并更新width为新结构kan/MultKAN.py。因此必须像教程那样写model model.prune_node(...)接收返回值。需要注意prune_node要求模型先有一次前向传播记录激活值self.acts若为空会自动调用get_act()kan/MultKAN.py教程中在fit之后额外执行了一次model(dataset[train_input])正是为了填充激活缓存。剪枝边prune_edge 与阈值设定边剪枝对应model.prune_edge()它作用于层与层之间的连接即样条激活函数φ(l,i,j)。教程使用稍短的训练步数steps6重新训练一个同构模型后调用# create a KAN: 2D inputs, 1D output, and 5 hidden neurons. cubic spline (k3), 5 grid intervals (grid5). model KAN(width[2,5,1], grid5, k3, seed1, devicedevice) # create dataset f(x,y) exp(sin(pi*x)y^2) f lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) x[:,[1]]**2) dataset create_dataset(f, n_var2, devicedevice) dataset[train_input].shape, dataset[train_label].shape # train the model model.fit(dataset, optLBFGS, steps6, lamb0.01); model(dataset[train_input]) model.plot() model.prune_edge() model.plot()对应训练日志steps6时损失与正则值略高于 20 步版本| train_loss: 7.84e-02 | test_loss: 7.80e-02 | reg: 7.26e00 | : 100%|█| 6/6 [00:0100:00, 3.72it/s参数语义与源码实现prune_edge的完整定义位于 kan/MultKAN.pydef prune_edge(self, threshold3e-2, log_historyTrue)threshold默认3e-2边归因分数阈值。低于该阈值的边视为死亡并被置零。注意prune_edge的默认阈值3e-2与prune_node的默认阈值1e-2不同这是教程刻意区分的两个独立超参数实现要点对每一层用self.edge_scores[i] threshold生成新掩码并与旧的mask逐元素相乘后写回act_fun[i].maskkan/MultKAN.py。也就是说边剪枝是软删除——被剪的边只是 mask 置零网络宽度结构保持不变这与prune_node返回新模型的行为有本质区别因此调用时不需要写model model.prune_edge()直接model.prune_edge()即可。边剪枝后的网络如图所示剩余连接被明显稀疏化节点与边联合剪枝一个 prune 全搞定实际使用中最常见的是同时剪节点和边直接调用model.prune()即可。教程代码如下from kan import * # create a KAN: 2D inputs, 1D output, and 5 hidden neurons. cubic spline (k3), 5 grid intervals (grid5). model KAN(width[2,5,1], grid5, k3, seed1, devicedevice) # create dataset f(x,y) exp(sin(pi*x)y^2) f lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) x[:,[1]]**2) dataset create_dataset(f, n_var2, devicedevice) dataset[train_input].shape, dataset[train_label].shape # train the model model.fit(dataset, optLBFGS, steps20, lamb0.01); model(dataset[train_input]) model.plot() model model.prune() model.plot()prune的定义位于 kan/MultKAN.pydef prune(self, node_th1e-2, edge_th3e-2): if self.acts None: self.get_act() self self.prune_node(node_th, log_historyFalse) self.forward(self.cache_data) self.attribute() self.prune_edge(edge_th, log_historyFalse) self.log_history(prune) return self从源码可以看出prune()的内部调用链以node_th默认1e-2调用prune_node完成节点级剪枝并返回新模型对剪枝后的模型执行一次前向self.forward(self.cache_data)并重新计算归因分数self.attribute()因为节点被删除后边的贡献分布已经改变必须重算边分数以edge_th默认3e-2调用prune_edge完成边级剪枝记录prune到历史日志并返回最终模型。由于prune_node会重建模型prune()同样需要接收返回值model model.prune()。联合剪枝后的最终结构如下隐藏层被压缩为单节点边也大幅稀疏化剪枝的底层机制与实操建议归因分数从哪来剪枝决策依赖的node_scores与edge_scores由attribute()计算。该方法以输出层为单位矩阵作为初始节点分数沿网络逐层反向传播先由节点分数展开为子节点分数score_node2subnode再结合edge_actscale与subnode_actscale通过 einsum 求出边分数最后对边分数求和回传得到下一层节点分数kan/MultKAN.py。因此归因分数直观地反映了每个节点/边对输出的贡献量级这正是阈值剪枝的依据。几个实操要点阈值是核心旋钮prune_node默认1e-2、prune_edge默认3e-2教程中均使用默认值。实际项目中建议先用model.attribute()查看分数分布再据此调整阈值避免过度剪枝导致精度骤降剪枝前后对比推荐在剪枝后重新评估train_loss / test_loss并配合model.plot()目检结构。教程示例中 20 步训练后损失约3.46e-02剪枝后结构虽大幅缩小但仍可继续微调恢复精度自动保存机制由于默认auto_saveTrue每次剪枝都会在./model目录生成新版本 checkpoint教程日志中的saving model version 0.2即剪枝后的版本。若不需要自动保存可在构造模型时设置auto_saveFalse或在需要时通过state_id/ckpt_path回溯旧版本手动剪枝的自由度prune_node(active_neurons_id...)允许完全绕开阈值由用户指定每层保留哪些神经元适合在已通过可解释性分析确定关键神经元时使用此外源码还提供了prune_input按输入特征归因裁剪输入维度kan/MultKAN.py可作为进阶工具。小结pykan 的剪枝 API 为 KAN 模型提供了从稠密到稀疏的完整工具链prune_node()负责按归因分数或人工指定删除低贡献节点重建更窄的模型prune_edge()负责将低贡献边置零掩码软删除prune()则按先节点后边的顺序一键完成联合剪枝。三者的默认阈值分别为1e-2、3e-2与node_th1e-2 / edge_th3e-2全部基于attribute()计算的归因分数工作。掌握这套 API你就能在训练 KAN 之后快速获得一个结构精简、推理高效、且更容易通过 model.plot() 与符号化解释symbolic_formula进一步分析的稀疏网络。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表