ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

pykan 可视化实战:掌握 KAN 网络 `plot()` 绘图 API 的完整参数指南

pykan 可视化实战:掌握 KAN 网络 `plot()` 绘图 API 的完整参数指南 pykan 可视化实战掌握 KAN 网络plot()绘图 API 的完整参数指南【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan本指南以 pykan 开源仓库Kolmogorov Arnold Networks的官方 API 文档 docs/API_demo/API_2_plotting.rst 为主体系统讲解 KAN 模型可视化接口plot()的全部核心参数。你将掌握如何控制激活函数透明度beta、切换重要性度量metric、调整画布尺寸scale、叠加样本散点sample以及通过剪枝prune、符号化fix_symbolic与混合模式set_mode改变连线颜色编码从而真正读懂并输出高质量的网络结构图。1. 环境准备初始化 KAN 模型并生成数据集在调用绘图接口之前需要先完成两件事构造一个 KAN 模型和准备一份用于前向传播的数据集。官方示例从kan包导入全部符号from kan import *该包在 kan/init.py 中导出了MultKAN/KAN类与utils工具函数并自动选择 CUDA 或 CPU 设备from kan import * device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) # create a KAN: 2D inputs, 1D output, and 5 hidden neurons. # cubic spline (k3), 3 grid intervals (grid3). model KAN(width[2,5,1], grid3, k3, seed1, devicedevice) # create dataset f(x,y) exp(sin(pi*x)y^2) 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运行后终端会输出设备信息以及自动保存的检查点提示checkpoint directory created: ./model、saving model version 0.0数据形状为(torch.Size([1000, 2]), torch.Size([1000, 1]))。这里的create_dataset定义于 kan/utils.py其默认参数为ranges[-1,1]、train_num1000、test_num1000、seed0返回的字典中包含train_input、train_label、test_input、test_label四个键。KAN 类本身在 kan/MultKAN.py 中实现width[2,5,1]表示输入 2 维、隐藏层 5 个神经元、输出 1 维。注意plot()依赖模型保存前向传播的中间激活值save_actTrue默认开启。文档源码注释明确指出 cannot plot since data are not saved. Set save_actTrue first.若在未做前向传播前直接绘图会抛出异常model hasnt seen any data yet.。因此官方示例总是先执行一次model(dataset[train_input])再进行绘图。2. 绘制初始化状态下的 KANbeta与标签参数2.1 基础绘图model.plot(beta100)对刚初始化未训练的模型调用# plot KAN at initialization model(dataset[train_input]); model.plot(beta100)图中从左到右依次为输入层2 个节点、隐藏层5 个节点与输出层1 个节点每一对节点之间的小图即为对应的 1D 激活函数spline 样条曲线连线粗细/透明度反映激活函数的相对重要性。2.2 添加变量名与标题in_vars/out_vars/title# if you want to add variable names and title model.plot(beta100, in_vars[r$\alpha$, x], out_vars[y], title My KAN)in_vars与out_vars接收字符串列表支持 LaTeX 语法如r$\alpha$title为整图标题。在plot()的源码签名中这些参数默认均为None另有varscale用于控制输入变量文字的缩放。3. 训练带稀疏正则的模型fit()与lamb为了观察训练后的结构变化官方示例使用 L-BFGS 优化器训练 20 步并施加 L1 正则# train the model model.fit(dataset, optLBFGS, steps20, lamb0.01);训练输出| train_loss: 5.20e-02 | test_loss: 5.35e-02 | reg: 4.93e00 | : 100%|█| 20/20 [00:0300:00, 5.22it saving model version 0.1fit()的完整签名定义于 kan/MultKAN.pyoptLBFGS、steps100、lamb0.、update_gridTrue、grid_update_num10等。这里的关键是lamb正则系数正则会抑制不重要的激活使后续绘图中重要连接与噪声连接的对比更明显。4. 理解beta参数控制激活函数透明度训练后beta直接决定哪些激活函数在图中浮现出来。官方说明如下beta控制激活函数的透明度。更大的beta意味着更多激活函数显示出来。我们通常希望设置一个合适的beta使得只有重要的连接在视觉上显著。透明度公式为$$\text{transparency} \tanh(\beta \cdot \phi)$$其中 $\phi$ 的取值取决于metricmetricforward_u激活函数本身的尺度scalemetricforward_n归一化后的尺度normalized scalemetricbackward特征归因分数feature attribution score默认在 kan/MultKAN.py 中对应实现为score2alpha(score) np.tanh(beta * score)随后alpha [score2alpha(...) for score in scores]连线透明度即由该alpha与 mask 共同决定。4.1 默认参数绘图model.plot()默认情况下beta3、metricbackwardmodel.plot()4.2 大betabeta100000几乎全部激活可见model.plot(beta100000)4.3 小betabeta0.1只保留最重要连接model.plot(beta0.1)三张图对比可以直观体会到beta越大被激活显示的连接越多beta越小图中保留的越接近模型真正的骨架。5. 切换重要性度量metricforward_n | forward_u | backwardmetric决定连线的透明度基于哪种分数。官方示例在beta100下逐一演示三种度量model.plot(metricforward_n, beta100) model.plot(metricforward_u, beta100) model.plot(metricbackward, beta100)从源码kan/MultKAN.py可以看到三种度量对应的内部数据metric内部数据含义forward_nself.acts_scale归一化激活尺度forward_uself.edge_actscale未归一化的边激活尺度backwardself.edge_scores反向归因特征重要性分数默认三者对应的输出图分别为若传入不支持的度量源码会抛出Exception(fmetric \{metric}\ not recognized)。需要说明使用backward度量时plot()内部会先自动调用self.attribute()计算边分数见 kan/MultKAN.py。6. 剪枝与结构精简prune()之后重新绘图训练 稀疏正则后大量不重要的神经元与边已被抑制。调用prune()可以将它们真正移除得到一个更紧凑的网络model model.prune() model.plot()输出saving model version 0.2。剪枝后图结构从 2-5-1 变为 2-1-1隐藏层仅剩 1 个神经元prune()在源码中同时调用prune_node(node_th1e-2)与prune_edge(edge_th3e-2)见 kan/MultKAN.py默认阈值分别为node_th1e-2、edge_th3e-2model model.prune()的赋值写法说明剪枝会返回一个新模型对象。7. 调整画布尺寸scalescale控制整张图的物理尺寸默认值为 0.5model.plot(scale0.5) # 默认值500x400 model.plot(scale0.2) # 缩小200x160 model.plot(scale2.0) # 放大2000x1600在源码中画布尺寸由figsize(10 * scale, 10 * scale * (neuron_depth - 1) * (y0z0))决定kan/MultKAN.py同时节点大小min_spacing ** 2 * 10000 * scale ** 2与连线宽度lw2 * scale也会随scale等比缩放。8. 叠加样本散点sampleTrue与get_act()默认绘图只画激活函数曲线折线。当希望同时看到数据样本在每条激活函数上的分布时设置sampleTruemodel.plot(sampleTrue)样本越多散点越密集、越难分辨。官方示例通过get_act()只喂入前 20 个样本让散点更清晰model.get_act(dataset[train_input][:20]) model.plot(sampleTrue)源码中散点由plt.scatter(..., s400 * scale ** 2)绘制kan/MultKAN.py即散点大小同样受scale控制。get_act(x)用于以指定输入x重新计算并缓存各层激活值供后续绘图使用。9. 颜色编码语义符号函数、数值函数与混合模式KAN 支持将单个激活替换为已知符号函数如x^2、sin(x)等绘图会以颜色区分边的类型。颜色规则在 kan/MultKAN.py 中由symbolic_mask与numeric_mask共同决定类型symbolic_masknumeric_mask连线颜色纯符号函数00红色red纯数值样条00黑色black符号 数值混合00紫色purple被 mask 移除00白色white9.1 将激活替换为符号函数fix_symbolic(0,1,0,x^2)model.fix_symbolic(0,1,0,x^2)输出r2 is 0.9992202520370483 saving model version 0.3 tensor(0.9992, devicecuda:0)fix_symbolic(l, i, j, fun_name)将第l层中第i个输入到第j个输出的激活替换为指定符号函数返回拟合的 R² 分数此处为 0.9992说明x^2与该激活拟合得非常好。之后调用model.plot()该边会显示为红色9.2 符号 数值混合模式set_mode(0,1,0,modens)当一条边同时启用符号函数与样条函数输出为两者之和时显示为紫色model.set_mode(0,1,0,modens) model.plot(beta100)set_mode在源码kan/MultKAN.py中支持四种模式modes纯符号mask_s1mask_n0moden纯数值样条mask_s0mask_n1modesn或modens符号与数值混合两者 mask 均为 1其他取值全部关闭mask 均为 010. 参数速查表与绘图工作流建议plot()完整签名见 kan/MultKAN.pydef plot(self, folder./figures, beta3, metricbackward, scale0.5, tickFalse, sampleFalse, in_varsNone, out_varsNone, titleNone, varscale1.0)参数默认值作用folder./figures保存各激活子图 PNG 的目录sp_{l}_{i}_{j}.pngbeta3控制透明度tanh(beta * score)越大显示越多metricbackward重要性度量forward_n/forward_u/backwardscale0.5画布整体缩放节点、连线、散点同步缩放tickFalse是否显示坐标轴刻度开启后可查看激活输入/输出范围sampleFalse是否叠加样本散点in_vars/out_varsNone输入/输出变量名支持 LaTeX如[r$\alpha$, x]titleNone整图标题varscale1.0变量文字大小缩放一个典型的训练 → 剪枝 → 解读工作流为model(dataset[train_input])做一次前向以缓存激活值model.fit(dataset, optLBFGS, steps20, lamb0.01)带稀疏正则训练用不同beta如0.1、3、100配合metricbackward观察连接重要性model model.prune()移除不重要神经元/边后再次绘图对关键边用fix_symbolic尝试符号拟合或结合suggest_symbolic/auto_symbolic得到可解释公式再用set_mode调整符号/数值混合显示。本文全部代码均可直接运行于 pykan 仓库环境依赖见 requirements.txt安装方式见 README.md通过反复调节beta、metric、scale与sample即可快速建立起对 KAN 内部结构的直觉为后续的符号回归、剪枝与可解释性分析打下基础。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表