FlashKAN:通过截断幂形式实现B样条KANs

arXiv cs.LG 论文

摘要

FlashKAN提出了一种加速柯尔莫哥洛夫-阿诺德网络(KANs)的方法,通过将Cox-de Boor递归替换为截断幂形式来评估B样条,提供了融合的GPU实现和一个开源包。

arXiv:2609.01956v1 Announce Type: new 摘要:柯尔莫哥洛夫-阿诺德网络 (KANs) 将可学习的B样条激活函数放置在网络边上,而不是节点上的固定激活函数。标准的Cox-de Boor递归通过k次连续传递评估这些激活函数用于k次样条,消耗了超过90%的前向传递时间。FlashKAN用截断幂形式替换了这种递归,这是逼近论中的一个经典结果,将每个均匀三次B样条表示为五个在移位节点位置的(x)_+^3项。本文做出了三项贡献:(1) 一个torch.compile融合的实现,将这些操作折叠成一个GPU内核,消除了所有递归、跨度查找和散射-聚集操作;(2) 一种有界坐标稳定化,将归一化输入限制在[0, k+1],防止了历史上促使Cox-de Boor递归出现的灾难性抵消;(3) 一个生产就绪的开源包(pip install flashkan),可作为现有KAN层的直接替代品。
查看原文
查看缓存全文

缓存时间: 2026/09/03 06:13

# FlashKAN:基于截断幂形式的 B 样条 KAN
来源:https://arxiv.org/html/2609.01956
###### 摘要

Kolmogorov-Arnold 网络 (KANs) 将可学习的 B 样条激活函数置于网络边上,而非将固定激活函数置于节点上。标准的 Cox-de Boor 递归通过对 k 次样条进行 k 次顺序传递来评估这些激活,消耗了超过 90% 的前向传播时间。

FlashKAN 用截断幂形式取代了这一递归,这是逼近理论中的一个经典结果,它将每个均匀三次 B 样条表示为在偏移节点位置处的五个 (x)_+^3 项。本文做出三项贡献:(1) 一个 torch.compile 融合实现,将这些操作合并为一个 GPU 内核,消除了所有递归、区间查找和散射-聚集操作;(2) 一种有界坐标稳定化方法,将归一化输入钳位到 [0, k+1],防止了历史上促使 Cox-de Boor 递归产生的灾难性抵消问题;(3) 一个可用于生产的开源包(pip install flashkan),可作为现有 KAN 层的直接替代品。

## 1 引言

Kolmogorov-Arnold 表示定理 (Kolmogorov, 1957 (https://arxiv.org/html/2609.01956#bib.bib11)) 保证任何多元连续函数都可以分解为一元函数与加法。KANs (Liu et al., 2024 (https://arxiv.org/html/2609.01956#bib.bib1)) 通过将一个可学习的激活函数置于每条边上来实践这一结果。每个激活函数是 B 样条基函数的加权和,这些基函数可以被检查、绘图,并在条件合适时进行符号恢复。

这种可解释性是有代价的。Cox-de Boor 递归 (Cox, 1972 (https://arxiv.org/html/2609.01956#bib.bib6); de Boor, 1971 (https://arxiv.org/html/2609.01956#bib.bib7)) 通过对 k 次样条进行 k 次顺序传递来评估 B 样条基函数。对于三次样条 (k=3),需要三次传递,每次依赖于前一次的输出。性能分析显示,仅基函数计算就占 KAN 层前向传播时间的 91% (表1 (https://arxiv.org/html/2609.01956#S2.T1))。

递归本身并非一直是标准。在 20 世纪 40 年代,Schoenberg 在阿伯丁试验场平滑弹道数据时引入了样条作为分段多项式 (Schoenberg, 1946 (https://arxiv.org/html/2609.01956#bib.bib3))。他的表示使用截断幂函数 max(0, x)^k 作为基础构建块,Curry 和 Schoenberg 通过这个基正式定义了 B 样条 (Curry and Schoenberg, 1947 (https://arxiv.org/html/2609.01956#bib.bib2))。截断幂形式在数学上很优雅:紧支集、单位分解和光滑性的证明直接得出。但 20 世纪 60 年代硬件上的早期实现因大幂次交替和中的数值抵消而受挫,并且截断幂基矩阵条件数较差 (de Boor, 2001 (https://arxiv.org/html/2609.01956#bib.bib8); Schumaker, 2007 (https://arxiv.org/html/2609.01956#bib.bib9))。Cox 和 de Boor 独立开发了递归算法 (Cox, 1972 (https://arxiv.org/html/2609.01956#bib.bib6); de Boor, 1971 (https://arxiv.org/html/2609.01956#bib.bib7)) 作为数值稳定的替代方案,而 de Boor 的教科书 (de Boor, 2001 (https://arxiv.org/html/2609.01956#bib.bib8)) 将递归确立为未来五十年的标准。

先前的工作已尝试三种策略来降低 KAN 中 B 样条评估的成本。重构 Cox-de Boor 递归 (Blealtan, 2024 (https://arxiv.org/html/2609.01956#bib.bib4)) 降低了常数因子,但保留了顺序依赖性。预计算每个区间的矩阵 (Coffman and Chen, 2025 (https://arxiv.org/html/2609.01956#bib.bib13)) 消除了递归,但引入了依赖数据的聚集操作,使 GPU 内存访问碎片化。高斯径向基函数 (Li, 2024 (https://arxiv.org/html/2609.01956#bib.bib5)) 用单个 exp() 调用完全替代了 B 样条,获得了速度但牺牲了紧支集和自动单位分解(高斯函数是非负且 C∞ 的,因此光滑性不是权衡因素)。

最近,Southworth et al. (2026) (https://arxiv.org/html/2609.01956#bib.bib10) 通过基变换矩阵,建立了样条 KAN 层与具有幂 ReLU 激活的多通道 MLP 之间的形式代数等价性。对于均匀节点,该矩阵是 Toeplitz 矩阵,其元素与经典的截断幂系数匹配。他们的工作主要利用这种等价性进行*多级训练*:样条基在不同细化级别上启用互补的松弛动态,在 PINN 精度方面带来了数量级的改进。他们指出,非递归形式快了一个等于样条次数的因子,但没有追求编译器级优化、数值稳定化或可部署的实现。

FlashKAN 弥合了代数等价性和实际部署之间的差距。贡献如下:

1. 1. 编译器融合评估。对于均匀三次 B 样条,截断幂形式简化为具有固定系数的五个 max(0,·)^3 项。torch.compile 将所有逐元素操作融合为单个 GPU 内核,消除了递归、区间查找和依赖数据的内存访问。
2. 2. 有界坐标稳定化。原始截断幂和当归一化输入 u 远在基支集 [0, k+1] 之外时会遭受灾难性抵消。在评估前将 u 钳位到此区间,可以限定所有中间项并消除抵消,而数学上正确的输出不变(因为 B 样条在其支集外恰好为零)。这解决了历史上促使 Cox-de Boor 递归的数值担忧 (de Boor, 2001 (https://arxiv.org/html/2609.01956#bib.bib8))。
3. 3. 即插即用包。FlashKAN 作为开源 Python 包(pip install flashkan,MIT 许可证)提供,与标准 KAN 层具有相同的 API。替换 Cox-de Boor 评估只需更改一个导入。

## 2 从递归到闭式

推导分四步进行。每一步都将算法转化为表达式,揭示算法所掩盖的结构。

### 2.1 步骤 1:从 de Casteljau 到 Bernstein

设 P_0, P_1, ..., P_n ∈ ℝ^d 是一系列*控制点*,t ∈ [0, 1] 是一个标量参数。两点之间的*线性插值* (lerp) 定义为 lerp(A, B, t) = (1-t) A + t B。de Casteljau 算法 (Farin, 2002 (https://arxiv.org/html/2609.01956#bib.bib12)) 通过递归应用 lerp 来评估 n 次 Bézier 曲线 Q(t):在每一级,相邻点成对插值,将点数减少一个,直到只剩一个值。

对于三次情况 (n=3),对 P_0, P_1, P_2, P_3 进行三级插值得到曲线点 Q(t)。展开并收集每个控制点上的项,得到 Bernstein 形式:
$$
Q(t) = \sum_{i=0}^3 B_{i,3}(t) \, P_i, \quad B_{i,3}(t) = \binom{3}{i} \, t^i \, (1-t)^{3-i} \tag{1}
$$
系数 $\binom{3}{i} = \{1,3,3,1\}$ 是帕斯卡三角形的第 3 行。每个插值级别乘以 $[(1-t)+t]=1$。三级产生 $((1-t)+t)^3$,由二项式定理展开为 Bernstein 基。单位分解 $\sum_{i=0}^3 B_{i,3}(t)=1$ 直接得出。

### 2.2 步骤 2:矩阵形式

将每个 Bernstein 多项式展开为 t 的幂,将 Bézier 曲线分解为三个因子:
$$
Q(t) = \underbrace{\begin{bmatrix}1 & t & t^{2} & t^{3}\end{bmatrix}}_{\mathbf{T}(t)} \underbrace{\begin{bmatrix}1 & 0 & 0 & 0 \\ -3 & 3 & 0 & 0 \\ 3 & -6 & 3 & 0 \\ -1 & 3 & -3 & 1\end{bmatrix}}_{\mathbf{M}_{\mathrm{B\acute{e}zier}}} \underbrace{\begin{bmatrix}P_{0} \\ P_{1} \\ P_{2} \\ P_{3}\end{bmatrix}}_{\mathbf{P}} \tag{2}
$$
幂向量 $\mathbf{T}(t)$ 依赖于输入。控制点 $\mathbf{P}$ 是自由参数。基矩阵 $\mathbf{M}_{\mathrm{B\acute{e}zier}}$ 是一个完全由次数决定的常数:它编码了 Bernstein 多项式的二项式展开。这种分解 $\mathbf{T}(t) \, \mathbf{M} \, \mathbf{P}$ 是通用的。不同样条类型共享相同的结构,只是 $\mathbf{M}$ 不同。依赖于输入的计算始终是 $\mathbf{T}(t)=[1,t,t^{2},t^{3}]$:四个值,极其廉价。

### 2.3 步骤 3:从 Bézier 段到 B 样条

一个三次 Bézier 曲线有四个控制点和一个段。它还有一个基本限制:*全局*控制。移动任何控制点都会影响整条曲线。为了在宽输入范围上建模函数,必须将多个 Bézier 段首尾相连,并且在每个连接点必须强制执行三个约束以确保光滑性:C^0 连续性(段相交)、C^1 连续性(切线匹配)和 C^2 连续性(曲率一致)。对于具有八个控制点的两个三次段,这六个约束(每个连接点三个,一个连接点)消耗了每个段三个自由度,八个中只剩五个自由参数。约束成本随段数线性增长:每个新连接点引入三个方程,将三个控制点锁定到其邻居。

B 样条通过构造解决了这个问题。B 样条不是拼接独立段并在事后强加约束,而是定义基函数,在整个域内*固有地*满足光滑性。给定一个非递减*节点向量* $\mathbf{t}=(t_{0},t_{1},\ldots,t_{m})$,将参数空间划分为若干区间,Cox-de Boor 递归 (de Boor, 2001 (https://arxiv.org/html/2609.01956#bib.bib8)) 定义了次数为 k 的基函数 $N_{i,k}(x)$,在每个内部节点处自动具有 C^{k-1} 连续性:
$$
\begin{aligned}
N_{i,0}(x) &= \begin{cases}1 & \text{if } t_{i} \leq x < t_{i+1} \\ 0 & \text{otherwise}\end{cases} \\
N_{i,k}(x) &= \frac{x-t_{i}}{t_{i+k}-t_{i}} \, N_{i,k-1}(x) + \frac{t_{i+k+1}-x}{t_{i+k+1}-t_{i+1}} \, N_{i+1,k-1}(x)
\end{aligned} \tag{3,4}
$$
每个基函数都具有*紧支集*:$N_{i,k}(x)$ 仅在 k+1 个连续节点区间上非零。移动一个控制点只改变该局部区域内的曲线。无需管理约束。没有自由度损失。每个控制点都是自由的。

对于间距为 h 的均匀节点向量,每个 B 样条段允许与 Bézier 相同的矩阵分解,但使用不同的基矩阵:
$$
\mathbf{M}_{\mathrm{B\text{-}spline}} = \frac{1}{6}\begin{bmatrix}1 & 4 & 1 & 0 \\ -3 & 0 & 3 & 0 \\ 3 & -6 & 3 & 0 \\ -1 & 3 & -3 & 1\end{bmatrix} \tag{5}
$$
结构 $\mathbf{T}(t) \, \mathbf{M} \, \mathbf{P}$ 相同;只有 $\mathbf{M}$ 中的系数改变,编码了 B 样条通过构造满足的光滑性约束。

该递归反映了 de Casteljau:相同的线性插值,但作用于节点区间而非 [0,1]。对于 k=3,需要三次顺序传递。每次传递依赖于前一次。表1 (https://arxiv.org/html/2609.01956#S2.T1) 显示这些传递消耗了 KAN 层前向传播时间的 91%。

表 1:KAN 层 (256×784→64, MPS GPU) 的前向传播成本分解。基计算占主导。
### 2.4 步骤 4:从 Cox-de Boor 到截断幂形式

将 de Casteljau 转化为 Bernstein 多项式的代数方法同样适用于 Cox-de Boor。对于具有常数间距 $h = t_{i+1}-t_i$ 的*均匀*节点向量,展开递归得到截断幂函数的有限差 (Curry and Schoenberg, 1947 (https://arxiv.org/html/2609.01956#bib.bib2); Schoenberg, 1946 (https://arxiv.org/html/2609.01956#bib.bib3)):
$$
N_{i,k}(x) = \frac{1}{k!\,h^{k}} \sum_{j=0}^{k+1}(-1)^{j}\binom{k+1}{j}\,\max(0,\;x-t_{i+j})^{k} \tag{6}
$$
对于三次样条 (k=3),代入 $u=(x-g_i)/h$ 得到:
$$
N_{i}(u) = \frac{1}{6}\Big[(u)_{+}^{3}-4(u-1)_{+}^{3}+6(u-2)_{+}^{3}-4(u-3)_{+}^{3}+(u-4)_{+}^{3}\Big] \tag{7}
$$
其中 $g_i$ 是第 i 个支集区间的起点,$(z)_+ = \max(0,z)$。系数 $\{1,-4,6,-4,1\}$ 是帕斯卡三角形第四行的交替符号除以 $3!=6$。这与 Southworth et al. (2026) (https://arxiv.org/html/2609.01956#bib.bib10) 通过其基变换分析确定的系数模板相同。

两种推导之间的平行性是精确的。de Casteljau 对应 Bernstein,Cox-de Boor 对应截断幂形式。在两种情况下,一个由嵌套插值构建的算法都允许一个由二项式系数构建的闭式表达式。算法是顺序的。表达式是并行的。

#### 计算性质.
公式 7 (https://arxiv.org/html/2609.01956#S2.E7) 继承了矩阵形式的分离性。依赖于输入的部分是 $u=(x-g_i)/h$:一次减法和一次乘法。基矩阵 $\mathbf{M}$ 被吸收到固定系数 $\{1,-4,6,-4,1\}/6$ 中。控制点是可学习的权重。没有顺序依赖,没有依赖数据的内存访问,只有逐元素操作,所有这些 torch.compile 都融合到一个内核中。

### 2.5 有界坐标稳定化

公式 7 (https://arxiv.org/html/2609.01956#S2.E7) 中的截断幂和在无限精度下是精确的。然而,在有限精度算术中,当 u 远在支集区间 [0,4] 之外时,五项大小为 Θ(u^3) 的交替和遭受灾难性抵消。当 u≥4 或 u≤0 时,数学上正确的值恰好为零,但浮点评估会产生一个残差,其增长为 O(ε_mach·u^3)。这一担忧正是促使 Cox 和 de Boor 开发递归算法的原因 (Cox, 1972 (https://arxiv.org/html/2609.01956#bib.bib6); de Boor, 1971 (https://arxiv.org/html/2609.01956#bib.bib7); de Boor, 2001 (https://arxiv.org/html/2609.01956#bib.bib8))。

对于 KAN 层,风险是具体的。隐藏层激活在训练过程中可能漂移到标称网格域之外,即使原始输入是归一化的,也会产生大的 u 值。在混合精度(float16)下,一旦 u 超过大约 40,u^3 就会溢出。

解决方法很简单。在评估公式 7 (https://arxiv.org/html/2609.01956#S2.E7) 之前,将归一化坐标钳位到支集内:
$$
\bar{u} = \operatorname{clamp}(u,\;0,\;k+1) \tag{8}
$$
并计算 $N_{i}(\bar{u})$。

相似文章

SechKAN: 基于双曲正割函数的Kolmogorov-Arnold网络

arXiv cs.LG

SechKAN 是一种新颖的 Kolmogorov-Arnold 网络架构,使用双曲正割函数作为基函数,在函数拟合、偏微分方程问题和图像分类任务中取得了具有竞争力的性能,同时保持了与多层感知机相当的参数效率。

通过Kolmogorov-Arnold网络在FPGA上实现超快机器学习

Hacker News Top

本文介绍了作者的硕士论文,该论文利用Kolmogorov-Arnold网络(KAN)在FPGA上实现超快机器学习,通过自定义硬件架构实现亚微秒级推理和在线学习。文章引用了两篇已接收的论文:基于LUT评估的KANELÉ(FPGA 2026最佳论文奖)以及一种在FPGA上进行在线学习的方法(ICML 2026)。