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_input、train_label、test_input、test_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 > 0而save_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)起步,观察reg与test_loss的走向再逐步调整。
参数二:λ_ent —— 熵正则的相对强度
熵正则的底层计算
除了整体强度lamb,fit()还提供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)并同步固定torch、numpy的全局随机状态; - 网格精度(
grid)、样条阶数(k)也属于会影响收敛的配置,改变它们时应作为独立变量对待; - 多 seed 平均(如 seed ∈ {0,1,42} 各跑一遍取均值)可以降低初始化偶然性对结论的干扰。
训练循环中的其他可调项
fit()的完整签名还包含一批与正则正交、但在实操中同样重要的参数,了解它们有助于把超参数实验设计得更严谨:
| 参数 | 默认值 | 作用 |
|---|---|---|
opt | "LBFGS" | 优化器,可选"LBFGS"或"Adam";LBFGS 使用 strong_wolfe 线搜索(见 kan/MultKAN.py) |
steps | 100 | 训练步数 |
log | 1 | 日志输出频率 |
lr | 1.0 | 学习率 |
lamb_coef/lamb_coefdiff | 0.0 / 0.0 | B 样条系数的幅度 / 平滑惩罚 |
update_grid/grid_update_num | True / 10 | 训练中是否自适应更新网格及更新次数 |
reg_metric | 'edge_forward_spline_n' | 正则度量,可选edge_forward_spline_n、edge_forward_spline_u、edge_forward_sum、edge_backward、node_backward(见 kan/MultKAN.py) |
其中reg_metric决定正则项作用的"激活尺度来源":默认的edge_forward_spline_n使用前向样条激活尺度;node_backward则需要先执行node_attribute()归因计算(见 kan/MultKAN.py),适合以节点重要性为导向的稀疏化场景。
超参数调优路线图
综合本教程的三组对照实验,可总结出如下实操流程:
- 固定基线:固定 seed、grid、k、opt、steps,确定一个合理的
lamb起点(如 0.01); - 先调 λ:观察 train/test loss 与 reg 的权衡。损失低但 reg 异常高 → 增大 λ;损失高但 reg 低 → 减小 λ;
- 再调 λ_ent:固定 λ 后,在 0.0~10.0 范围内小步调整
lamb_entropy,以稀疏结构不损害测试精度为界; - 必要时开启系数正则:若激活曲线抖动剧烈,可引入
lamb_coefdiff增加平滑;若想强制样条归零,可引入lamb_coef; - 多 seed 验证:对选定的超参数组合,用多个 seed 复跑确认结论稳定;
- 用 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),仅供参考