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

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

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_var2, devicedevice)会在[-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_var2, devicedevice) 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 输出、网格grid5、B 样条阶数k3、随机种子seed1采用 LBFGS 优化器训练 20 步正则强度lamb0.01# train the model model KAN(width[2,5,1], grid5, k3, seed1, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb0.01); model.plot()训练日志要点checkpoint directory created: ./model saving model version 0.0 | train_loss: 3.34e-02 | test_loss: 3.29e-02 | reg: 4.93e00 | : 100%|█| 20/20 [00:0500: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_actTrueKAN 构造函数默认开启来缓存激活值用于计算正则项源码 kan/MultKAN.py 明确提示若lamb 0而save_actFalse正则项会退化为 0相当于lamb被静默置零。实验 Aλ 0关闭正则# train the model model KAN(width[2,5,1], grid5, k3, seed1, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb0.00); model.plot()| train_loss: 5.51e-03 | test_loss: 6.14e-03 | reg: 1.52e01 | : 100%|█| 20/20 [00:0300:00, 5.84it对比基线lamb0.01训练损失从 3.34e-02 大幅降至 5.51e-03拟合精度显著提升但正则项reg从 4.93 飙升至 1.52e01。这正是无正则化的典型表现模型自由地使用所有边网络结构稠密、曲线复杂虽然误差低但可解释性差。实验 Bλ 1正则过强# train the model model KAN(width[2,5,1], grid5, k3, seed0, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb1.0); model.plot()| train_loss: 1.70e00 | test_loss: 1.73e00 | reg: 1.08e01 | : 100%|█| 20/20 [00:0400:00, 4.59it正则项权重放大 100 倍后训练损失恶化到 1.70e00比基线高约 50 倍因为优化器被迫优先压低正则项而非拟合数据。注意本例中种子换成了seed0这也是损失差异的来源之一——控制变量时种子必须保持一致这一点将在参数三中详述。λ 的取值建议从三个实验可以清晰看到一条权衡曲线λ0 时欠正则、结构过密λ0.01 时在拟合与稀疏之间取得平衡λ1 时过正则、欠拟合。实践中应从较小的 λ如 0.0010.01起步观察reg与test_loss的走向再逐步调整。参数二λ_ent —— 熵正则的相对强度熵正则的底层计算除了整体强度lambfit()还提供lamb_l1L1 惩罚、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, axisrow) 1) p_col vec / (sum(vec, axiscol) 1) entropy -(mean(sum(p_row * log2(p_row 1e-4), axisrow)) mean(sum(p_col * log2(p_col 1e-4), axiscol))) 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固定lamb0.01将熵惩罚关掉# train the model model KAN(width[2,5,1], grid5, k3, seed1, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb0.01, lamb_entropy0.0); model.plot()| train_loss: 4.20e-02 | test_loss: 4.50e-02 | reg: 2.57e00 | : 100%|█| 20/20 [00:0400: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], grid5, k3, seed1, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb0.01, lamb_entropy10.0); model.plot()| train_loss: 7.83e-02 | test_loss: 7.74e-02 | reg: 1.54e01 | : 100%|█| 20/20 [00:0500:00, 3.77it将熵权重放大到 10 后正则项猛增至 1.54e01损失也涨到 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固定初始化从而保证同一超参数、同一种子得到可复现的结果。实验 Eseed 42model KAN(width[2,5,1], grid3, k3, seed42, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb0.01); model.plot()| train_loss: 5.67e-02 | test_loss: 5.72e-02 | reg: 5.81e00 | : 100%|█| 20/20 [00:0400:00, 4.81it注意本例与基线有两处不同种子从 1 改为 42且网格从grid5改为grid3。训练损失 5.67e-02 高于基线3.34e-02正则 5.81e00 略高于基线4.93e00。在网格更粗grid3的前提下不同种子得到不同的初始化最终收敛到不同的局部最优——这说明在比较超参数时必须固定 seed 与其余配置否则无法把差异归因于目标参数。seed 的实践要点论文级实验请固定 seed如seed1并同步固定torch、numpy的全局随机状态网格精度grid、样条阶数k也属于会影响收敛的配置改变它们时应作为独立变量对待多 seed 平均如 seed ∈ {0,1,42} 各跑一遍取均值可以降低初始化偶然性对结论的干扰。训练循环中的其他可调项fit()的完整签名还包含一批与正则正交、但在实操中同样重要的参数了解它们有助于把超参数实验设计得更严谨参数默认值作用optLBFGS优化器可选LBFGS或AdamLBFGS 使用 strong_wolfe 线搜索见 kan/MultKAN.pysteps100训练步数log1日志输出频率lr1.0学习率lamb_coef/lamb_coefdiff0.0 / 0.0B 样条系数的幅度 / 平滑惩罚update_grid/grid_update_numTrue / 10训练中是否自适应更新网格及更新次数reg_metricedge_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.010.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),仅供参考
分享:

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

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