用 Hessian 特征值分析 KAN 与 MLP 的损失景观与有效参数量pykan 可解释性实战【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan导读本篇文章基于 pykan 官方可解释性教程 docs/Interp/Interp_10_hessian.rst对应 Notebook 见 docs/Interp/Interp_10_hessian.ipynb教程副本位于 tutorials/Interp/Interp_10_hessian.ipynb讲解如何通过计算「损失函数对模型参数的 Hessian 矩阵」并提取其特征值来理解损失景观loss landscape的结构进而对比 KAN 与 MLP 两类网络的有效参数量差异。读完本文你将掌握get_derivative这一核心 API 的完整用法、Hessian 特征值的可视化分析方法以及「KAN 的非零特征值通常多于 MLP」这一结论背后的源码级实现原理。背景为什么要计算 Hessian 特征值在深度学习可解释性研究中损失景观是理解模型优化行为的重要视角。损失函数在参数空间中的局部几何性质由Hessian 矩阵刻画——它是损失对模型参数的二阶偏导数矩阵。对 Hessian 做特征分解后非零特征值的数量对应损失函数在该点真正发生变化的参数方向数目可以粗略理解为模型的有效参数量effective number of parameters大量接近零的特征值意味着参数空间中存在许多平坦方向模型在这些方向上冗余特征值的量级分布尤其是用对数坐标观察反映了各参数方向对损失影响的强弱差异。原文档给出的核心实验结论是在相同任务上分别训练 KAN 与 MLP通常会发现 KAN 的 Hessian 非零特征值更多这意味着 KAN 的有效参数量大于 MLP。这一观察为「KAN 用更紧凑的结构表达更复杂函数」提供了损失景观层面的证据。实验准备数据集与模型在 pykan 中本实验基于一个最简单的单变量目标函数f lambda x: x[:,[0]]**2即 $f(x) x_1^2$。完整的环境准备与数据生成代码如下from kan.utils import get_derivative import torch from kan.MLP import MLP from kan.MultKAN import KAN from kan.utils import create_dataset, model2param import copy device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) f lambda x: x[:,[0]]**2 dataset create_dataset(f, n_var1, train_num1000, devicedevice) inputs dataset[train_input] labels dataset[train_label]要点说明create_dataset定义在 kan/utils.py 中其默认签名为create_dataset(f, n_var2, f_modecol, ranges[-1,1], train_num1000, test_num1000, normalize_inputFalse, normalize_labelFalse, devicecpu, seed0)。这里显式设置n_var1输入维度为 1、train_num1000训练样本 1000 个并将数据放到当前设备CUDA 优先回退 CPU。返回的dataset是一个字典包含train_input/train_label/test_input/test_label四个键本文只使用训练集部分。训练一个 KAN 模型实验中先训练一个宽度为[1,5,1]的 KAN1 个输入、5 个隐藏单元、1 个输出# model MLP(width [1,30,1]) model KAN(width[1,5,1], devicedevice) model.fit(dataset, optAdam, lr1e-2, lamb0.000, steps1000);运行输出的关键信息与 Notebook 记录一致cuda checkpoint directory created: ./model saving model version 0.0 | train_loss: 8.51e-04 | test_loss: 8.26e-04 | reg: 1.11e01 | : 100%|█| 1000/1000 [00:0800:00, 114] saving model version 0.1对上述调用做几点展开KAN(width[1,5,1], ...)KAN类在 kan/MultKAN.py 中定义构造参数包括grid3B 样条网格数、k3样条阶数、noise_scale0.3、base_funsilu、auto_saveTrue默认自动保存检查点等。默认检查点路径ckpt_path./model因此训练时会创建./model目录并写入0.0、0.1等版本文件仓库根目录下的 model/ 即为该实验留下的检查点产物。model.fit(dataset, optAdam, lr1e-2, lamb0.000, steps1000)fit方法同样定义在 kan/MultKAN.py默认优化器为optLBFGS此处显式切换为 Adamlamb0.000表示关闭正则项文档代码刻意置零以保证 Hessian 反映的是纯预测损失steps1000为训练步数。文档中注释掉的MLP(width[1,30,1])是留给读者做对照实验的MLP 定义在 kan/MLP.py其激活函数默认actsilu网络由若干nn.Linear堆叠而成。把KAN(width[1,5,1], ...)换成MLP(width[1,30,1], ...)即可重复「KAN vs MLP」的对比实验。训练完成后可以可视化模型结构model.plot()该调用会输出 KAN 的结构图节点、边与激活函数形态如下图所示。教程中这一步的作用是确认模型已收敛到一个合理的函数表示为后续 Hessian 分析提供基线。计算 Hessian 并提取特征值核心分析代码只有两行hess get_derivative(model, inputs, labels, derivativehessian) values, vectors torch.linalg.eigh(hess)get_derivative(model, inputs, labels, derivativehessian)返回损失对模型全部参数的二阶导数矩阵其形状为(1, P, P)其中P是模型可训练参数总数model2param将所有参数展平为一维向量后拼接见 kan/utils.py 中的model2param与get_derivative。torch.linalg.eigh(hess)是对称矩阵的特征分解values是特征值升序排列vectors是对应特征向量。Hessian 天然对称因此使用eigh而非eig。特征值分布的可视化沿用文档代码import matplotlib.pyplot as plt plt.plot(values.cpu().numpy()[0], markero); plt.yscale(log)注意这里values.cpu().numpy()[0]由于get_derivative返回的是带 batch 维大小为 1的 Hessian因此取第[0]个样本plt.yscale(log)将对数刻度便于同时观察跨数量级的特征值。运行结果如下图所示横轴为特征值索引纵轴为对数刻度的特征值大小可见绝大多数特征值落在较小量级而少量特征值显著偏大。源码剖析get_derivative 是如何工作的要理解「非零特征值数量 有效参数量」这一论断需要深入 kan/utils.py 中get_derivative的实现kan/utils.py。其内部流程可拆解为四个环节建立参数名映射get_mapping遍历model.state_dict()的所有 key用正则把KANLayer.0.spline_fun.0.0.weight这类命名转换为形如model1.KANLayer[0].spline_fun[0][0].weight的可执行表达式从而能够按名称把展平向量写回模型内部张量。参数展平与还原model2param(model)kan/utils.py把model.parameters()逐参数reshape(-1)后拼接成一维向量pparam2statedict(p, keys, shapes)再按各参数的shape把一维向量切分还原成 state_dict实现「向量 ⇄ 参数」的双向转换。构造「参数 → 损失」的可微函数param2loss_fun该函数接收展平参数向量p先把参数写回模型副本再计算损失。get_derivative通过loss_mode参数支持三种损失口径源码 kan/utils.pyloss_modepred默认仅预测损失 $\text{MSE} \text{mean}((\text{model}(x) - y)^2)$loss_modereg仅正则损失model1.get_reg(reg_metric..., lamb_l1..., lamb_entropy...)loss_modeall预测损失 lamb加权的正则损失。注意一个关键实现细节这里用model.copy()复制模型再通过differentiable_load_state_dict写入参数从而保证「参数 → 损失」的映射全程可微这是二阶导数计算的前提。调用自动微分求二阶导get_derivative依据derivative参数分发derivativehessian时调用batch_hessian(fun, p)kan/utils.py其实现是先对参数向量求一阶 Jacobianbatch_jacobian(fun, p, create_graphTrue)再对 Jacobian 结果再求一次 Jacobian本质是「Jacobian of Jacobian」——这也解释了为什么需要create_graphTrue保留计算图derivativejacobian时则直接调用batch_jacobiankan/utils.py基于torch.autograd.functional.jacobian实现。整个链路可概括为展平参数 → 可微地写回模型 → 计算损失 → 双重自动微分 → 得到 P×P 的 Hessian。model2param与get_derivative同时被导入正是因为前者承担了「参数向量化」这一前置步骤。实验结果解读KAN vs MLP 的有效参数量文档给出的核心结论值得反复推敲Try both KAN and MLP, you will usually see that KANs have more non-zero eigenvalues than MLPs, meaning that KANs have more effective number of parameters than MLP.结合本实验可以这样理解训练结束后模型收敛到损失景观中的一个近似极小值点。在该点计算 Hessian特征值的大小反映了「沿对应特征向量方向移动参数时损失变化的剧烈程度」。特征值数量级接近 0 的方向对损失几乎没有影响对应冗余/无效的参数方向非零特征值在数值意义上显著大于 0的方向才是真正被数据约束住、贡献模型表达能力的参数方向。因此非零特征值数量 ≈ 有效参数量。文档的经验结论是 KAN 在相同甚至更小宽度下拥有更多有效参数这与其「每条边都是一条可学习的 B 样条函数、参数分布在样条系数与网格上」的结构特性一致。MLP 则相反其权重矩阵中存在较多的线性相关/冗余方向表现为更多近零特征值。需要补充的严谨性说明从源码与实验设定可以推断「非零」是数值意义上的判断取决于数值阈值由于浮点计算与torch.linalg.eigh的精度限制理论上为 0 的特征值实际会落在1e-7乃至更小的量级因此实践中通常以对数坐标图观察「特征值的悬崖式跌落」来区分有效与冗余方向而不是机械地统计严格非零的个数。本实验用lamb0.000关闭了正则意味着 Hessian 完全来自预测损失若改用loss_modereg或all并开启正则lamb0特征谱的形状会随之变化——正则项会改变损失景观的曲率。get_derivative提供的lamb_l1、lamb_entropy、reg_metric参数正是为此类扩展实验预留的接口。实验中 MLP 的宽度取[1,30,1]注释代码总参数规模与[1,5,1]的 KAN 处于可对比的量级这保证了「有效参数量差异」这一结论不是由简单参数总数差异造成的。进阶实验建议在 docs/Interp/Interp_10_hessian.ipynb 基础上读者可以沿着以下方向把分析做深直接替换模型类型把KAN(width[1,5,1], ...)换成MLP(width[1,30,1], ...)或MLP(width[1,10,1], ...)其余代码不变即可复现「KAN 非零特征值更多」的对比也可以把目标函数换成更复杂的多变量函数如x[:,[0]]**2 x[:,[1]]**3观察有效参数量的变化。探究正则的影响在model.fit(...)中设置lamb为非零值并配合get_derivative(..., loss_modeall, lamb..., lamb_l1..., lamb_entropy...)计算含正则的 Hessian对比特征谱如何随正则强度移动。观察训练过程中的特征谱演化利用 KAN 的检查点机制model.saveckpt/model.loadckpt默认写入./model在不同训练阶段载入模型并分别计算 Hessian观察非零特征值数量随训练步数的变化。扩展到 Jacobian将derivativehessian改为derivativejacobian可得到损失对参数的一阶导数向量用于梯度层面的诊断。小结本文围绕 docs/Interp/Interp_10_hessian.rst 展开完整复现了「训练 KAN → 计算损失 Hessian → 特征分解 → 对数坐标可视化」的完整流程并从 kan/utils.py 源码层面解释了get_derivative的实现机制参数展平、可微参数回写、双重自动微分与loss_mode/derivative等参数的实际含义。Hessian 特征谱是理解 KAN 有效参数量的有力工具它把「KAN 的表达能力来自何处」这一问题落到了损失景观的几何语言上为后续的可解释性分析如特征归因、剪枝、符号化提供了定量依据。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考