pykan 训练超参数调优实战:λ、熵正则与随机种子如何塑造 KAN 的可解释性
2026/9/14 3:58:05 网站建设 项目流程

pykan 训练超参数调优实战:λ、熵正则与随机种子如何塑造 KAN 的可解释性

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

正则化(Regularization)是让 KAN 网络更稀疏、更可解释的关键手段,而正则化的效果高度依赖超参数的选择。本篇基于 pykan 仓库的 API 6 教程,通过控制变量实验逐一剖析lamb(整体惩罚强度)、lamb_entropy(熵正则相对强度)与seed(随机种子)对训练损失、正则项及最终网络结构的影响,并深入到 kan/MultKAN.py 的fit()reg()源码,说明每个超参数在底层是如何参与目标函数构造的。读完本文,你将掌握一套系统的 KAN 超参数排查思路,能够根据"损失偏高 / 结构过密 / 复现失败"等具体症状快速定位应调整的参数。

实验环境准备:构造目标函数与数据集

教程选用的目标函数是一个两变量复合函数:

f(x) = exp( sin(π·x₁) + x₂² )

在 kan/utils.py 中,create_dataset(f, n_var=2, device=device)会在[-1,1]范围内分别随机采样 1000 个训练样本和 1000 个测试样本,返回包含train_inputtrain_labeltest_inputtest_label四个键的字典。若目标函数输出是标量,会自动unsqueeze为单列形状。

from kan import * import torch device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(device) f = lambda x: torch.exp(torch.sin(torch.pi*x[:,[0]]) + x[:,[1]]**2) dataset = create_dataset(f, n_var=2, device=device) dataset['train_input'].shape, dataset['train_label'].shape

输出结果(默认 1000 训练样本、2 输入、1 输出):

cuda (torch.Size([1000, 2]), torch.Size([1000, 1]))

基线实验:默认超参数下的训练

先建立一条可对比的基线。模型使用宽度[2,5,1](2 输入、5 隐藏、1 输出)、网格grid=5、B 样条阶数k=3、随机种子seed=1,采用 LBFGS 优化器训练 20 步,正则强度lamb=0.01

# train the model model = KAN(width=[2,5,1], grid=5, k=3, seed=1, device=device) model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01); model.plot()

训练日志要点:

checkpoint directory created: ./model saving model version 0.0 | train_loss: 3.34e-02 | test_loss: 3.29e-02 | reg: 4.93e+00 | : 100%|█| 20/20 [00:05<00:00, 3.73it saving model version 0.1

从日志中可以看到三类关键指标:train_loss/test_loss(均方误差 RMSE)和reg(正则项数值)。训练过程会每隔log步输出一次,并在每个检查点自动保存模型版本。正则项reg的数值来自 kan/MultKAN.py 中self.get_reg(...)的计算结果,它直接参与了目标函数的构造。

参数一:λ —— 整体正则惩罚强度

λ 在目标函数中的位置

lamb(源码中写作lamb,对应fit()签名的第 4 个参数)控制的是"正则项在整个目标函数中所占的权重"。在 kan/MultKAN.py 的闭包函数中,优化目标被构造为:

objective = train_loss + lamb * reg_

即总损失 = 数据拟合损失 +lamb× 正则项。因此lamb越大,模型越倾向于牺牲拟合精度来换取稀疏、平滑的网络结构。一个关键细节是:当lamb > 0时,fit()内部必须依赖save_act=True(KAN 构造函数默认开启)来缓存激活值用于计算正则项;源码 kan/MultKAN.py 明确提示,若lamb > 0save_act=False,正则项会退化为 0,相当于lamb被静默置零。

实验 A:λ = 0(关闭正则)

# train the model model = KAN(width=[2,5,1], grid=5, k=3, seed=1, device=device) model.fit(dataset, opt="LBFGS", steps=20, lamb=0.00); model.plot()
| train_loss: 5.51e-03 | test_loss: 6.14e-03 | reg: 1.52e+01 | : 100%|█| 20/20 [00:03<00:00, 5.84it

对比基线(lamb=0.01),训练损失从 3.34e-02 大幅降至 5.51e-03,拟合精度显著提升,但正则项reg从 4.93 飙升至 1.52e+01。这正是"无正则化"的典型表现:模型自由地使用所有边,网络结构稠密、曲线复杂,虽然误差低但可解释性差。

实验 B:λ = 1(正则过强)

# train the model model = KAN(width=[2,5,1], grid=5, k=3, seed=0, device=device) model.fit(dataset, opt="LBFGS", steps=20, lamb=1.0); model.plot()
| train_loss: 1.70e+00 | test_loss: 1.73e+00 | reg: 1.08e+01 | : 100%|█| 20/20 [00:04<00:00, 4.59it

正则项权重放大 100 倍后,训练损失恶化到 1.70e+00(比基线高约 50 倍),因为优化器被迫优先压低正则项而非拟合数据。注意本例中种子换成了seed=0,这也是损失差异的来源之一——控制变量时种子必须保持一致,这一点将在参数三中详述。

λ 的取值建议

从三个实验可以清晰看到一条权衡曲线:λ=0 时欠正则、结构过密;λ=0.01 时在拟合与稀疏之间取得平衡;λ=1 时过正则、欠拟合。实践中应从较小的 λ(如 0.001~0.01)起步,观察regtest_loss的走向再逐步调整。

参数二:λ_ent —— 熵正则的相对强度

熵正则的底层计算

除了整体强度lambfit()还提供lamb_l1(L1 惩罚)、lamb_entropy(熵惩罚)、lamb_coef(系数幅度惩罚)、lamb_coefdiff(系数平滑惩罚)四个细分项,默认值分别为 1.0、2.0、0.0、0.0。它们的实际生效幅度是lamb与自身取值的乘积:如文档所述,熵正则的绝对强度为λ × λ_ent

在 kan/MultKAN.py 的reg()方法中,正则项对每一层激活尺度向量vec计算:

l1 = sum(vec) p_row = vec / (sum(vec, axis=row) + 1) p_col = vec / (sum(vec, axis=col) + 1) entropy = -(mean(sum(p_row * log2(p_row + 1e-4), axis=row)) + mean(sum(p_col * log2(p_col + 1e-4), axis=col))) reg += lamb_l1 * l1 + lamb_entropy * entropy

熵项使用 log2 信息熵衡量"激活在行/列方向上的分布是否集中":熵越低,说明激活越集中到少数边/节点上,网络越稀疏。这里的+1+1e-4都是数值稳定项,防止除零与 log(0)。此外,reg()还会对每个样条激活函数的 B 样条系数coef施加lamb_coef(系数 L1)与lamb_coefdiff(相邻系数差分 L1,鼓励样条曲线平滑)两类惩罚。

实验 C:λ_ent = 0(仅保留 L1)

固定lamb=0.01,将熵惩罚关掉:

# train the model model = KAN(width=[2,5,1], grid=5, k=3, seed=1, device=device) model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01, lamb_entropy=0.0); model.plot()
| train_loss: 4.20e-02 | test_loss: 4.50e-02 | reg: 2.57e+00 | : 100%|█| 20/20 [00:04<00:00, 4.68it

熵惩罚关闭后,正则项reg降到 2.57(基线 4.93 的一半左右),但损失略升到 4.20e-02。说明仅有 L1 惩罚时,模型不再被强制"集中激活",正则总量下降,稀疏性主要靠 L1 的软阈值效果维持。

实验 D:λ_ent = 10(熵惩罚过强)

# train the model model = KAN(width=[2,5,1], grid=5, k=3, seed=1, device=device) model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01, lamb_entropy=10.0); model.plot()
| train_loss: 7.83e-02 | test_loss: 7.74e-02 | reg: 1.54e+01 | : 100%|█| 20/20 [00:05<00:00, 3.77it

将熵权重放大到 10 后,正则项猛增至 1.54e+01,损失也涨到 7.83e-02。此时熵项主导了整个正则目标,优化器会把大量精力用于"压平激活分布",导致拟合精度受损。这组对照说明:lamb_entropy是一个灵敏度很高的旋钮,默认值 2.0 是相对均衡的选择,调参时应小步试探。

参数三:seed —— 随机种子与可复现性

seed 影响哪些随机性

seed同时控制 KAN 模型初始化(样条网格噪声、基函数初始化)与create_dataset的采样随机性。从 kan/utils.py 可见,数据集生成时会执行np.random.seed(seed)torch.manual_seed(seed);模型构造时也会用seed固定初始化,从而保证"同一超参数、同一种子"得到可复现的结果。

实验 E:seed = 42

model = KAN(width=[2,5,1], grid=3, k=3, seed=42, device=device) model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01); model.plot()
| train_loss: 5.67e-02 | test_loss: 5.72e-02 | reg: 5.81e+00 | : 100%|█| 20/20 [00:04<00:00, 4.81it

注意本例与基线有两处不同:种子从 1 改为 42,且网格从grid=5改为grid=3。训练损失 5.67e-02 高于基线(3.34e-02),正则 5.81e+00 略高于基线(4.93e+00)。在网格更粗(grid=3)的前提下,不同种子得到不同的初始化,最终收敛到不同的局部最优——这说明在比较超参数时,必须固定 seed 与其余配置,否则无法把差异归因于目标参数。

seed 的实践要点

  • 论文级实验请固定 seed(如seed=1)并同步固定torchnumpy的全局随机状态;
  • 网格精度(grid)、样条阶数(k)也属于会影响收敛的配置,改变它们时应作为独立变量对待;
  • 多 seed 平均(如 seed ∈ {0,1,42} 各跑一遍取均值)可以降低初始化偶然性对结论的干扰。

训练循环中的其他可调项

fit()的完整签名还包含一批与正则正交、但在实操中同样重要的参数,了解它们有助于把超参数实验设计得更严谨:

参数默认值作用
opt"LBFGS"优化器,可选"LBFGS""Adam";LBFGS 使用 strong_wolfe 线搜索(见 kan/MultKAN.py)
steps100训练步数
log1日志输出频率
lr1.0学习率
lamb_coef/lamb_coefdiff0.0 / 0.0B 样条系数的幅度 / 平滑惩罚
update_grid/grid_update_numTrue / 10训练中是否自适应更新网格及更新次数
reg_metric'edge_forward_spline_n'正则度量,可选edge_forward_spline_nedge_forward_spline_uedge_forward_sumedge_backwardnode_backward(见 kan/MultKAN.py)

其中reg_metric决定正则项作用的"激活尺度来源":默认的edge_forward_spline_n使用前向样条激活尺度;node_backward则需要先执行node_attribute()归因计算(见 kan/MultKAN.py),适合以节点重要性为导向的稀疏化场景。

超参数调优路线图

综合本教程的三组对照实验,可总结出如下实操流程:

  1. 固定基线:固定 seed、grid、k、opt、steps,确定一个合理的lamb起点(如 0.01);
  2. 先调 λ:观察 train/test loss 与 reg 的权衡。损失低但 reg 异常高 → 增大 λ;损失高但 reg 低 → 减小 λ;
  3. 再调 λ_ent:固定 λ 后,在 0.0~10.0 范围内小步调整lamb_entropy,以稀疏结构不损害测试精度为界;
  4. 必要时开启系数正则:若激活曲线抖动剧烈,可引入lamb_coefdiff增加平滑;若想强制样条归零,可引入lamb_coef
  5. 多 seed 验证:对选定的超参数组合,用多个 seed 复跑确认结论稳定;
  6. 用 plot() 目检:每一步都调用model.plot()观察网络结构与激活函数形状,正则效果最终要体现在"肉眼可见的稀疏"上。

本教程对应的完整可运行 Notebook 位于 docs/API_demo/API_6_training_hyperparameter.ipynb,训练过程中产生的检查点文件会写入仓库根目录的 model/ 目录(对应ckpt_path='./model'默认值)。关于检查点的保存、加载与版本管理机制,可进一步参考 API 12 检查点教程。

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

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

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

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

立即咨询