pykan 可解释性实战:用 model.tree() 检验训练后 KAN 网络学到的对称结构
2026/9/14 17:42:49 网站建设 项目流程

pykan 可解释性实战:用 model.tree() 检验训练后 KAN 网络学到的对称结构

【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan

导读

本文基于 pykan 仓库的docs/Interp/Interp_6_test_symmetry_NN.rst文档,讲解如何对一个训练完成的 KAN(Kolmogorov-Arnold Network)模型进行结构对称性检验:先画出目标函数自身的树图作为“标准答案”,再训练一个KAN(width=[4,5,5,1])拟合该函数,最后用model.tree(sym_th, sep_th)把模型内部的模块结构还原成树图,与真实结构对比,验证网络是否学到了变量之间的对称分组关系。读完本文,你将掌握kan.hypothesis.plot_treeKAN.tree的完整用法、两个关键阈值参数(对称阈值sym_th、可分离阈值sep_th)的含义,以及从源码层面理解这套“假设检验 → 树图还原”流水线的实现原理。


1. 背景:为什么要在训练后检验对称性

在 Interp_5_test_symmetry.rst 中,pykan 介绍了如何对符号函数直接做三类假设检验(加性可分离、乘性可分离、广义可分离)并画出树图,其思路部分受 AI Feynman 项目启发。然而真实场景中,我们手头往往只有一个已经训练好的黑盒模型(KAN 或 MLP),并不知道它内部究竟以怎样的模块结构组合输入。

Interp_6把这一思路推进了一步:对一个训练过的神经网络做同样的结构探测。此时模型不再是显式公式,而是一堆可求导的样条(spline)与符号函数,但只要我们能用自动微分算出它对输入的梯度/海森矩阵,前面定义的对称性、可分离性检验就依然适用——这正是本教程的核心价值所在。

2. 第一步:画出目标函数的“真实”树图作为参照

教程首先构造了一个 4 变量目标函数:

from kan import * from kan.hypothesis import plot_tree f = lambda x: (x[:,[0]]**2 + x[:,[1]]**2) ** 2 + (x[:,[2]]**2 + x[:,[3]]**2) ** 2 x = torch.rand(100,4) * 2 - 1 plot_tree(f, x)

该函数的结构一目了然:输入被分成(x1, x2)(x3, x4)两个对称分组,各自先平方求和、再整体平方,最后两组相加。plot_tree对采样点x上的函数值做逐层假设检验,输出如下树图(真实结构):

树图顶层是红色+,左右两支分别由蓝色连线汇聚两个变量,直观表达了“两组对称变量先各自组合、再相加”的模块化结构。这张图就是后续判断模型“学得对不对”的参照物。

3. 第二步:训练一个 KAN 拟合该函数

接着按教程在数据集上训练一个 4→5→5→1 的 KAN:

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(device) dataset = create_dataset(f, n_var=4, device=device) model = KAN(width=[4,5,5,1], seed=0, device=device) model.fit(dataset, steps=100)

关键要素说明:

  • create_dataset:在 kan/utils.py 中实现,默认生成 1000 个训练样本与 1000 个测试样本(train_num/test_num可调),输入范围默认[-1,1],并返回含train_input/train_label/test_input/test_label四个键的字典。
  • KAN(width=[4,5,5,1]):输入维度 4,两个宽度为 5 的隐藏层,输出 1 维;seed=0保证可复现。
  • model.fit(dataset, steps=100):默认使用 L-BFGS 优化器迭代 100 步,每步自动记录train_loss / test_loss / reg(正则项)并打印进度条。

参考训练输出(教程原始运行结果,硬件为 CUDA 环境):

cuda checkpoint directory created: ./model saving model version 0.0 | train_loss: 1.58e-03 | test_loss: 4.79e-03 | reg: 2.38e+01 | : 100%|█| 100/100 [00:20<00:00, 4.93 saving model version 0.1

100 步训练后train_loss降到约1.58e-03,说明模型已较精确地拟合了目标函数。日志中的saving model version 0.0 / 0.1来自 kan/MultKAN.py 的log_history:模型在./model目录自动保存检查点并写入history.txt,这也是model.fit训练闭环的一部分。

4. 第三步:用 model.tree() 还原模型内部结构

训练完成后,直接对模型调用树图方法:

model.tree(sym_th=1e-2, sep_th=5e-1)

对比两张树图可以发现,训练后的模型同样把输入分成了(x1,x2)(x3,x4)两个对称分组,即网络从数据中自行学到了与真实函数一致的变量分组/对称结构——这正是本教程要检验的核心结论。同时,模型树图顶层出现的操作符可能与真实函数略有出入,这正说明树图还原的是“模型当前实际表达的模块结构”,而非理想化的目标公式,可作为后续剪枝(pruning)或符号化(symbolization)的依据。

model.tree的实现位于 kan/MultKAN.py(MLP 亦有同构实现,见 kan/MLP.py),本质上是把自身缓存的数据self.cache_data交给kan.hypothesis.plot_tree

def tree(self, x=None, in_var=None, style='tree', sym_th=1e-3, sep_th=1e-1, skip_sep_test=False, verbose=False): if x == None: x = self.cache_data plot_tree(self, x, in_var=in_var, style=style, sym_th=sym_th, sep_th=sep_th, skip_sep_test=skip_sep_test, verbose=verbose)

5. 参数详解:sym_th 与 sep_th 在源码中的真实作用

tree的完整签名与plot_tree一致,核心参数作用如下:

参数默认值含义与作用
xNone(取cache_data输入采样点,2D 张量,形状(Batch, n_var);传None时使用模型缓存的数据
in_varNone输入变量的符号名列表,用于树图叶子标注(默认为x_1, x_2, ...
style'tree'树图风格:'tree'用连线 + 操作符符号;'box'用矩形框标注属性
sym_th1e-3对称性阈值:判断“一组变量是否构成一个对称分子(molecule)”时使用的依赖度阈值
sep_th1e-1可分离性阈值:判断某个模块是 Add / Mul / GS(广义可分离)时的阈值
skip_sep_testFalseTrue时跳过每个模块属性的精细测试以节省时间(教程中未开启)
verboseFalse是否打印逐层组装分子的详细过程

sym_th的底层逻辑(对应 kan/hypothesis.py 的test_symmetry→ get_dependence):

  1. 对候选变量组group,计算模型 Jacobian 中该组方向的归一化梯度,再对其余变量求导,得到“组内方向随组外变量变化”的依赖度矩阵;
  2. 用输入标准差做归一化,取中位数聚合,得到依赖度;
  3. 若组与其余变量的最大依赖度< sym_th,则判定该组是“对称的”(输出只依赖该组的某个标量函数,不依赖组内各变量的具体取值)。

sep_th的底层逻辑(对应 test_separability 与 test_general_separability):

  • 对给定的分组groups,用批式海森矩阵计算“组间耦合强度”得分矩阵score_mat
  • 加性可分离要求组间交叉块的最大得分< sep_th(乘性可分离则先对函数取log|f+bias|再检验);
  • 广义可分离则进一步检验两组之间“梯度比值函数”是否乘性可分离。

plot_tree的整体流水线(kan/hypothesis.py)分两步:

  1. get_molecule(kan/hypothesis.py):把单变量视为“原子”,循环尝试把原子/分子组装成更大的分子,每次组装都用test_symmetry(..., dependence_th=sym_th)判断是否成立,从而得到从原子到全变量集合的层级分组列表moleculess
  2. get_tree_node(kan/hypothesis.py):对每一层的每个分子,用test_general_separability/test_separability(阈值sep_th)判定其属性,最终标记为AddMulGSId,并据此绘制矩形/连线与操作符。

值得注意:模型训练时也会自动保存检查点与history.txt(kan/MultKAN.py),如果你希望保存树图分析结果,可在调用model.tree后自行plt.savefig(...),不会影响模型本身。

6. 小结与实践建议

  • 操作流程plot_tree(f, x)得到真实结构 →KAN(...).fit(dataset)训练 →model.tree(sym_th, sep_th)得到模型结构,逐层对比。
  • 调参经验sym_th控制“哪些变量会被合并成一组”,取值过小会导致分组过度细化、树过深;sep_th控制“模块属性(Add/Mul/GS)的判定”,教程示例中采用sym_th=1e-2, sep_th=5e-1,默认值分别为1e-31e-1,可按数据量级调整。
  • 适用范围plot_tree/treemodel参数支持 KAN、MLP 乃至任意可求导的 Python 函数(文档中即同时给出了真值函数与训练模型两种用法),因此该检验方法可以推广到用 pykan 训练出的任意模型。

如需深入阅读源码,建议依次查看 kan/hypothesis.py(核心检验算法)、kan/MultKAN.py(tree方法)与 kan/utils.py(数据集生成),并与姊妹篇 Interp_5_test_symmetry.rst(对符号函数的直接检验)对照学习,即可完整掌握 pykan 的结构可解释性工具链。

【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询