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: ./model、saving model version 0.0),数据形状为(torch.Size([1000, 2]), torch.Size([1000, 1]))。
这里的create_dataset定义于 kan/utils.py,其默认参数为:ranges=[-1,1]、train_num=1000、test_num=1000、seed=0,返回的字典中包含train_input、train_label、test_input、test_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_vars与out_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.1fit()的完整签名定义于 kan/MultKAN.py(opt="LBFGS"、steps=100、lamb=0.、update_grid=True、grid_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=3、metric='backward':
model.plot()4.2 大beta(beta=100000):几乎全部激活可见
model.plot(beta=100000)4.3 小beta(beta=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_n | self.acts_scale | 归一化激活尺度 |
forward_u | self.edge_actscale | 未归一化的边激活尺度 |
backward | self.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-2、edge_th=3e-2,model = 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=True与get_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^2、sin(x)等),绘图会以颜色区分边的类型。颜色规则在 kan/MultKAN.py 中由symbolic_mask与numeric_mask共同决定:
| 类型 | symbolic_mask | numeric_mask | 连线颜色 |
|---|---|---|---|
| 纯符号函数 | >0 | 0 | 红色(red) |
| 纯数值(样条) | 0 | >0 | 黑色(black) |
| 符号 + 数值混合 | >0 | >0 | 紫色(purple) |
| 被 mask 移除 | 0 | 0 | 白色(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=1,mask_n=0)mode='n':纯数值样条(mask_s=0,mask_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) |
beta | 3 | 控制透明度,tanh(beta * score);越大显示越多 |
metric | 'backward' | 重要性度量:forward_n/forward_u/backward |
scale | 0.5 | 画布整体缩放(节点、连线、散点同步缩放) |
tick | False | 是否显示坐标轴刻度(开启后可查看激活输入/输出范围) |
sample | False | 是否叠加样本散点 |
in_vars/out_vars | None | 输入/输出变量名(支持 LaTeX),如[r'$\alpha$', 'x'] |
title | None | 整图标题 |
varscale | 1.0 | 变量文字大小缩放 |
一个典型的"训练 → 剪枝 → 解读"工作流为:
model(dataset['train_input'])做一次前向以缓存激活值;model.fit(dataset, opt="LBFGS", steps=20, lamb=0.01)带稀疏正则训练;- 用不同
beta(如0.1、3、100)配合metric='backward'观察连接重要性; 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),仅供参考