pykan 网格机制深度解析:B 样条参数化、update_grid_from_samples 与自适应网格更新
2026/9/14 2:50:43 网站建设 项目流程

pykan 网格机制深度解析:B 样条参数化、update_grid_from_samples 与自适应网格更新

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

导读

KAN(Kolmogorov-Arnold Networks)的核心特性之一,是把 B 样条(B-spline)嵌入神经网络作为可学习的激活函数。但样条只在已知的有界区间内对函数有效,而神经网络中激活值的取值范围在训练过程中会不断变化——因此必须有一套机制,让网格(grid)跟随样本分布动态调整。本文以 docs/API_demo/API_5_grid.rst 为骨架,结合 kan/spline.py 与 kan/KANLayer.py 的源码实现,讲清三件事:KAN 如何用 B 样条参数化激活函数、激活函数为何是"残差 + 样条"两部分、以及update_grid_from_samplesgrid_eps如何让网格自适应数据。读完后你将能够熟练地通过gridkgrid_eps等参数控制 KAN 的样条拟合精度,并知道在数据分布与默认网格不符时如何正确调整网格。

为什么 KAN 需要更新网格

KAN 将样条嵌入神经网络。样条(spline)本质上是一种分段多项式逼近方法,它只能在已知的有界区域内近似函数。然而在神经网络训练过程中:

  • 初始阶段激活值可能落在[-1, 1]之类的默认范围内;
  • 随着参数更新,各层的输入激活(pre-activations)范围会漂移、扩大或收缩;
  • 如果网格范围与激活值的实际范围严重不匹配,样条在数据区间外将失去拟合能力(或退化)。

因此,KAN 需要根据当前样本动态更新网格,让网格节点始终覆盖激活值的真实分布。这就是update_grid_from_samples存在的原因。下面的官方文档演示(docs/API_demo/API_5_grid.ipynb)将从最基础的样条参数化讲起。

B 样条如何参数化:G 个区间、k 阶多项式,共 G+k 个基函数

网格与基函数

在 pykan 中,一条 1D 样条由以下要素决定:

  • G:网格区间数(grid intervals),即把定义域分成G段;
  • k:样条阶数(piecewise polynomial order,即样条的度数);
  • grid:网格节点,共有G+1个内部节点,G+k个 B 样条基函数。

样条函数是所有基函数的线性组合:

$${\rm spline}(x)=\sum_{i=0}^{G+k-1} c_i B_i(x)$$

其中 $B_i(x)$ 是 B 样条基函数,$c_i$ 是对应的系数。在 pykan 中:

  • B_batch 负责在给定输入x、网格grid和阶数k时计算所有基函数值,输出形状为(batch, in_dim, G+k)
  • extend_grid 负责把网格向两端各扩展k个点(以均匀步长h外推),保证样条在边界处的基函数完整定义;
  • coef2curve 通过torch.einsum('ijk,jlk->ijl', b_splines, coef)完成"系数 × 基函数"的求和,即实现上面的线性组合公式。

官方示例:绘制基函数

下面的代码构造一个定义在[-1, 1]上、G=5k=3的网格,并绘制全部G+k = 8个基函数:

from kan.spline import B_batch import torch import matplotlib.pyplot as plt import numpy as np from kan.spline import extend_grid # consider a 1D example. # Suppose we have grid in [-1,1] with G intervals, spline order k G = 5 k = 3 grid = torch.linspace(-1,1,steps=G+1)[None,:] grid = extend_grid(grid, k_extend=k) # and we have sample range in [-1,1] x = torch.linspace(-1,1,steps=1001)[None,:] basis = B_batch(x, grid, k=k) for i in range(G+k): plt.plot(x[0].detach().numpy(), basis[0,:,i].detach().numpy()) plt.legend(['B_{}(x)'.format(i) for i in np.arange(G+k)]) plt.xlabel('x') plt.ylabel('B_i(x)')

运行后可以看到 8 条形状各异的基函数曲线:

从源码看,B_batch采用递归方式计算基函数:k=0时为区间指示函数(0 阶);k>0时由k-1阶基函数递推得到(见 kan/spline.py),并对退化网格做了torch.nan_to_num保护(kan/spline.py)。

验证 KAN 确实实现了样条:激活 = 残差 + 样条

由于样条计算已经内置于 KAN,我们不需要自己实现。但官方文档给出了一段可验证的代码,确认 KAN 内部确实按上述公式运作。初始化一个[1,1]的 KAN——它本质上就是一条 1D 样条:

from kan import KAN model = KAN(width=[1,1], grid=G, k=k) # obtain coefficients c_i model.act_fun[0].coef assert(model.act_fun[0].coef[0].shape[1] == G+k) # the model forward model_output = model(x[0][:,None]) # spline output spline_output = torch.einsum('j,ij->i',model.act_fun[0].coef[0][0], basis[0])[:,None] torch.mean((model_output - spline_output)**2)

输出:

checkpoint directory created: ./model saving model version 0.0 tensor(0.0099, grad_fn=<MeanBackward0>)

注意两点:

  1. 构造KAN时会自动创建./model检查点目录并保存0.0版本(auto_save=True的默认行为,见 kan/MultKAN.py 的构造函数参数);
  2. coef的最后一维确实是G+k = 8,断言通过——系数个数与基函数个数严格对应
  3. 但模型输出与纯样条输出并不相等(MSE = 0.0099)。为什么?

激活函数的两部分结构

因为 pykan 把激活函数建模为残差函数 + 样条函数两部分的和:

$$\phi(x)={\rm scale_base}b(x)+{\rm scale_sp}{\rm spline}(x)$$

默认情况下残差函数 $b(x)={\rm silu}(x)=x/(1+e^{-x})$,即 SiLU 激活。在 KANLayer.forward 中对应实现为:

base = self.base_fun(x) # 残差部分 b(x) y = coef2curve(x_eval=x, grid=self.grid, coef=self.coef, k=self.k) # 样条部分 y = self.scale_base[None,:,:] * base[:,:,None] + self.scale_sp[None,:,:] * y

其中base_fun的默认值就是torch.nn.SiLU()(kan/KANLayer.py),scale_basescale_sp是默认可训练的缩放系数。把残差部分也加入后,MSE 变为 0:

# residual output residual_output = torch.nn.SiLU()(x[0][:,None]) scale_base = model.act_fun[0].scale_base scale_sp = model.act_fun[0].scale_sp torch.mean((model_output - (scale_base * residual_output + scale_sp * spline_output))**2)

输出tensor(0., grad_fn=<MeanBackward0>)——完全吻合。这验证了:KAN 的前向输出 = scale_base × SiLU(x) + scale_sp × 样条(x)。这也解释了为什么在实际使用中,KAN 即使样条部分退化,残差部分(SiLU)仍能提供基础的函数拟合能力。

网格与数据不匹配时:用 update_grid_from_samples 对齐

问题场景

默认情况下网格范围是[-1, 1]grid_range=[-1, 1],见 kan/KANLayer.py)。但如果你的数据在[-10, 10][-0.5, 0.5],网格与数据就不匹配了。此时应调用update_grid_from_samples让网格向样本对齐。

用法示例:数据在 [-10, 10]

model = KAN(width=[1,1], grid=G, k=k) print(model.act_fun[0].grid) # by default, the grid is in [-1,1] x = torch.linspace(-10,10,steps = 1001)[:,None] model.update_grid_from_samples(x) print(model.act_fun[0].grid) # now the grid becomes in [-10,10]. We add a 0.01 margin in case x have zero variance

输出:

Parameter containing: tensor([[-2.2000, -1.8000, -1.4000, -1.0000, -0.6000, -0.2000, 0.2000, 0.6000, 1.0000, 1.4000, 1.8000, 2.2000]]) Parameter containing: tensor([[-22., -18., -14., -10., -6., -2., 2., 6., 10., 14., 18., 22.]])

用法示例:数据在 [-0.5, 0.5]

model = KAN(width=[1,1], grid=G, k=k) print(model.act_fun[0].grid) # by default, the grid is in [-1,1] x = torch.linspace(-0.5,0.5,steps = 1001)[:,None] model.update_grid_from_samples(x) print(model.act_fun[0].grid) # now the grid becomes in [-10,10]. We add a 0.01 margin in case x have zero variance

输出:

Parameter containing: tensor([[-2.2000, -1.8000, -1.4000, -1.0000, -0.6000, -0.2000, 0.2000, 0.6000, 1.0000, 1.4000, 1.8000, 2.2000]]) Parameter containing: tensor([[-1.1000, -0.9000, -0.7000, -0.5000, -0.3000, -0.1000, 0.1000, 0.3000, 0.5000, 0.7000, 0.9000, 1.1000]])

可以观察到两个规律:

  • 更新后网格的内部区间精确落在数据范围[-10, 10][-0.5, 0.5]上(5 个区间、步长分别为 4 和 0.2);
  • 网格两端各额外扩展了 3 个点(例如[-22, -18, -14][14, 18, 22]),这正是extend_grid(grid, k_extend=k)在 kan/KANLayer.py 中的行为——为保证边界处 $k$ 阶基函数完整,必须向外扩展 $k$ 个节点。

源码视角:网格更新如何作用到所有层

KAN.update_grid_from_samples是模型级 API(kan/MultKAN.py),它会遍历所有层

for l in range(self.depth): self.get_act(x) self.act_fun[l].update_grid_from_samples(self.acts[l])

即:先把输入逐层前传得到每一层的激活值self.acts[l],再调用每个KANLayerupdate_grid_from_samples(kan/KANLayer.py)逐层更新网格。因此一次调用即可让所有层的所有样条网格与样本对齐

KANLayer层内,更新流程是:

  1. 对样本x按列排序得到x_pos
  2. 在当前网格上对排序后的样本求值y_eval = coef2curve(x_pos, self.grid, self.coef, self.k)
  3. 依据grid_eps计算新网格(见下一节);
  4. 用最小二乘curve2coef在新网格上重新拟合系数coef(kan/KANLayer.py),保证更新网格前后函数值(近似)不变。

均匀网格还是自适应网格:grid_eps 参数

两种极端与插值

官方文档指出有两种网格设计思路:

  1. 均匀网格(uniform grid):节点等间距分布;
  2. 自适应网格(adaptive grid):基于样本分布划分,使每个区间内大致有相同数量的样本(即按分位数划分)。

pykan 提供参数grid_eps在两者之间插值:

  • grid_eps = 1:完全均匀网格;
  • grid_eps = 0:完全自适应(分位数)网格;
  • 0 < grid_eps < 1:两者线性混合。

在 KANLayer.update_grid_from_samples 的源码中可以看到核心实现:

ids = [int(batch / num_interval * i) for i in range(num_interval)] + [-1] grid_adaptive = x_pos[ids, :].permute(1,0) # 按分位数取节点 ... grid_uniform = grid_adaptive[:,[0]] - margin + h * torch.arange(num_interval+1,)[None, :] grid = self.grid_eps * grid_uniform + (1 - self.grid_eps) * grid_adaptive

即:自适应网格节点直接取排序后样本的等分位点(保证每个区间样本数大致相等),均匀网格则在自适应网格的首尾之间等距展开,最终结果按grid_eps加权混合。

数值演示:正态分布样本下的两种网格

torch.normal(0, 1, size=(1000,1))采样,对比两种配置:

均匀网格(默认 grid_eps 语义)

# uniform grid model = KAN(width=[1,1], grid=G, k=k) print(model.act_fun[0].grid) # by default, the grid is in [-1,1] x = torch.normal(0,1,size=(1000,1)) model.update_grid_from_samples(x) print(model.act_fun[0].grid)
Parameter containing: tensor([[-2.2000, -1.8000, -1.4000, -1.0000, -0.6000, -0.2000, 0.2000, 0.6000, 1.0000, 1.4000, 1.8000, 2.2000]]) Parameter containing: tensor([[-8.3431, -6.8772, -5.4114, -3.9455, -2.4797, -1.0138, 0.4520, 1.9179, 3.3837, 4.8496, 6.3154, 7.7813]])

均匀网格下节点等间距:内部节点为-8.34, -6.88, -5.41, -3.95, -2.48, -1.01, 0.45, 1.92, 3.38, 4.85, 6.32, 7.78(步长 1.56)。

自适应网格(grid_eps=0)

# adaptive grid based on sample distribution model = KAN(width=[1,1], grid=G, k=k, grid_eps = 0.) print(model.act_fun[0].grid) # by default, the grid is in [-1,1] x = torch.normal(0,1,size=(1000,1)) model.update_grid_from_samples(x) print(model.act_fun[0].grid)
Parameter containing: tensor([[-2.2000, -1.8000, -1.4000, -1.0000, -0.6000, -0.2000, 0.2000, 0.6000, 1.0000, 1.4000, 1.8000, 2.2000]]) Parameter containing: tensor([[-8.3431, -6.8772, -5.4114, -3.9455, -0.8148, -0.2487, 0.2936, 0.8768, 3.3837, 4.8496, 6.3154, 7.7813]])

对比可见:自适应网格的首尾节点与均匀网格一致(都覆盖了样本的极值),但中间节点明显向 0 附近聚集-0.81, -0.25, 0.29, 0.88),因为正态分布的大多数样本集中在均值附近——每个区间内的样本数大致相等,从而在数据密集处获得更高的分辨率。

当前仓库的默认值说明

官方教程中说明"默认grid_eps = 1(均匀网格)"。需要指出的是:在当前仓库源码中,KAN/MultKANKANLayer构造函数的实际默认值为grid_eps=0.02(见 kan/MultKAN.py 与 kan/KANLayer.py),即默认近似均匀、略带回退到自适应的余量。在使用时如需严格的均匀或自适应网格,请显式传入grid_eps=1.0grid_eps=0.0

实践要点小结

  1. 参数对应关系G(区间数)与k(样条阶数)共同决定基函数个数G+k,也决定了coef的维度(in_dim, out_dim, G+k),在构造KAN(width=[...], grid=G, k=k)时直接传入。
  2. 激活函数结构φ(x) = scale_base · b(x) + scale_sp · spline(x),默认b(x)=SiLU。调试时若发现模型输出与"纯样条"不一致,属正常现象,请把残差部分一并计入。
  3. 网格对齐:数据分布与默认[-1, 1]网格不符时,调用model.update_grid_from_samples(x),它会作用于所有层,并按样本重新拟合系数,更新前后函数值保持一致。
  4. 网格风格grid_eps=1均匀网格适合样本近似均匀分布;grid_eps=0分位数自适应网格适合样本集中在局部区域的场景;中间值做线性插值。当前仓库默认grid_eps=0.02,请按需显式指定。
  5. 边界外推:无论哪种网格,更新后都会向两端各扩展k个节点(由extend_grid完成),这是保证样条边界定义完整所必需的。

相关源码入口:kan/spline.py(基函数与系数转换)、kan/KANLayer.py(单层网格更新)、kan/MultKAN.py(全模型网格更新)。完整可运行示例见 docs/API_demo/API_5_grid.ipynb 与 tutorials/API_demo/API_5_grid.ipynb。

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

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询