FEATURED · 精选文章

KAN 特征归因实战:用 feature_score、attribute 与 prune_input 量化输入重要性并剪枝

发布时间 / 2026/9/14 12:16:14
来源 / 创域科博编辑部
栏目 / 资讯中心
KAN 特征归因实战:用 feature_score、attribute 与 prune_input 量化输入重要性并剪枝 KAN 特征归因实战用 feature_score、attribute 与 prune_input 量化输入重要性并剪枝【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan本文基于 pykanKolmogorov-Arnold Networks官方可解释性教程系统讲解如何在 KAN 模型中量化每个输入特征对输出的贡献feature attribution。你将掌握三个核心能力读取并理解model.feature_score输入级归因分数、用model.attribute(l, i)检查任意隐藏节点对各特征的依赖关系以及通过prune_input自动或手动剔除无关输入让高维模型变得可解释、可读。一、特征归因在 KAN 中的意义与实现机制在符号回归、科学发现等场景中我们不仅关心模型拟合精度更关心哪些输入变量真正驱动了输出。这一需求被称为特征归因feature attribution。在 KAN 中得益于其仿射-free 的分层结构加法节点与乘法节点交替归因分数可以通过沿计算图反向传播自然获得无需额外的代理模型。在 pykan 的MultKAN类中归因机制由三个 API 组成model.feature_score一个property本质是调用attribute()后返回node_scores[0]即第一层输入层每个特征的归因分数见 kan/MultKAN.py#L358-L367model.attribute(l, i, out_score, plot)核心归因引擎支持按层、按神经元查询并可绘制柱状图见 kan/MultKAN.py#L1913-L2036model.prune_input(threshold, active_inputs)依据输入归因分数剪掉无关输入返回剪枝后的新模型见 kan/MultKAN.py#L1818-L1883。从源码看attribute()的反向传播算法可归纳为三步见 kan/MultKAN.py#L1985-L2013初始化输出分数从查询层出发用单位矩阵或用户指定的out_score作为该层各节点的初始分数逐层回传按node → subnode → edge → node的顺序自后向前传播分数其中从子节点到边的分数通过torch.einsum(ij,ki,i-kij, ...)结合边激活尺度edge_actscale与子节点激活尺度subnode_actscale加1e-4防除零计算取均值对所有输出维度取平均torch.mean(l, dim0)得到每个节点的最终归因分数其中node_scores[0]即输入特征分数。二、低维示例从构造数据集到读取特征分数文档首先构造了一个含 4 个输入、但重要性差异悬殊的合成任务from kan import * from sympy import * device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) # 构造数据集x0 影响最大x3 完全不参与 f lambda x: x[:,0]**2 0.3*x[:,1] 0.1*x[:,2]**3 0.0*x[:,3] dataset create_dataset(f, n_var4, devicedevice) input_vars [r$x_str(i)$ for i in range(4)] model KAN(width[4,5,1], devicedevice) model.fit(dataset, steps40, lamb0.001);这里create_dataset是 pykan 内置的合成数据生成器见 kan/utils.py#L62-L101关键参数包括f生成标签的符号函数、n_var输入维数、ranges输入取值范围默认[-1,1]、train_num/test_num默认各 1000、device与seed。输出目标函数f(x) x0² 0.3·x1 0.1·x2³ 0·x3意味着x0 起主导作用x3 对输出零贡献——这正是检验归因算法是否准确的标准试金石。训练 40 步后文档中的实际输出train_loss ≈ 8.00e-03test_loss ≈ 8.47e-03调用model.plot()可以可视化模型结构三、输入级归因model.feature_score直接读取输入层归因分数model.feature_score文档中的实际输出为tensor([0.8916, 0.5155, 0.1079, 0.0040], devicecuda:0, grad_fnMeanBackward1)这份分数与目标函数高度吻合x0 分数最高0.89x10.52与 x20.11次之而x3 的分数仅 0.004趋近于零——准确地识别出 x3 是无关特征。值得注意的是归因分数并不等同于回归系数它反映的是节点在计算图中对输出的整体贡献强度因此 x0 的二次项贡献被合理放大。四、隐藏节点归因model.attribute(l, i)除输入级分数外我们还可以检查某个隐藏神经元分别依赖哪些输入特征从而理解网络内部的分工# 第 1 层第 2 个神经元索引从 0 开始 model.attribute(1, 2)输出tensor([0.8915, 0.5146, 0.1079, 0.0040])与全局feature_score几乎一致说明该神经元综合使用了所有有效特征同时attribute会绘制一张柱状图再看第 1 层第 3 个神经元# 第 1 层第 3 个神经元索引从 0 开始 # 注意 y 轴尺度非常小 model.attribute(1, 3)输出tensor([4.6616e-05, 8.2072e-04, 3.2453e-06, 1.3511e-05])——所有值都在1e-3以下即该神经元对输出几乎没有贡献在训练中可能已被稀疏化正则压制。这正是 KAN 可解释性的体现通过逐神经元归因可以定位死神经元并判断哪些结构值得保留。从源码看attribute(l, i)的实现是当l ! None时先计算全模型归因再取出self.node_scores[l]中第i个神经元对应的分数向量见 kan/MultKAN.py#L1944-L1946并可选地绘制plt.bar柱状图见 kan/MultKAN.py#L2030-L2036。五、输入剪枝model.prune_input()既然知道 x3 不重要就可以直接从网络中剔除它model model.prune_input() model.plot(in_varsinput_vars)文档中的运行输出为keep: [True, True, True, False] saving model version 0.2prune_input的自动模式实现逻辑见 kan/MultKAN.py#L1857-L1863先执行attribute()拿到node_scores[0]再以input_score threshold生成布尔掩码threshold默认值为1e-2低于阈值的输入被判定为无关特征并打印keep:列表。除了自动模式它还支持手动模式传入active_inputs[0, 1]之类的索引列表即可按指定保留输入完全忽略归因分数见 kan/MultKAN.py#L1864-L1865。剪枝的底层操作并非直接修改原模型而是新建一个MultKAN并加载当前权重然后通过act_fun[0].get_subset(input_id, ...)只保留被选中输入对应的激活函数子集同时更新width[0]并记录input_id见 kan/MultKAN.py#L1867-L1877。剪枝后调用model.plot(in_varsinput_vars)图形中仅保留 x0、x1、x2 三个输入模型显著简化。六、高维案例100 维输入的归因、回退与剪枝当输入维度很高如 100 维而真正重要的特征很少时直接绘制全图既耗时又难以解读。文档给出了一个指数衰减权重的经典场景from kan import * # 构造数据集 n_var 100 def f(x): y 0 for i in range(n_var): # 指数衰减越靠后的特征越不重要 y x[:,[i]]**2*0.5**i return y dataset create_dataset(f, n_varn_var, devicedevice) input_vars [r$x_{str(i)}$ for i in range(n_var)] model KAN(width[n_var,10,10,1], seed2, devicedevice) model.fit(dataset, steps50, lamb1e-3);这里目标函数y Σ x_i² · 0.5^i的权重随索引 i 呈指数衰减0.5^i理论上只有前几个特征显著。训练完成后文档实际输出train_loss ≈ 3.20e-02test_loss ≈ 5.46e-02先回退到训练结束时的模型版本model model.rewind(0.1)rewind(model_id)将当前模型保存为新的一轮round再加载model_id指向的历史版本见 kan/MultKAN.py#L637-L665文档输出rewind to model version 0.1, renamed as 1.1。随后将 100 个特征的归因分数按排名绘制成 log-log 散点图plt.scatter(np.arange(n_var)1, model.feature_score.cpu().detach().numpy()) plt.xscale(log) plt.yscale(log) plt.xlabel(rank of input features, fontsize15) plt.ylabel(feature attribution score, fontsize15)图中可见明显的长尾分布排名前 10 的特征分数较高之后快速衰减到1e-2以下并趋于零——与构造数据时指数衰减的权重设计一致验证了feature_score在高维场景下的可靠性。最后先整体剪枝隐藏节点 边再按阈值剪掉无关输入model model.prune() model model.prune_input(threshold3e-2) model.plot(in_varsinput_vars)prune()默认参数为node_th1e-2, edge_th3e-2流程是先prune_node剔除归因分数低于阈值的节点再前向计算并重新归因最后prune_edge剔除低分边见 kan/MultKAN.py#L1782-L1816。紧接着的prune_input(threshold3e-2)输出keep: [True, True, True, True, True, False, ...]——100 个输入仅保留前 5 个有效特征第 64 个因分数恰好略高也短暂保留最终模型宽度从[100,10,10,1]收敛为仅 5 个输入的紧凑结构可视化后一目了然。七、实践要点与建议归因分数是相对强度而非绝对值feature_score与attribute的返回值反映节点/特征对输出的贡献强度适合用于排序与相对比较解读时应结合任务背景。阈值选取决定剪枝力度prune_input默认threshold1e-2高维场景可先用 log-log 散点图观察分数分布再针对性提高阈值如文档中的3e-2避免误删边缘有效特征。自动与手动模式互补自动模式适合快速筛选当领域知识明确知道某特征必须保留时用active_inputs[...]手动指定更稳妥。剪枝会返回新模型prune()与prune_input()返回的是剪枝后的新MultKAN实例记得用model model.prune_input()接收返回值配合rewind/checkout可以随时回退到历史版本。与稀疏化正则协同文档示例统一使用lamb0.001或1e-3的 L1 稀疏化正则它能让不重要的边/节点归因分数更快归零使后续剪枝更干净。完整的可复现代码见 docs/Interp/Interp_4_feature_attribution.rst 及同名 Jupyter Notebook docs/Interp/Interp_4_feature_attribution.ipynb相关实现源码位于 kan/MultKAN.py。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻