news 2026/9/14 10:23:06

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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 (k=3), 3 grid intervals (grid=3). model = KAN(width=[2,5,1], grid=3, k=3, seed=1, device=device) # 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_var=2, device=device) dataset['train_input'].shape, dataset['train_label'].shape

运行后终端会输出设备信息以及自动保存的检查点提示(checkpoint directory created: ./modelsaving model version 0.0),数据形状为(torch.Size([1000, 2]), torch.Size([1000, 1]))

这里的create_dataset定义于 kan/utils.py,其默认参数为:ranges=[-1,1]train_num=1000test_num=1000seed=0,返回的字典中包含train_inputtrain_labeltest_inputtest_label四个键。KAN 类本身在 kan/MultKAN.py 中实现,width=[2,5,1]表示输入 2 维、隐藏层 5 个神经元、输出 1 维。

注意plot()依赖模型保存前向传播的中间激活值(save_act=True,默认开启)。文档源码注释明确指出 "cannot plot since data are not saved. Set save_act=True first.",若在未做前向传播前直接绘图,会抛出异常model hasn't seen any data yet.。因此官方示例总是先执行一次model(dataset['train_input'])再进行绘图。

2. 绘制初始化状态下的 KAN:beta与标签参数

2.1 基础绘图:model.plot(beta=100)

对刚初始化(未训练)的模型调用:

# plot KAN at initialization model(dataset['train_input']); model.plot(beta=100)

图中从左到右依次为输入层(2 个节点)、隐藏层(5 个节点)与输出层(1 个节点),每一对节点之间的小图即为对应的 1D 激活函数(spline 样条曲线),连线粗细/透明度反映激活函数的相对重要性。

2.2 添加变量名与标题:in_vars/out_vars/title

# if you want to add variable names and title model.plot(beta=100, in_vars=[r'$\alpha$', 'x'], out_vars=['y'], title = 'My KAN')

in_varsout_vars接收字符串列表(支持 LaTeX 语法如r'$\alpha$'),title为整图标题。在plot()的源码签名中,这些参数默认均为None,另有varscale用于控制输入变量文字的缩放。

3. 训练带稀疏正则的模型:fit()lamb

为了观察训练后的结构变化,官方示例使用 L-BFGS 优化器训练 20 步并施加 L1 正则:

# train the model model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01);

训练输出:

| train_loss: 5.20e-02 | test_loss: 5.35e-02 | reg: 4.93e+00 | : 100%|█| 20/20 [00:03<00:00, 5.22it saving model version 0.1

fit()的完整签名定义于 kan/MultKAN.py(opt="LBFGS"steps=100lamb=0.update_grid=Truegrid_update_num=10等)。这里的关键是lamb正则系数:正则会抑制不重要的激活,使后续绘图中"重要连接"与"噪声连接"的对比更明显。

4. 理解beta参数:控制激活函数透明度

训练后,beta直接决定哪些激活函数在图中"浮现"出来。官方说明如下:

beta控制激活函数的透明度。更大的beta意味着更多激活函数显示出来。我们通常希望设置一个合适的beta,使得只有重要的连接在视觉上显著。

透明度公式为:

$$\text{transparency} = \tanh(\beta \cdot \phi)$$

其中 $\phi$ 的取值取决于metric

  • metric='forward_u':激活函数本身的尺度(scale)
  • metric='forward_n':归一化后的尺度(normalized scale)
  • metric='backward':特征归因分数(feature attribution score,默认)

在 kan/MultKAN.py 中对应实现为score2alpha(score) = np.tanh(beta * score),随后alpha = [score2alpha(...) for score in scores],连线透明度即由该alpha与 mask 共同决定。

4.1 默认参数绘图:model.plot()

默认情况下beta=3metric='backward'

model.plot()

4.2 大betabeta=100000):几乎全部激活可见

model.plot(beta=100000)

4.3 小betabeta=0.1):只保留最重要连接

model.plot(beta=0.1)

三张图对比可以直观体会到:beta越大,被"激活显示"的连接越多;beta越小,图中保留的越接近模型真正的"骨架"

5. 切换重要性度量:metric='forward_n' | 'forward_u' | 'backward'

metric决定连线的透明度基于哪种分数。官方示例在beta=100下逐一演示三种度量:

model.plot(metric='forward_n', beta=100) model.plot(metric='forward_u', beta=100) model.plot(metric='backward', beta=100)

从源码(kan/MultKAN.py)可以看到三种度量对应的内部数据:

metric内部数据含义
forward_nself.acts_scale归一化激活尺度
forward_uself.edge_actscale未归一化的边激活尺度
backwardself.edge_scores反向归因(特征重要性)分数(默认)

三者对应的输出图分别为:

若传入不支持的度量,源码会抛出Exception(f'metric = \'{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_th=1e-2)prune_edge(edge_th=3e-2)(见 kan/MultKAN.py),默认阈值分别为node_th=1e-2edge_th=3e-2model = model.prune()的赋值写法说明剪枝会返回一个新模型对象。

7. 调整画布尺寸:scale

scale控制整张图的物理尺寸,默认值为 0.5

model.plot(scale=0.5) # 默认值,500x400 model.plot(scale=0.2) # 缩小(200x160) model.plot(scale=2.0) # 放大(2000x1600)

在源码中,画布尺寸由figsize=(10 * scale, 10 * scale * (neuron_depth - 1) * (y0+z0))决定(kan/MultKAN.py),同时节点大小(min_spacing ** 2 * 10000 * scale ** 2)与连线宽度(lw=2 * scale)也会随scale等比缩放。

8. 叠加样本散点:sample=Trueget_act()

默认绘图只画激活函数曲线(折线)。当希望同时看到数据样本在每条激活函数上的分布时,设置sample=True

model.plot(sample=True)

样本越多散点越密集、越难分辨。官方示例通过get_act()只喂入前 20 个样本,让散点更清晰:

model.get_act(dataset['train_input'][:20]) model.plot(sample=True)

源码中散点由plt.scatter(..., s=400 * scale ** 2)绘制(kan/MultKAN.py),即散点大小同样受scale控制。get_act(x)用于以指定输入x重新计算并缓存各层激活值,供后续绘图使用。

9. 颜色编码语义:符号函数、数值函数与混合模式

KAN 支持将单个激活替换为已知符号函数(如x^2sin(x)等),绘图会以颜色区分边的类型。颜色规则在 kan/MultKAN.py 中由symbolic_masknumeric_mask共同决定:

类型symbolic_masknumeric_mask连线颜色
纯符号函数>00红色(red)
纯数值(样条)0>0黑色(black)
符号 + 数值混合>0>0紫色(purple)
被 mask 移除00白色(white)

9.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, device='cuda: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,mode='ns')

当一条边同时启用符号函数与样条函数(输出为两者之和)时显示为紫色

model.set_mode(0,1,0,mode='ns') model.plot(beta=100)

set_mode在源码(kan/MultKAN.py)中支持四种模式:

  • mode='s':纯符号(mask_s=1mask_n=0
  • mode='n':纯数值样条(mask_s=0mask_n=1
  • mode='sn'mode='ns':符号与数值混合(两者 mask 均为 1)
  • 其他取值:全部关闭(mask 均为 0)

10. 参数速查表与绘图工作流建议

plot()完整签名(见 kan/MultKAN.py):

def plot(self, folder="./figures", beta=3, metric='backward', scale=0.5, tick=False, sample=False, in_vars=None, out_vars=None, title=None, varscale=1.0)
参数默认值作用
folder"./figures"保存各激活子图 PNG 的目录(sp_{l}_{i}_{j}.png
beta3控制透明度,tanh(beta * score);越大显示越多
metric'backward'重要性度量:forward_n/forward_u/backward
scale0.5画布整体缩放(节点、连线、散点同步缩放)
tickFalse是否显示坐标轴刻度(开启后可查看激活输入/输出范围)
sampleFalse是否叠加样本散点
in_vars/out_varsNone输入/输出变量名(支持 LaTeX),如[r'$\alpha$', 'x']
titleNone整图标题
varscale1.0变量文字大小缩放

一个典型的"训练 → 剪枝 → 解读"工作流为:

  1. model(dataset['train_input'])做一次前向以缓存激活值;
  2. model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01)带稀疏正则训练;
  3. 用不同beta(如0.13100)配合metric='backward'观察连接重要性;
  4. model = model.prune()移除不重要神经元/边后再次绘图;
  5. 对关键边用fix_symbolic尝试符号拟合,或结合suggest_symbolic/auto_symbolic得到可解释公式,再用set_mode调整符号/数值混合显示。

本文全部代码均可直接运行于 pykan 仓库环境(依赖见 requirements.txt,安装方式见 README.md),通过反复调节betametricscalesample,即可快速建立起对 KAN 内部结构的直觉,为后续的符号回归、剪枝与可解释性分析打下基础。

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

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 10:19:41

分布式爬虫架构设计与性能优化实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/14 10:19:15

FP0R系列PLC在SSD精密组装中的多轴同步控制方案

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华