资讯动态

pykan 正则化实战指南:用 reg_metric 与 lamb 让 KAN 更稀疏、更可解释

发布时间:2026/9/14 10:03:37 来源:尧图企业网站定制
pykan 正则化实战指南用 reg_metric 与 lamb 让 KAN 更稀疏、更可解释【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan本篇技术指南以 pykanKolmogorov-Arnold Networks的 API 8 正则化教程为主体系统讲解如何通过 L1/熵entropy正则化让 KAN 网络变稀疏、从而获得更强的可解释性。你将掌握 pykan 中五种reg_metric选择的具体含义与适用场景、fit()中lamb等超参数对训练的影响方式以及如何用model.plot()的不同metric直观检验稀疏化效果为后续的剪枝pruning与公式提取symbolic regression打下基础。一、为什么 KAN 需要正则化稀疏性是可解释性的前提KAN 将网络表示为可学习的样条spline激活函数其可解释性建立在结构足够简单之上如果网络里每个边、每个激活都处于活跃状态很难判断哪些输入真正驱动了输出。正则化Regularization的核心目标就是通过惩罚项迫使大部分边/节点的贡献趋近于零让网络自动长成一个稀疏图——只保留真正有作用的连接其余边在可视化中近乎透明从而帮助研究者解读模型学到的函数关系。pykan 官方文档指出Regularization helps interpretability by making KANs sparser. This may require some hyperparameter tuning.也就是说稀疏化效果与超参数尤其是lamb与正则化度量方式reg_metric强相关需要针对具体任务做调优这正是本篇要解决的核心问题。二、准备数据构造二输入回归数据集正则化实验的第一步与普通 KAN 训练一致导入kan包确定计算设备并借助create_dataset生成训练/测试数据。以下代码来自原教程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输出结果cuda (torch.Size([1000, 2]), torch.Size([1000, 1]))这里选择的目标函数为f(x) exp(sin(π·x₁) x₂²)自变量维度n_var2。create_dataset会按默认规模生成 1000 个训练样本与同等规模的测试样本输入形状为(1000, 2)、标签形状为(1000, 1)具体生成逻辑可查看 kan/utils.py 中的create_dataset实现。后续所有正则化实验都在这份数据集上进行。三、五种 reg_metric对哪个张量施加 L1 正则正则化的第一步是明确对什么施加惩罚。pykan 并不直接对网络权重做 L1而是对边的强度度量edge attribution / activation scale做惩罚。fit()的reg_metric参数提供了五种选择见 kan/MultKAN.py 中reg()方法的实现reg_metric含义底层张量说明edge_forward_spline_n默认边的范数归一化输出 std / 输入 std仅考虑 spline 部分忽略 symbolicacts_scale_spline训练时最常用只惩罚样条部分的归一化强度edge_forward_sum边的范数归一化输出 std / 输入 std同时包含 spline symbolicacts_scale把符号函数symbolic的贡献也计入边的强度edge_forward_spline_u边的范数未归一化输出 std仅考虑 spline 部分edge_actscale不除以输入 std量纲为原始输出标准差edge_backward边的归因分数edge attribution scoreedge_scores基于反向传播的归因得分与可视化中的 backward 一致node_backward节点的归因分数node attribution scorenode_attribute_scores对节点神经元层面做归因惩罚源码依据在 kan/MultKAN.py 中reg()方法根据reg_metric把acts_scale指向对应的强度张量if reg_metric edge_forward_spline_n: acts_scale self.acts_scale_spline elif reg_metric edge_forward_sum: acts_scale self.acts_scale elif reg_metric edge_forward_spline_u: acts_scale self.edge_actscale elif reg_metric edge_backward: acts_scale self.edge_scores elif reg_metric node_backward: acts_scale self.node_attribute_scores else: raise Exception(freg_metric {reg_metric} not recognized!)而这些张量是在前向传播forward()中缓存得到的kan/MultKAN.pyinput_range std(preacts) 0.1、output_range_spline std(postacts_numerical)仅样条部分、output_range std(postacts)含符号部分由此分别构造出acts_scale_spline归一化、仅 spline、acts_scale归一化、splinesymbolic与edge_actscale未归一化。理解这一链路有助于判断当你的模型大量使用符号函数如x²、sin时edge_forward_spline_n会看不见符号层的贡献此时改用edge_forward_sum更合理。需要说明的是edge_backward与node_backward两种模式在训练时还会额外触发归因计算fit()内部在reg_metric为这两者时分别调用self.attribute()与self.node_attribute()见 kan/MultKAN.py计算开销相对更高但能直接对解释性得分做惩罚语义上与最终可视化目标最一致。四、实战训练在 fit 中启用正则化原教程使用KAN(width[2,5,1], grid3, k3, seed1)初始化一个 2 输入、单隐层 5 神经元、单输出的 KAN并用 LBFGS 优化器训练 20 步# train the model model KAN(width[2,5,1], grid3, k3, seed1, devicedevice) model.fit(dataset, optLBFGS, steps20, lamb0.01, reg_metricedge_forward_spline_n); # default #model.fit(dataset, optLBFGS, steps20, lamb0.01, reg_metricedge_forward_sum); #model.fit(dataset, optLBFGS, steps20, lamb0.01, reg_metricedge_forward_spline_u); #model.fit(dataset, optLBFGS, steps20, lamb0.01, reg_metricedge_backward); #model.fit(dataset, optLBFGS, steps20, lamb0.01, reg_metricnode_backward); model.plot()训练过程输出含 checkpoint 与进度条信息checkpoint directory created: ./model saving model version 0.0 | train_loss: 4.57e-02 | test_loss: 4.35e-02 | reg: 7.15e00 | : 100%|█| 20/20 [00:0400:00, 4.58it saving model version 0.1关键超参数说明源自fit()的完整签名见 kan/MultKAN.pylamb正则化总强度最终优化目标为objective train_loss lamb * reg_见 kan/MultKAN.py。lamb0表示完全关闭正则化原教程取0.01作为温和的惩罚强度。reg_metric正则化度量方式取值即第三节表格中的五种默认edge_forward_spline_n。lamb_l1默认 1.0L1 惩罚强度作用于边强度向量的元素求和。lamb_entropy默认 2.0熵惩罚强度作用于边强度按行/列归一化后的熵用于鼓励均匀或集中的结构。lamb_coef默认 0.0样条系数幅度惩罚强度鼓励 spline 系数整体趋近于零。lamb_coefdiff默认 0.0相邻样条系数差值的 L1 惩罚平滑性鼓励系数曲线光滑。其余参数如optLBFGS或Adam、steps、lr、update_grid、grid_update_num等与常规训练一致不在正则化讨论范围内。训练机制细节从fit()源码可以看到正则化项的累计方式为对每一层边强度向量vec计算lamb_l1 * sum(vec) lamb_entropy * (entropy_row entropy_col)其中行/列熵基于p_row vec / (sum(vec, dim1) 1)、p_col vec / (sum(vec, dim0) 1)计算再叠加样条系数的lamb_coef与lamb_coefdiff惩罚kan/MultKAN.py。这意味着lamb并不是唯一的旋钮——即使lamb固定调节lamb_l1、lamb_entropy也会显著改变稀疏化的形态。五、解读训练指标train_loss / test_loss / reg训练进度条展示的三项指标分别来自results字典kan/MultKAN.pytrain_loss训练集上的 RMSEsqrt(mean((pred - label)²))本示例为4.57e-02说明 20 步内已能较好拟合目标函数test_loss测试集上的 RMSE本示例为4.35e-02与训练损失接近未出现明显过拟合reg正则化项reg_的数值本示例为7.15e00。fit()返回的results字典中这三项均为按步记录的一维数组可用来观察损失下降与稀疏化推进的权衡曲线若reg下降过慢说明惩罚不足图仍会显得稠密若train_loss明显劣于无正则化训练则说明lamb过大、过度压缩了模型容量需要下调。六、用 plot 可视化稀疏化效果三种 metricmodel.plot()在绘图时同样提供了与正则化度量对应的选项原教程说明 源码 kan/MultKAN.py 双重印证plot 的 metric对应张量语义forward_uedge_actscale同reg_metricedge_forward_spline_u未归一化输出 stdforward_nacts_scale归一化输出 std / 输入 std含 splinesymbolic对应edge_forward_sumbackward默认edge_scores同reg_metricedge_backward边的归因分数绘图时各边透明度由alpha tanh(beta * score)决定beta默认 3kan/MultKAN.py分数越低边越透明正则化做得越好可视化图中弱边越隐去网络结构越清晰。运行model.plot(metricforward_u) #model.plot(metricforward_n) #model.plot(metricbackward) # default以下是使用默认reg_metricedge_forward_spline_n、lamb0.01训练 20 步后绘制的 KAN 结构图默认 backward metric改用metricforward_u重新绘制同一模型边的不透明度改为基于未归一化输出 std可以交叉验证不同强度定义下的稀疏结构是否一致一个值得注意的细节原教程注释中写有 forward_n: same asedge_forward_spline_u但从 kan/MultKAN.py 的源码看forward_n实际读取的是self.acts_scale归一化、含 splinesymbolic与edge_forward_sum对应而forward_u才对应edge_forward_spline_u。对比源码后可以推断原注释此处存在笔误实际使用时应以源码映射为准。七、超参数调优建议与注意事项从默认组合起步reg_metricedge_forward_spline_nlamb0.01是一个稳妥的起点。若发现网络仍稠密可逐步增大lamb如 0.01 → 0.1 → 1观察reg项下降与train_loss上升的平衡点。符号函数参与时切换度量当模型中 symbolic 部分占比高时edge_forward_spline_n会忽略符号层贡献建议改用edge_forward_sum让惩罚覆盖全量贡献。区分归一化与未归一化edge_forward_spline_u未归一化对不同输入尺度的边一视同仁地按原始标准差惩罚可能偏向压制输出幅值大的边归一化版本则更能反映边的相对重要性。归因类度量的代价edge_backward/node_backward与最终可视化语义一致但训练中每步都要额外做归因计算attribute()/node_attribute()训练更慢适合对可解释性有高要求的小规模任务。lamb0时的行为fit()中当lamb0时会自动关闭激活缓存与符号层见disable_symbolic_in_fitkan/MultKAN.py此时正则项恒为 0如需启用正则化请保持lamb 0并确认模型处于可缓存激活的状态。与其他 API 串联稀疏化是剪枝pruning见 API 7与公式提取symbolic regression见 API 12 等的前置步骤——剪枝正是基于edge_scores等强度张量按阈值掩码实现的kan/MultKAN.py正则化先把弱边压下去剪枝再将其删干净二者配合可获得极简的可解释模型。八、小结正则化是 pykan 可解释性工作流中的关键一环通过reg_metric选定对哪类边强度做惩罚通过lamb及其细分项lamb_l1、lamb_entropy、lamb_coef、lamb_coefdiff控制惩罚力度训练后用plot()的forward_u/forward_n/backward三种视图检验稀疏化效果。本教程示例在 20 步 LBFGS 训练内将回归损失压到4.57e-02train/4.35e-02test同时让reg项参与优化为后续剪枝与公式提取铺平了道路。实践时建议结合训练输出的reg曲线与结构图透明度反复微调lamb找到拟合精度与结构稀疏性的最佳平衡。延伸阅读本教程对应的可运行 Notebook 位于 docs/API_demo/API_8_regularization.ipynbreg()、fit()、plot()、attribute()的完整实现参见 kan/MultKAN.py与正则化配套的剪枝、符号化等后续流程可参考 docs/API_demo 下的 API 7、API 12 教程。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价