资讯动态

如何用 PyKAN 对结理论数据集做监督分类并用 feature_score 提取特征重要性

发布时间:2026/9/15 19:05:54 来源:尧图企业网站定制
如何用 PyKAN 对结理论数据集做监督分类并用 feature_score 提取特征重要性【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykanPyKAN 仓库的 Example 14: Knot supervised 演示了一个完整的监督分类场景把结理论knot theory数据集按signature列分成若干类别用KAN模型训练一个分类器训练结束后通过model.feature_score拿到每个输入特征的归因分数画成柱状图判断哪些特征对分类贡献最大。本文按该示例的操作路径走一遍并在最后说明feature_score在 Interp 4: Feature attribution 中的定位和可用参数。准备环境与数据集依赖安装。按照 README 的说明PyKAN 要求 Python 3.9.7 或更高版本并可通过 pip 安装pip install pykan如果需要源码安装git clone https://github.com/KindXiaoming/pykan.git cd pykan pip install -e .README 列出的依赖版本参考requirements.txt 内容# python3.9.7 matplotlib3.6.2 numpy1.24.4 scikit_learn1.1.3 setuptools65.5.0 sympy1.11.1 torch2.2.2 tqdm4.66.2 pandas2.0.1 seaborn pyyaml也可以按 README 用pip install -r requirements.txt一次性装齐。数据集。示例脚本从本地路径./knot_data.csv读取数据代码中注明该数据来自 DeepMind 结理论 notebookmathematics_conjectures/blob/main/knot_theory.ipynb具体位置见 notebook 内注释。仓库本身不附带这个 CSV你需要先按注释里的来源把knot_data.csv放到工作目录。如果直接运行而没有该文件示例记录中出现的报错正是FileNotFoundError: [Errno 2] No such file or directory: ./knot_data.csv出现这个报错时先确认文件已就位而不是排查代码问题。预处理划分特征、标签并做归一化示例中的数据约定是knot_data.csv的第一列是结的名称后面用df.keys()[1:-1]跳过最后一列signature是目标。中间各列全部作为输入特征X。X 的归一化按列做标准化并把均值和标准差保存下来后续复用时可取用input_normalierX df[df.keys()[1:-1]].to_numpy() Y df[[signature]].to_numpy() # normalize X X_mean np.mean(X, axis0) X_std np.std(X, axis0) X (X - X_mean[np.newaxis,:])/X_std[np.newaxis,:] input_normalier [X_mean, X_std]Y 的类别化。signature是奇数取值示例用((Y - min_signature) / 2).astype(int)把它映射成0..n_class-1的整数标签类别数由n_class int((max_signature - min_signature)/2 1)算出# normalize Y max_signature np.max(Y) min_signature np.min(Y) Y ((Y-min_signature)/2).astype(int) n_class int((max_signature-min_signature)/21) output_normalier [min_signature, 2]训练/测试划分按 80/20 随机切分并用torch.manual_seed(42)、np.random.seed(42)固定种子dataset {} num X.shape[0] n_feature X.shape[1] train_ratio 0.8 train_id_ np.random.choice(num, int(num*train_ratio), replaceFalse) test_id_ np.array(list(set(range(num))-set(train_id_))) dtype torch.get_default_dtype() dataset[train_input] torch.from_numpy(X[train_id_]).type(dtype).to(device) dataset[train_label] torch.from_numpy(Y[train_id_][:,0]).type(torch.long).to(device) dataset[test_input] torch.from_numpy(X[test_id_]).type(dtype).to(device) dataset[test_label] torch.from_numpy(Y[test_id_][:,0]).type(torch.long).to(device)dataset必须是含train_input/train_label/test_input/test_label四个键的字典这是后面model.fit(dataset, ...)的直接输入。设备的选择沿用示例写法device torch.device(cuda if torch.cuda.is_available() else cpu)。训练 KAN 分类器分类任务用一个浅结构即可示例配置为输入n_feature维、单隐层1、输出n_class类的KAN(width[n_feature,1,n_class], grid5, k3, ...)配合交叉熵损失def train_acc(): return torch.mean((torch.argmax(model(dataset[train_input]), dim1) dataset[train_label]).float()) def test_acc(): return torch.mean((torch.argmax(model(dataset[test_input]), dim1) dataset[test_label]).float()) model KAN(width[n_feature,1,n_class], grid5, k3, seedseed, devicedevice) model.fit(dataset, lamb0.005, batch1024, loss_fnnn.CrossEntropyLoss(), metrics[train_acc, test_acc], display_metrics[train_loss, reg, train_acc, test_acc]);几个要点loss_fnnn.CrossEntropyLoss()是分类任务的损失model.fit会把标签按多类处理所以前面要把Y转成torch.long的类别索引。metrics传入train_acc/test_acc两个闭包display_metrics指定训练日志中显示train_loss、reg、train_acc、test_acc四项。lamb0.005是 L1 稀疏正则系数。按 README 的调参建议稀疏化有助于可解释性如果你只关心分类精度可以先从小值甚至lamb0起步再逐步加大。lamb的正则强度会影响后续feature_score的可读性Interp 4 示例中lamb0.001下分数已经能明显区分重要与不重要的特征示例 14 用稍大的lamb0.005。具体数值请以你自己的训练日志为准文档没有给出固定成功阈值。如果训练时报错找不到./knot_data.csv先回到上一节确认数据集文件存在。用 feature_score 提取特征重要性训练完成后直接读取model.feature_score它返回与输入特征顺序一一对应的归因分数源码 kan/MultKAN.py 中feature_score属性先调用attribute()再返回node_scores[0]即第一层对输入的归因scores model.feature_score features list(df.keys()[1:-1]) y_pos range(len(features)) plt.bar(y_pos, scores) plt.xticks(y_pos, features, rotation90); plt.ylabel(feature importance)示例 14 同时还画了模型结构图并给每个输入分支标注列名n 17对应该数据集中 17 个特征列如果你的列数不同需要相应调整model.plot(scale1.0, beta0.2) n 17 for i in range(n): plt.gcf().get_axes()[0].text(1/(2*n)i/n-0.005,-0.02,df.keys()[1:-1][i], rotation270, rotation_modeanchor)如何解读结果。feature_score是归因分数分数越高表示该输入对输出的贡献越大。Interp 4 文档给出了一个对照性验证方式用已知系数的函数x[:,0]**2 0.3*x[:,1] 0.1*x[:,2]**3 0.0*x[:,3]训练后model.feature_score的示例输出为tensor([0.8916, 0.5155, 0.1079, 0.0040])文档示例数值随训练环境略有差异——系数为 0 的x3得分接近 0与函数定义一致。在结理论场景里你可以同样检查柱状图中分数接近 0 的特征判断它们对signature的预测是否确实可以忽略。进阶检查中间节点与裁剪输入拿到整体分数后还可以按 Interp 4 的方法深入查看单个隐节点的依赖model.attribute(1, 2)返回第 0 层第 2 个隐节点对各个输入特征的归因分数并画出条形图文档示例中一个活跃节点的分数形如tensor([0.8915, 0.5146, 0.1079, 0.0040])文档示例而一个近乎无用的节点分数会全部落在1e-05量级示例输出为tensor([4.6616e-05, 8.2072e-04, 3.2453e-06, 1.3511e-05])文档示例。裁剪低分输入model model.prune_input()按默认阈值删掉分数低的输入Interp 4 文档示例的输出为keep: [True, True, True, False]文档示例高维场景可用model.prune_input(threshold3e-2)显式指定阈值再model.plot(in_varsinput_vars)画出裁剪后的图。文档还建议对高维网络先model.prune()再prune_input()否则整图难以阅读。恢复历史版本prune_input等操作会产生新的模型版本如文档中的saving model version 1.2可用model.rewind(0.1)回到训练完成的版本重新查看feature_score。裁剪会改变模型结构属于可选步骤如果只需要一份特征重要性清单读取model.feature_score即可无需裁剪。限制与下一步该任务的数据不在仓库内必须先按 notebook 注释的来源准备knot_data.csv否则第一步就会FileNotFoundError。示例 14 的 notebook 中除数据文件缺失报错外未记录训练完成的日志与最终精度因此文中所有数值结果均来自 Interp 4 文档示例只能作为量级参考不要当成固定预期。若后续任务要写自定义训练循环而不使用model.fit()README 提示调用model.speed()可关闭 symbolic 分支以提升效率属于可选优化与本场景无强依赖。同一数据集还有无监督变体 Example 15: Knot unsupervised需要时可在完成本文后再阅读。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

免费获取报价