用 Hessian 特征值分析 KAN 与 MLP 的损失景观与有效参数量(pykan 可解释性实战)
2026/9/14 4:33:58 网站建设 项目流程

用 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_var=1, train_num=1000, device=device) inputs = dataset['train_input'] labels = dataset['train_label']

要点说明:

  • create_dataset定义在 kan/utils.py 中,其默认签名为create_dataset(f, n_var=2, f_mode='col', ranges=[-1,1], train_num=1000, test_num=1000, normalize_input=False, normalize_label=False, device='cpu', seed=0)。这里显式设置n_var=1(输入维度为 1)、train_num=1000(训练样本 1000 个),并将数据放到当前设备(CUDA 优先,回退 CPU)。
  • 返回的dataset是一个字典,包含train_input/train_label/test_input/test_label四个键,本文只使用训练集部分。

训练一个 KAN 模型

实验中先训练一个宽度为[1,5,1]的 KAN(1 个输入、5 个隐藏单元、1 个输出):

# model = MLP(width = [1,30,1]) model = KAN(width=[1,5,1], device=device) model.fit(dataset, opt='Adam', lr=1e-2, lamb=0.000, steps=1000);

运行输出的关键信息(与 Notebook 记录一致):

cuda checkpoint directory created: ./model saving model version 0.0 | train_loss: 8.51e-04 | test_loss: 8.26e-04 | reg: 1.11e+01 | : 100%|█| 1000/1000 [00:08<00:00, 114] saving model version 0.1

对上述调用做几点展开:

  • KAN(width=[1,5,1], ...)KAN类在 kan/MultKAN.py 中定义,构造参数包括grid=3(B 样条网格数)、k=3(样条阶数)、noise_scale=0.3base_fun='silu'auto_save=True(默认自动保存检查点)等。默认检查点路径ckpt_path='./model',因此训练时会创建./model目录并写入0.00.1等版本文件(仓库根目录下的 model/ 即为该实验留下的检查点产物)。
  • model.fit(dataset, opt='Adam', lr=1e-2, lamb=0.000, steps=1000)fit方法同样定义在 kan/MultKAN.py,默认优化器为opt="LBFGS",此处显式切换为 Adam;lamb=0.000表示关闭正则项(文档代码刻意置零,以保证 Hessian 反映的是纯预测损失);steps=1000为训练步数。
  • 文档中注释掉的MLP(width=[1,30,1])是留给读者做对照实验的:MLP 定义在 kan/MLP.py,其激活函数默认act='silu',网络由若干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, derivative='hessian') values, vectors = torch.linalg.eigh(hess)
  • get_derivative(model, inputs, labels, derivative='hessian')返回损失对模型全部参数的二阶导数矩阵,其形状为(1, P, P),其中P是模型可训练参数总数(model2param将所有参数展平为一维向量后拼接,见 kan/utils.py 中的model2paramget_derivative)。
  • torch.linalg.eigh(hess)是对称矩阵的特征分解:values是特征值(升序排列),vectors是对应特征向量。Hessian 天然对称,因此使用eigh而非eig

特征值分布的可视化沿用文档代码:

import matplotlib.pyplot as plt plt.plot(values.cpu().numpy()[0], marker='o'); 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)。其内部流程可拆解为四个环节:

  1. 建立参数名映射(get_mapping:遍历model.state_dict()的所有 key,用正则把KANLayer.0.spline_fun.0.0.weight这类命名转换为形如model1.KANLayer[0].spline_fun[0][0].weight的可执行表达式,从而能够按名称把展平向量写回模型内部张量。

  2. 参数展平与还原

    • model2param(model)(kan/utils.py)把model.parameters()逐参数reshape(-1)后拼接成一维向量p
    • param2statedict(p, keys, shapes)再按各参数的shape把一维向量切分还原成 state_dict,实现「向量 ⇄ 参数」的双向转换。
  3. 构造「参数 → 损失」的可微函数(param2loss_fun:该函数接收展平参数向量p,先把参数写回模型副本,再计算损失。get_derivative通过loss_mode参数支持三种损失口径(源码 kan/utils.py):

    • loss_mode='pred'(默认):仅预测损失 $\text{MSE} = \text{mean}((\text{model}(x) - y)^2)$;
    • loss_mode='reg':仅正则损失model1.get_reg(reg_metric=..., lamb_l1=..., lamb_entropy=...)
    • loss_mode='all':预测损失 +lamb加权的正则损失。

    注意一个关键实现细节:这里用model.copy()复制模型再通过differentiable_load_state_dict写入参数,从而保证「参数 → 损失」的映射全程可微,这是二阶导数计算的前提。

  4. 调用自动微分求二阶导get_derivative依据derivative参数分发:

    • derivative='hessian'时调用batch_hessian(fun, p)(kan/utils.py),其实现是先对参数向量求一阶 Jacobian(batch_jacobian(fun, p, create_graph=True)),再对 Jacobian 结果再求一次 Jacobian,本质是「Jacobian of Jacobian」——这也解释了为什么需要create_graph=True保留计算图;
    • derivative='jacobian'时则直接调用batch_jacobian(kan/utils.py),基于torch.autograd.functional.jacobian实现。

整个链路可概括为:展平参数 → 可微地写回模型 → 计算损失 → 双重自动微分 → 得到 P×P 的 Hessianmodel2paramget_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乃至更小的量级,因此实践中通常以对数坐标图观察「特征值的悬崖式跌落」来区分有效与冗余方向,而不是机械地统计严格非零的个数。
  • 本实验用lamb=0.000关闭了正则,意味着 Hessian 完全来自预测损失;若改用loss_mode='reg''all'并开启正则(lamb>0),特征谱的形状会随之变化——正则项会改变损失景观的曲率。get_derivative提供的lamb_l1lamb_entropyreg_metric参数正是为此类扩展实验预留的接口。
  • 实验中 MLP 的宽度取[1,30,1](注释代码),总参数规模与[1,5,1]的 KAN 处于可对比的量级,这保证了「有效参数量差异」这一结论不是由简单参数总数差异造成的。

进阶实验建议

在 docs/Interp/Interp_10_hessian.ipynb 基础上,读者可以沿着以下方向把分析做深:

  1. 直接替换模型类型:把KAN(width=[1,5,1], ...)换成MLP(width=[1,30,1], ...)(或MLP(width=[1,10,1], ...)),其余代码不变,即可复现「KAN 非零特征值更多」的对比;也可以把目标函数换成更复杂的多变量函数(如x[:,[0]]**2 + x[:,[1]]**3),观察有效参数量的变化。
  2. 探究正则的影响:在model.fit(...)中设置lamb为非零值,并配合get_derivative(..., loss_mode='all', lamb=..., lamb_l1=..., lamb_entropy=...)计算含正则的 Hessian,对比特征谱如何随正则强度移动。
  3. 观察训练过程中的特征谱演化:利用 KAN 的检查点机制(model.saveckpt/model.loadckpt,默认写入./model),在不同训练阶段载入模型并分别计算 Hessian,观察非零特征值数量随训练步数的变化。
  4. 扩展到 Jacobian:将derivative='hessian'改为derivative='jacobian',可得到损失对参数的一阶导数向量,用于梯度层面的诊断。

小结

本文围绕 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),仅供参考

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

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

立即咨询