用假设检验揭示 KAN 的内在对称性:pykan 的kan.hypothesis可分离性检测与树图绘制实战指南
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
面对一个训练好的 KAN 模型,解析出其背后的精确符号公式(如f(x)=sin(x1·x2)+x3²)是最理想的结果,但往往过于困难。退而求其次,我们仍可以通过一系列假设检验,判断模型所表示的函数是否具备某种模块化结构——例如加法/乘法可分离性、变量之间的对称性,以及变量如何逐层组合成完整表达式的树状结构。本指南基于 pykan 仓库的 Interp_5_test_symmetry.rst 与 hypothesis.py 源码,系统讲解detect_separability、test_symmetry、test_symmetry_var、plot_tree等工具的使用方法与底层原理。读完本文,你将能够对任意黑盒函数(或训练好的 KAN/MLP 模型)自动检测可分离性、验证对称性假设,并绘制出变量的组合树图。
本文对应的可运行 Notebook 位于 docs/Interp/Interp_5_test_symmetry.ipynb,姊妹篇 Interp_6_test_symmetry_NN.rst 则将同样的方法应用于神经网络(NN)模型。
1. 背景:从"精确公式"到"结构假设"
可解释性的终极目标是恢复模型的符号公式,但这在大多数场景下难以实现。正如原文档开篇所述:"Figuring out the symbolic formula represented by a model is ideal but sometimes too challenging. In this case, we might be content with simply figuring out some modular structures or symmetries."
这类假设检验的思路部分受到AI Feynman项目的启发:与其直接猜测整个函数,不如先回答一系列更简单的问题:
- 函数是否可以写成若干子函数相加或相乘的形式?
- 输出是否只依赖某几个变量的标量组合(对称性),而不再依赖变量的个体取值?
- 变量之间以怎样的层级结构组合在一起?
pykan 将这一整套能力封装在 kan/hypothesis.py 中,其核心 API 包括:
| 函数 | 作用 | 核心参数 |
|---|---|---|
detect_separability(model, x, mode, ...) | 自动检测加性/乘性可分离性并输出变量分组 | mode='add'/'mul'、score_th=1e-2、res_th=1e-2、n_clusters=None、bias=0.、verbose=False |
test_separability(model, x, groups, mode, ...) | 在给定分组下验证可分离性(返回布尔值) | groups(分组列表)、mode='add'、threshold=1e-2、bias=0 |
test_general_separability(model, x, groups, ...) | 验证广义可分离性h(p+q) | groups、threshold=1e-2 |
test_symmetry(model, x, group, ...) | 验证某组变量是否只以标量组合影响输出 | group(变量索引列表)、dependence_th=1e-3 |
test_symmetry_var(model, x, input_vars, symmetry_var) | 用 SymPy 表达式显式假设对称变量并检验 | input_vars(sympy 符号)、symmetry_var(sympy 表达式) |
plot_tree(model, x, style, ...) | 迭代应用上述检验并绘制变量组合树图 | style='tree'/'box'、sym_th=1e-3、sep_th=1e-1、skip_sep_test=False、verbose=False |
所有函数接受model(MultKAN、MLP 或任意 Python 函数)与x(torch.float类型的 2D 输入张量)作为前两个参数,因此不仅适用于 KAN 模型,也适用于任何可微的"黑盒"。
2. 准备工作与可分离性的数学定义
首先导入工具模块并构造测试输入:
from kan.hypothesis import * import torch本文示例全部使用 Pythonlambda函数作为被测对象。需要说明的是,hypothesis.py中的接口对输入x的形状约定是(Batch, Length)(批量维在前、特征维在后),且所有函数都要求模型可微——因为底层依赖一阶与二阶自动微分(详见后文源码分析)。
Case 1 聚焦于可分离性(separability),原文档给出了三类定义:
- 加性可分离:
f(x1, x2, ...) = g1(x1,x2) + g2(x3) + g3(x4,x5,x6) + ... - 乘性可分离:
f(x1, x2, ...) = g1(x1,x2) * g2(x3) * g3(x4,x5,x6) * ... - 广义可分离:
f(x1, x2, x3, ...) = h(p(x1,x2) + q(x3,...))(注意:广义加性可分离 = 广义乘性可分离,因为h(p+q)的形式允许内部子结构任意互换)
3. Case 1:自动检测可分离性(detect_separability)
3.1 加性可分离检测
考虑函数f(x) = x1*x2 + x3*x4 + x5*x6,它显然是三个二元乘积子函数相加:
f = lambda x: x[:,[0]] * x[:,[1]] + x[:,[2]] * x[:,[3]] + x[:,[4]] * x[:,[5]] x = torch.rand(100,6) * 2 - 1 detect_separability(f, x, 'add')运行后输出:
add separability detected并返回如下字典:
{'hessian': tensor([[0.0000, 0.3147, 0.0000, 0.0000, 0.0000, 0.0000], [0.3147, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000], [0.0000, 0.0000, 0.0000, 0.3619, 0.0000, 0.0000], [0.0000, 0.0000, 0.3619, 0.0000, 0.0000, 0.0000], [0.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.3358], [0.0000, 0.0000, 0.0000, 0.0000, 0.3358, 0.0000]]), 'n_groups': 3, 'labels': [2, 2, 1, 1, 0, 0], 'groups': [[4, 5], [2, 3], [0, 1]]}输出解读:
hessian:6×6 的 Hessian 分数矩阵。注意非零元素全部集中在(0,1)、(2,3)、(4,5)这些交叉位置,而同一子组内的元素(如(0,0))为零——这正是"组间无交叉导数"的体现;n_groups = 3:自动发现了 3 个互不相干的变量组;labels = [2, 2, 1, 1, 0, 0]:每个变量所属分组的编号(从 0 开始);groups = [[4, 5], [2, 3], [0, 1]]:实际的分组结果,即x5,x6、x3,x4、x1,x2各成一组,与真实函数结构完全吻合。
3.2 乘性可分离检测
将加法换成乘法:f(x) = (x1+x2) * (x3+x4) * (x5+x6):
f = lambda x: (x[:,[0]] + x[:,[1]]) * (x[:,[2]] + x[:,[3]]) * (x[:,[4]] + x[:,[5]]) x = torch.rand(100,6) * 2 - 1 detect_separability(f, x, 'mul');输出:
mul separability detected3.3 源码原理:Hessian 矩阵、归一化与层次聚类
detect_separability的实现位于 kan/hypothesis.py,其核心逻辑分三步:
- 计算 Hessian:
mode='add'时直接调用batch_hessian(model, x)计算输出关于输入的批量二阶导数矩阵;mode='mul'时则先对模型输出做log|f(x)+bias|复合变换再求 Hessian(源码第 59-60 行的compose(torch.log, torch.abs, lambda x: x+bias, model)),从而把乘法结构转化为"对数域中的加法"; - 归一化打分:用输入各维的标准差对 Hessian 做缩放
hessian * std[:,None] * std[None,:],再沿批量维取中位数(median),得到对称的分数矩阵score_mat,并据此构建硬阈值掩码score_mat < score_th; - 层次聚类分组:以掩码矩阵为距离,使用
sklearn.cluster.AgglomerativeClustering(metric='precomputed'、linkage='complete')在n_cluster_try范围内尝试不同分组数,计算每个候选分组下的残差比例residual_ratio = (total_sum - block_sum) / total_sum(源码第 89-93 行);当residual_ratio < res_th时记录该分组。若最终找到的分组数n_groups > 1,打印'{mode} separability detected'。
batch_hessian本身定义在 kan/utils.py:它借助torch.autograd.functional.jacobian对"批量雅可比之和"再求一次雅可比,从而得到批量 Hessian,返回张量形状为(Batch, Length, Length)。
3.4 给定分组验证(test_separability)
有时候我们已经有具体的变量分组假设,只需验证它是否正确。此时使用test_separability:
f = lambda x: (x[:,[0]] + x[:,[1]]) * (x[:,[2]] + x[:,[3]]) * (x[:,[4]] + x[:,[5]]) x = torch.rand(100,6) * 2 - 1 groups = [[0,1],[2,3],[4,5]] test_separability(f, x, groups, 'mul')输出tensor(True),正确分组通过验证。而错误分组会被拒绝:
test_separability(f, x, [[0,1],[2,4],[3,5]], 'mul')输出tensor(False)。
test_separability(源码 kan/hypothesis.py)的计算方式与detect_separability相同(同样计算归一化 Hessian 分数矩阵),区别在于不再聚类,而是执行两类检查:
- 内部测试:对任意两个不同分组
groups[i]与groups[j],检查其交叉块的分数最大值torch.max(score_mat[groups[i]][:,groups[j]]) < threshold; - 外部测试:若存在不属于任何分组的变量(
nongroup_id),检查分组与外部变量之间的交叉分数同样低于阈值。
只有所有检查都通过,才返回True。
3.5 广义可分离性(test_general_separability)
如果变量组合的外层函数不是简单的加或乘,而是任意可逆函数h呢?例如:
f = lambda x: torch.sin((x[:,[0]] + x[:,[1]]) * (x[:,[2]] + x[:,[3]]) * (x[:,[4]] + x[:,[5]])) x = torch.rand(100,6) * 2 - 1 test_separability(f, x, [[0,1],[2,3],[4,5]], 'mul')输出tensor(False):外层sin破坏了对数域的乘性结构,所以普通的乘性可分离测试失败。但广义可分离测试能够识别出"组内先组合、组间再组合"的结构:
test_general_separability(f, x, [[0,1],[2,3],[4,5]])输出tensor(True)。
其实现(kan/hypothesis.py)非常巧妙:对任意两个组 A、B 及组内成员member_A、member_B,考察函数grad[member_B] / grad[member_A](两组梯度的比值)。如果该比值函数是乘性可分离的,则说明两组变量通过某个外部函数h组合——从而验证广义可分离性。这正是"广义加性可分离 = 广义乘性可分离"这一性质的直接利用。
4. Case 2:对称性测试(test_symmetry与test_symmetry_var)
4.1 对称性的定义
原文档对对称性的定义是:输出y只依赖某几个变量的标量函数,而不依赖这些变量的个体取值。形式化地说,称函数具有对称性h(x1, x2),如果:
f(x1, x2, x3, ...) = g(h(x1, x2), x3, ...)
例如f = (x1+x2)·(x3+x4)·(x5+x6)对{x1,x2}具有对称性(只依赖x1+x2),但{x1,x3}不具备对称性。
4.2 使用test_symmetry检验候选分组
test_symmetry(model, x, group)接受一个变量索引列表group,返回布尔张量:
f = lambda x: (x[:,[0]] + x[:,[1]]) * (x[:,[2]] + x[:,[3]]) * (x[:,[4]] + x[:,[5]]) x = torch.rand(100,6) * 2 - 1 print('[0,1]:', test_symmetry(f, x, [0,1])) print('[0,2]:', test_symmetry(f, x, [0,2])) print('[2,3]:', test_symmetry(f, x, [2,3]))输出:
[0,1]: tensor(True) [0,2]: tensor(False) [2,3]: tensor(True)原理(源码 kan/hypothesis.py):
- 把变量划分为组内
group_A与组外group_B; - 计算模型对
group_A的梯度,并按group_A的梯度范数归一化得到"单位梯度方向"input_grad_A / ||input_grad_A||; - 再求该归一化方向关于
group_B的雅可比(batch_grad_normgrad,源码第 111-126 行),得到依赖度矩阵dependence,乘以输入标准差归一化后取中位数; - 若依赖度的最大值
< dependence_th(默认1e-3),说明组内变量的梯度方向不随组外变量变化——即组内变量以固定标量组合方式影响输出,返回True;否则返回False。
直觉上:若f只依赖h(x1,x2),则∂f/∂x1与∂f/∂x2的方向(比值)恒定,不受x3...影响;若组外变量能改变这一方向,则对称性假设不成立。源码第 163-164 行还有一个边界处理:当group覆盖全部变量或为空时直接返回True(此时对称性定义退化为平凡情形)。
4.3 使用test_symmetry_var检验任意 SymPy 表达式
test_symmetry只能检验"存在某个对称组合",而test_symmetry_var允许你显式给出假设的对称变量表达式(用 SymPy 符号定义),并输出证据强度:
from sympy import * # 该函数只依赖 b/c,而不依赖 b、c 的个体取值 f = lambda x: x[:,[0]] * torch.sqrt(1 + (x[:,[1]]/x[:,[2]])**2) input_vars = a, b, c = symbols('a b c') symmetry_var = b/c x = torch.rand(100,3) * 2 - 1 test_symmetry_var(f, x, input_vars, symmetry_var);输出:
100.0% data have more than 0.9 cosine similarity suggesting symmetry而错误的假设b*c会被拒绝:
not_symmetry_var = b * c test_symmetry_var(f, x, input_vars, not_symmetry_var);输出:
20.0% data have more than 0.9 cosine similarity not suggesting symmetry原理(源码 kan/hypothesis.py):
- 用
batch_jacobian计算模型关于输入的梯度input_grad; - 用
sympy.utilities.lambdify把symmetry_var编译为 numpy 函数,再计算该"对称变量"关于输入的梯度sym_grad; - 只保留出现在
symmetry_var.free_symbols中的变量维度(idx),计算两组梯度的余弦相似度cossim = |Σ(g1·g2)| / (||g1||·||g2||); - 统计余弦相似度 > 0.9 的数据比例
ratio:若ratio > 0.9(即 90% 以上的样本支持),打印suggesting symmetry,否则打印not suggesting symmetry,并返回完整的余弦相似度向量供进一步分析。
这一检验的直觉是:如果f确实只通过h(x1,x2)=b/c依赖b,c,那么模型梯度∂f/∂b、∂f/∂c的方向应与∂h/∂b、∂h/∂c的方向平行(沿相同等高线方向变化),从而余弦相似度接近 1。
5. Case 3:绘制变量组合树图(plot_tree)
将前述假设检验迭代应用,就能逐步还原变量如何自底向上组合成完整表达式,并以树图可视化。plot_tree(model, x, style=...)会依次调用get_molecule(逐层组装变量分子)与get_tree_node(标注每个节点的运算属性),最后用 matplotlib 绘制。
5.1 嵌套平方和结构(8 个变量)
考虑一个 4 层嵌套的平方和结构:
f = lambda x: ((x[:,[0]]**2 + x[:,[1]]**2) ** 2 + (x[:,[2]]**2 + x[:,[3]]**2) ** 2) ** 2 + ((x[:,[4]]**2 + x[:,[5]]**2) ** 2 + (x[:,[6]]**2 + x[:,[7]]**2) ** 2) ** 2 x = torch.rand(100,8) * 2 - 1 plot_tree(f, x, style='tree') # 默认 style = 'tree'换用style='box'后,每个中间节点被绘制为带属性标签的矩形框:
plot_tree(f, x, style='box')可以看到,box风格比tree风格多出每个节点的属性文字,便于直接读出每个组合模块的类型。
5.2 非对称结构(5 个变量)
把第二个分支替换为单个变量x5²,形成不对称结构:
f = lambda x: ((x[:,[0]]**2 + x[:,[1]]**2) ** 2 + (x[:,[2]]**2 + x[:,[3]]**2) ** 2) ** 2 + x[:,[4]]**2 x = torch.rand(100,5) * 2 - 1 plot_tree(f, x, style='tree') # 默认 style = 'tree'5.3 树图的生成原理
树图并非简单的绘图,其背后是两阶段分析(源码 kan/hypothesis.py):
阶段一:get_molecule(分子组装)。从每个变量作为独立"原子"开始,反复扫描当前原子列表,尝试用test_symmetry(model, x, current_molecule+atom, dependence_th=sym_th)判断某个原子能否并入当前分子(即并入后整体仍满足对称性假设)。能并入则合并,不能则开启新分子。每一轮扫描结束后,把当前分子作为下一轮的"原子",直到只剩一个分子。结果moleculess是一个分层列表,例如对 8 变量嵌套平方和会得到:
[[[0],[1],[2],[3],[4],[5],[6],[7]], [[0,1],[2,3],[4,5],[6,7]], [[0,1,2,3],[4,5,6,7]], [[0,1,2,3,4,5,6,7]]]阶段二:get_tree_node(属性标注)。对相邻两层的分子关系计算每个节点的"元数(arity)"并判定属性:
'Id':元数为 1(直连,无组合);'GS':元数 > 1 且通过test_general_separability——广义可分离节点(可被外部函数h包住);'Add'/'Mul':仅在最后一层(l == depth-1)进一步用test_separability区分加性/乘性;'':以上皆非(未知组合,绘制为空白矩形)。
对应到plot_tree的绘制逻辑(kan/hypothesis.py):
style='tree':Add/Mul节点用蓝色斜线汇聚并标注红色+或*;GS节点用蓝色斜线但不标注符号;Id节点绘制竖直黑线;未知属性''绘制矩形框;style='box':所有非叶子节点统一绘制矩形,框内直接写属性文字(Add、GS、Id等)。
plot_tree的其他可选参数:in_var(输入变量名列表或 sympy 符号列表,默认自动生成x_1, x_2, ...)、sym_th=1e-3(对称性阈值)、sep_th=1e-1(可分离性阈值,注意这里默认比detect_separability的1e-2宽松)、skip_sep_test=False(设为True可跳过属性测试以节省时间,此时除Id外所有节点属性为空)、verbose=False。
6. 从假设检验到模型结构推断的完整工作流
综合上述三个 Case,可以总结出一条从"训练好的模型"到"结构理解"的通用流程:
- 可分离性粗筛:用
detect_separability(model, x, mode='add')与mode='mul'自动发现变量分组数n_groups与分组groups,了解函数的大致分解形态; - 对称性精查:对粗筛出的候选分组(或领域知识给出的候选组合),用
test_symmetry快速验证,或用test_symmetry_var验证具体 SymPy 表达式(如b/c、b+c); - 结构树还原:调用
plot_tree(model, x, style='box')一键还原完整的变量组合层级,结合节点属性(Add/Mul/GS/Id)读出每个组合层的运算类型; - 验证与修正:若树图结构与预期不符,可调整阈值参数(
sym_th、sep_th)重新分析,或回到第 2 步补充验证更精细的对称性假设。
值得强调的是,这套工具并不绑定 KAN——model参数可以是 MultKAN、MLP 或任意可微 Python 函数,这也正是姊妹篇 Interp_6_test_symmetry_NN.rst 将同一套hypothesis工具应用于神经网络的原因。对于 KAN 用户而言,这意味着训练完成后无需人工审视激活函数图像,即可借助自动化的结构假设检验快速获得模型行为的高层理解,为后续的公式提取、剪枝与科学发现提供依据。
7. 小结
本文基于 kan/hypothesis.py 的源码实现,完整复现并深入讲解了 pykan 文档 Interp_5_test_symmetry.rst 中的三组核心工具:基于 Hessian 的加性/乘性/广义可分离性检测(detect_separability/test_separability/test_general_separability)、基于归一化梯度的对称性假设检验(test_symmetry/test_symmetry_var),以及由二者驱动的变量组合树图绘制(plot_tree)。底层依赖的批量雅可比/黑塞计算实现在 kan/utils.py,全部示例均可直接在 Interp_5_test_symmetry.ipynb 中运行验证。
掌握这些工具后,面对任何一个"黑盒"函数,你都能快速回答三个关键问题:它能不能分解?它对称吗?它的变量是怎样组合起来的?——这正是 AI Feynman 式科学发现工作流在 KAN 生态中的落地实践。
【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考