拓冰建站拓冰建站
首页 / 资讯中心 / 正文

用假设检验揭示 KAN 的内在对称性:pykan 的 `kan.hypothesis` 可分离性检测与树图绘制实战指南

用假设检验揭示 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, ...)自动检测加性/乘性可分离性并输出变量分组modeadd/mul、score_th1e-2、res_th1e-2、n_clustersNone、bias0.、verboseFalsetest_separability(model, x, groups, mode, ...)在给定分组下验证可分离性返回布尔值groups分组列表、modeadd、threshold1e-2、bias0test_general_separability(model, x, groups, ...)验证广义可分离性h(pq)groups、threshold1e-2test_symmetry(model, x, group, ...)验证某组变量是否只以标量组合影响输出group变量索引列表、dependence_th1e-3test_symmetry_var(model, x, input_vars, symmetry_var)用 SymPy 表达式显式假设对称变量并检验input_varssympy 符号、symmetry_varsympy 表达式plot_tree(model, x, style, ...)迭代应用上述检验并绘制变量组合树图styletree/box、sym_th1e-3、sep_th1e-1、skip_sep_testFalse、verboseFalse所有函数接受modelMultKAN、MLP 或任意 Python 函数与xtorch.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(pq)的形式允许内部子结构任意互换3. Case 1自动检测可分离性detect_separability3.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]]}输出解读hessian6×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) (x1x2) * (x3x4) * (x5x6)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其核心逻辑分三步计算 Hessianmodeadd时直接调用batch_hessian(model, x)计算输出关于输入的批量二阶导数矩阵modemul时则先对模型输出做log|f(x)bias|复合变换再求 Hessian源码第 59-60 行的compose(torch.log, torch.abs, lambda x: xbias, model)从而把乘法结构转化为对数域中的加法归一化打分用输入各维的标准差对 Hessian 做缩放hessian * std[:,None] * std[None,:]再沿批量维取中位数median得到对称的分数矩阵score_mat并据此构建硬阈值掩码score_mat score_th层次聚类分组以掩码矩阵为距离使用sklearn.cluster.AgglomerativeClusteringmetricprecomputed、linkagecomplete在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_separabilityf 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_var4.1 对称性的定义原文档对对称性的定义是输出y只依赖某几个变量的标量函数而不依赖这些变量的个体取值。形式化地说称函数具有对称性h(x1, x2)如果f(x1, x2, x3, ...) g(h(x1, x2), x3, ...)例如f (x1x2)·(x3x4)·(x5x6)对{x1,x2}具有对称性只依赖x1x2但{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, styletree) # 默认 style tree换用stylebox后每个中间节点被绘制为带属性标签的矩形框plot_tree(f, x, stylebox)可以看到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, styletree) # 默认 style tree5.3 树图的生成原理树图并非简单的绘图其背后是两阶段分析源码 kan/hypothesis.py阶段一get_molecule分子组装。从每个变量作为独立原子开始反复扫描当前原子列表尝试用test_symmetry(model, x, current_moleculeatom, dependence_thsym_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.pystyletreeAdd/Mul节点用蓝色斜线汇聚并标注红色或*GS节点用蓝色斜线但不标注符号Id节点绘制竖直黑线未知属性绘制矩形框stylebox所有非叶子节点统一绘制矩形框内直接写属性文字Add、GS、Id等。plot_tree的其他可选参数in_var输入变量名列表或 sympy 符号列表默认自动生成x_1, x_2, ...、sym_th1e-3对称性阈值、sep_th1e-1可分离性阈值注意这里默认比detect_separability的1e-2宽松、skip_sep_testFalse设为True可跳过属性测试以节省时间此时除Id外所有节点属性为空、verboseFalse。6. 从假设检验到模型结构推断的完整工作流综合上述三个 Case可以总结出一条从训练好的模型到结构理解的通用流程可分离性粗筛用detect_separability(model, x, modeadd)与modemul自动发现变量分组数n_groups与分组groups了解函数的大致分解形态对称性精查对粗筛出的候选分组或领域知识给出的候选组合用test_symmetry快速验证或用test_symmetry_var验证具体 SymPy 表达式如b/c、bc结构树还原调用plot_tree(model, x, stylebox)一键还原完整的变量组合层级结合节点属性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),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门