KAN神经网络:从可学习激活函数到高精度函数逼近的架构革新
1. 项目概述:从MLP到KAN,一次神经网络架构的根本性反思
最近在复现和思考一些前沿的神经网络架构时,Kolmogorov–Arnold Networks(KAN)这个概念让我眼前一亮。它不像Transformer或者扩散模型那样,在既有框架上做加法,而是直接对神经网络最基础的构建块——全连接层(MLP)——提出了一个根本性的质疑和重构。简单来说,KAN用可学习的、光滑的“样条函数”替换了MLP中固定的、简单的激活函数(如ReLU),并将权重参数从连接线上移到了这些激活函数内部。这个看似简单的“交换”,带来的却是模型精度、可解释性和参数效率的显著提升。如果你对神经网络底层原理感兴趣,或者在实际项目中遇到了MLP拟合复杂函数能力不足、参数量爆炸、像个黑盒子一样难以理解的问题,那么KAN绝对值得你花时间深入研究。它不是一个简单的“新模型”,而是一种全新的、更接近函数逼近本质的建模思路。
2. 核心思路拆解:为什么是“Kolmogorov–Arnold”?
要理解KAN,必须先理解它名字背后的两个数学巨人:Kolmogorov和Arnold。他们在1950年代证明了一个著名的定理,即Kolmogorov–Arnold表示定理。这个定理指出,任何多元连续函数都可以表示为有限个单变量连续函数和加法运算的组合。这是一个非常强大的存在性定理,它从理论上保证了用单变量函数的叠加来精确表示复杂多元函数的可能性。
然而,经典的MLP架构虽然也受此启发,但走了一条不同的路。MLP将可学习的参数放在线性权重矩阵中,而使用固定的、简单的非线性激活函数(如sigmoid, ReLU)。这相当于把“学习”的重担完全交给了线性组合部分,而非线性部分则是一个固定的、粗糙的“模具”。KAN的思路则是对这个定理更直接、更忠实的实现:它将网络中的每一个“边”都看作一个可学习的单变量函数,通常用B-样条(B-spline)来实现。网络中的节点则只进行简单的求和操作。这样一来,可学习的、复杂的非线性变换被分配到了每一条连接上,而节点只是信息的聚合点。
这种架构上的“交换”带来了几个根本性的优势:
- 更高的表达能力:样条函数可以光滑地拟合任意复杂的单变量函数,远比ReLU这样的分段线性函数强大。这使得KAN理论上能以更少的参数达到更高的精度。
- 内在的可解释性:由于每个激活函数都是可学习的单变量函数,我们可以直接可视化这些函数,观察网络在每一层、每一个特征维度上学到了什么样的变换。这为理解神经网络内部工作机制打开了一扇窗。
- 更优的缩放定律:论文中的实验表明,KAN在逼近复杂函数时,其精度随参数增加的提升速度(缩放定律)远优于MLP,这意味着它用更少的计算资源就能获得更好的效果。
注意:KAN不是要完全取代MLP。MLP在特征学习、大规模数据拟合上经过了几十年的优化,其工程实践非常成熟。KAN目前更像一个“特长生”,在需要高精度函数逼近、科学发现(如从数据中发现物理公式)和模型可解释性的场景下,展现出独特的潜力。
2.1 从固定激活到可训练激活函数:核心范式转移
传统MLP中,激活函数是固定的、不可训练的。你选择ReLU,那它在整个训练过程中就是f(x)=max(0,x),不会改变。而在KAN中,所谓的“激活函数”本身就是一个由参数控制的复杂函数,通常是样条函数。这些参数会随着梯度下降一起被优化。
这带来了一个根本性的变化:网络的学习能力不再仅仅依赖于堆叠更多的层和神经元(宽度/深度),而是可以精细地调整每一条信息通路上的非线性变换形状。举个例子,想象一下在拟合一个物理过程时,网络可能需要在某个中间变量上施加一个类似正弦波的变换,在另一个变量上施加一个对数变换。传统的MLP需要很多神经元和层来“拼凑”出这种效果,而KAN可能直接在对应的边上学习出接近正弦或对数的样条函数,结构更清晰,参数更少。
2.2 样条函数:如何实现“可训练”的光滑函数?
那么,如何用一组可训练的参数来表示一个任意形状的光滑单变量函数呢?KAN选择的是B-样条(B-spline)。这是一种在计算机图形学和数值分析中非常成熟的技术,用于构造光滑的曲线。
简单理解,B-样条是通过一组“控制点”和“基函数”来定义一条曲线。基函数是固定的、局部支撑的(只在一小段区间内非零),而控制点的y坐标(或高度)就是我们的可训练参数。通过调整这些控制点的高度,我们就可以改变这条样条曲线的形状。在KAN中,每一条边上的激活函数,就是由这样一组可训练的控制点参数定义的B-样条曲线。
使用样条的优势在于:
- 光滑性:可以确保函数是
k-1次可导的(k是样条阶数),这对于优化和表示光滑物理过程很重要。 - 局部性:调整一个控制点,只会影响曲线局部的形状,这使得学习更稳定。
- 计算高效:样条函数的值可以通过德布尔算法快速计算,其前向传播和反向传播(求导)都有成熟高效的实现。
3. KAN网络架构与实操定义
一个KAN层可以看作是对传统线性层+激活函数组合的彻底重构。假设输入维度为n_in,输出维度为n_out。
- 传统MLP层:
Y = σ(W * X + b),其中W是n_out x n_in的权重矩阵,b是偏置,σ是固定的逐元素激活函数。 - KAN层:
Y = Φ(X),这里Φ本身就是一个函数矩阵。具体计算是,对于输出Y的第j个分量y_j,有:y_j = Σ_{i=1}^{n_in} φ_{j,i}(x_i)其中,φ_{j,i}就是连接输入x_i到输出y_j的那条边上的可学习单变量函数(通常为样条函数)。所有φ_{j,i}的集合就构成了函数矩阵Φ。
3.1 网络分层与堆叠
和MLP一样,KAN也可以堆叠成深度网络。一个L层的KAN网络可以表示为:KAN(x) = (Φ_{L-1} ◦ Φ_{L-2} ◦ ... ◦ Φ_1 ◦ Φ_0)(x)其中,每一层Φ_l都是一个函数矩阵。值得注意的是,由于每一层的函数输入输出都是标量到标量的映射,深度KAN的复合在数学上仍然是良定义的,并且得益于样条函数的光滑性,整个网络也是光滑的。
在实操中,我们需要为网络中每一个φ_{j,i}函数定义其样条参数。这包括:
- 定义域网格:将函数的输入范围(通常经过归一化)划分为若干个小区间(网格)。
- 样条阶数:例如3阶(二次样条)或4阶(三次样条),阶数越高曲线越光滑。
- 控制点系数:每个网格区间上的控制点高度,这就是主要的可训练参数。
3.2 前向传播计算详解
前向传播时,对于输入x_i,我们需要计算φ_{j,i}(x_i)。这个过程是:
- 定位:根据
x_i的值,确定它落在定义域网格的哪一个区间。 - 基函数求值:根据样条的阶数
k,取出该区间及其附近k个区间对应的B-样条基函数。 - 加权求和:用该位置上的
k个基函数的值,乘以对应的k个可训练的控制点系数,求和即得到函数值φ_{j,i}(x_i)。 - 聚合:对所有输入维度
i的φ_{j,i}(x_i)求和,得到输出y_j。
这个过程虽然描述起来复杂,但可以通过向量化和高效的样条求值库(如torch中的torch.spline或自定义CUDA核)来实现,确保训练和推理的效率。
4. 实战:从零构建一个KAN层(PyTorch思路)
理解了原理,我们来看看如何用PyTorch实现一个简易的KAN层。这里我们聚焦于核心思想,省略一些工程优化细节。
首先,我们需要实现B-样条函数。为了简化,我们可以使用一维卷积来高效实现B-样条基函数的求值和与系数的加权求和。
import torch import torch.nn as nn import torch.nn.functional as F import math class BSplineActivation(nn.Module): """ 实现一个一维的B-样条可学习激活函数。 假设输入x已归一化到[0, 1]区间。 """ def __init__(self, num_control_points=10, spline_order=3): super().__init__() self.num_control_points = num_control_points # 控制点数量(包括区间端点) self.spline_order = spline_order # 样条阶数,k # 可训练参数:控制点系数(高度) # 初始化为一个平坦的线,例如从-0.1到0.1的均匀分布 self.control_coeffs = nn.Parameter( torch.zeros(num_control_points).uniform_(-0.1, 0.1) ) # 构造均匀网格。grid的范围需要稍微扩展,以处理边界处的基函数。 self.grid = torch.linspace(-spline_order/(num_control_points-1), 1+spline_order/(num_control_points-1), num_control_points + spline_order) def forward(self, x): """ x: 任意形状的张量,值应在[0,1]附近。 返回:与x同形状的张量,每个元素通过样条函数映射。 """ # 将x从[0,1]映射到样条节点索引空间。 # 这里是一个简化实现,实际需要更精确的德布尔算法。 # 我们使用线性插值来近似。 x_scaled = x * (self.num_control_points - self.spline_order) + self.spline_order/2 x_scaled = x_scaled.clamp(min=self.grid[self.spline_order].item(), max=self.grid[-self.spline_order-1].item()) # 为了简化,我们使用PyTorch的grid_sample进行一维插值(将控制点视为一维图像)。 # 这并非标准的B-样条求值,但作为一个概念演示。 # 首先,将控制点系数reshape为“一维图像” (1, 1, num_control_points) control_image = self.control_coeffs.view(1, 1, -1) # 将x的坐标归一化到[-1, 1](grid_sample的要求) sample_grid = (x_scaled.unsqueeze(-1) * 2 - 1).view(-1, 1, 1, 1) # 形状: (N, 1, 1, 1) # 使用双线性插值模式,模拟一个近似的样条插值。 # 注意:这只是一个粗糙的近似,真正的B-样条需要更复杂的基函数卷积。 output = F.grid_sample(control_image.unsqueeze(0), sample_grid, mode='bilinear', align_corners=False) return output.squeeze() class KANLayer(nn.Module): """ 一个完整的KAN层。 """ def __init__(self, input_dim, output_dim, num_control_points=10, spline_order=3): super().__init__() self.input_dim = input_dim self.output_dim = output_dim self.spline_order = spline_order self.num_control_points = num_control_points # 为每一对 (输入神经元i, 输出神经元j) 创建一个样条激活函数。 # 我们使用ModuleList来存储。 self.spline_functions = nn.ModuleList() for _ in range(output_dim): row = nn.ModuleList([BSplineActivation(num_control_points, spline_order) for _ in range(input_dim)]) self.spline_functions.append(row) # 可选的偏置项 self.bias = nn.Parameter(torch.zeros(output_dim)) def forward(self, x): # x shape: (batch_size, input_dim) batch_size = x.shape[0] output = torch.zeros(batch_size, self.output_dim, device=x.device) for j in range(self.output_dim): for i in range(self.input_dim): # 取出第j个输出神经元对第i个输入维度的样条函数 phi = self.spline_functions[j][i] # 计算 phi(x[:, i]) 并累加 output[:, j] += phi(x[:, i]) # 加上偏置 output[:, j] += self.bias[j] return output实操心得:上面的实现为了清晰牺牲了效率。在真实应用中,
forward函数里的双重循环是性能杀手。必须进行向量化优化。一种思路是将所有样条函数的控制点系数存储在一个大的张量中(output_dim, input_dim, num_control_points),然后利用一维卷积或自定义CUDA核,一次性计算所有φ_{j,i}(x_i)。这是实现高性能KAN的关键。
4.1 初始化与归一化技巧
KAN对初始化比较敏感。如果所有样条函数初始化为零附近,那么梯度可能很小。论文中采用了一种巧妙的初始化方法:将样条函数初始化为一个近似silu(x) = x * sigmoid(x)的函数。这个函数在零点附近近似线性,远离零点时饱和,是一个很好的默认非线性。
另一个关键点是输入归一化。由于样条函数通常定义在一个固定的区间(如[-1, 1]或[0, 1]),我们必须确保每一层的输入值大致落在这个范围内。这可以通过在每一层KAN前添加一个可学习的仿射归一化层(学习数据的均值和方差)来实现,或者使用批量归一化(BatchNorm)。否则,输入值超出样条定义域会导致外推,而样条在外推区域的行为可能不稳定。
5. KAN的训练策略与调参经验
训练KAN与训练MLP有相似之处,也有其特殊性。
- 损失函数:和MLP一样,根据任务选择(MSE用于回归,交叉熵用于分类等)。
- 优化器:AdamW通常是安全的选择。由于KAN参数可能更多(样条控制点),权重衰减(
weight_decay)对于防止过拟合很重要。 - 学习率:可以尝试与MLP相似的学习率调度,如余弦退火。由于样条参数需要精细调整,学习率不宜过大。
- 正则化:
- L1正则化:对样条函数的控制点系数施加L1正则,可以促使许多控制点系数变为零,从而实现网络稀疏化。这是KAN一个非常强大的特性,可以自动学习出简洁的网络结构,极大提升可解释性。
- 样条平滑性正则:可以添加一个惩罚项,惩罚相邻控制点系数之间的二阶差分,鼓励学习出更光滑的函数,防止过拟合噪声。
5.1 网格细化:动态增加表达能力
这是KAN训练中一个非常独特的技巧。我们可以从一个小网格(较少的控制点)开始训练。当损失平台期后,我们可以对每个样条函数进行网格细化:将当前的样条函数用更高分辨率的网格(更多的控制点)重新参数化,然后继续训练。这相当于在不改变网络函数表达能力的前提下,增加了其拟合细节的能力。这个过程可以迭代进行,类似于自适应增加模型容量。
# 伪代码:网格细化思路 def refine_spline_grid(spline_func, new_num_control_points): # 1. 评估旧样条在精细网格上的值 fine_grid = torch.linspace(0, 1, new_num_control_points) with torch.no_grad(): fine_values = spline_func(fine_grid) # 2. 用这些值作为新的、更精细的样条控制点的初始值(可能需要插值) # 3. 替换旧的样条函数参数 spline_func.num_control_points = new_num_control_points spline_func.control_coeffs.data = ... # 根据fine_values初始化新参数6. 可解释性实践:如何“看懂”一个KAN?
KAN最大的魅力之一在于其可解释性。训练完成后,我们可以可视化每一层的样条函数φ_{j,i}。
- 单函数可视化:对于某个特定的
φ_{j,i},我们可以在其定义域内均匀采样,画出输入x和输出φ(x)的关系图。这直接告诉我们,从输入特征i到中间特征j,网络学到了什么样的非线性变换。是线性?是二次型?是周期函数?还是某种复杂的饱和函数?一目了然。 - 函数矩阵可视化:将一层中所有的
φ_{j,i}画成一个小图像矩阵,行对应输出神经元j,列对应输入神经元i。这可以让我们快速浏览该层整体的变换模式,发现哪些连接是重要的(函数变化剧烈),哪些是次要的(函数近似为零或常数)。 - 符号化尝试:对于学到的光滑样条函数,我们可以尝试用符号回归工具(如
PySR)去拟合一个简单的数学表达式(如sin,exp,polynomial)。如果成功,我们甚至可以用一个简单的公式来近似描述该连接的作用,从而实现极高层次的模型解释。
例如,在一个用于发现物理公式的KAN中,你可能会发现某个φ函数非常接近sin(5.0*x),这强烈暗示了数据中存在的周期性成分。
7. 常见问题与避坑指南
在实际尝试KAN的过程中,我遇到了不少坑,这里总结一下:
训练不稳定或发散
- 可能原因:输入未归一化,导致值域远超样条定义域
[0,1],在外推区域梯度爆炸。 - 解决:在每一层KAN前强制添加输入归一化层(如
LayerNorm或可学习的仿射变换)。务必监控每一层输入的均值和方差。 - 可能原因:学习率过高。样条控制点参数对学习率比较敏感。
- 解决:使用较小的学习率(例如
1e-3或1e-4),并配合学习率热身(Warmup)策略。
- 可能原因:输入未归一化,导致值域远超样条定义域
模型表现不如简单MLP
- 可能原因:网格分辨率(控制点数量)太低,模型表达能力不足。
- 解决:尝试增加
num_control_points,或采用网格细化策略,从粗到细训练。 - 可能原因:样条阶数太低(如1阶就是分段线性,和ReLU类似),无法拟合光滑函数。
- 解决:使用3阶(二次)或4阶(三次)样条,这是光滑性和灵活性的较好折衷。
- 可能原因:任务本身不适合。KAN在低维、高精度函数逼近和可解释性任务上优势明显,但在极高维、非结构化的数据(如图像、文本)的特征提取上,目前可能不如经过高度优化的CNN或Transformer。不要把它当作万能药。
训练速度慢
- 可能原因:使用了低效的实现(如Python循环)。这是初期实现最大的瓶颈。
- 解决:必须实现向量化的样条求值。可以寻找开源的优化实现(如
pykan库),或自己用torch的einsum和grid_sample进行向量化,甚至编写CUDA扩展。
稀疏化后模型崩溃
- 可能原因:L1正则化系数
λ太大,过早地将太多重要的控制点或连接置零。 - 解决:采用渐进式稀疏化。先在不加或加很小L1正则的情况下训练模型至收敛,然后缓慢增加
λ,并配合较小的学习率进行微调,让网络有机会将重要的信息重新分配到剩下的连接中。
- 可能原因:L1正则化系数
过拟合
- 可能原因:网格分辨率过高,控制点过多,而训练数据不足。
- 解决:使用更强的L1/L2正则化,或采用
Dropout(虽然不常用,但可以尝试在求和后添加)。最根本的是使用网格细化,从较粗的网格开始,仅在需要时增加分辨率。
8. KAN vs MLP:场景选择与未来展望
经过一段时间的实践,我对KAN和MLP的适用场景有了更深的体会:
选择KAN当:
- 你的问题本质是高精度函数逼近,例如求解偏微分方程、科学计算、符号回归。
- 模型的可解释性至关重要,你需要知道模型是如何做出预测的,甚至从中发现新的知识(如物理定律)。
- 数据维度相对较低,但函数关系非常复杂、非线性。
- 你希望用更少的参数获得比MLP更好的精度。
坚持使用MLP当:
- 处理非常高维的数据(如图像、文本嵌入),MLP和CNN/Transformer的组合在特征提取方面工程化程度极高。
- 任务对极致推理速度有要求,MLP的GPU矩阵乘法已被极度优化。
- 你有一个现成的、在特定领域调校得非常出色的MLP流程,更换架构成本过高。
- 问题更依赖于大规模的表示学习而非精确的函数映射。
我个人认为,KAN最有前景的方向不在于全面替代深度学习,而在于开辟一个新的赛道:科学智能。它将神经网络从一个纯粹的黑箱预测工具,转变为一个可以与人类先验知识互动、可以被检查和理解的“白箱”或“灰箱”模型。未来,我们可能会看到更多KAN与符号计算、因果发现结合的混合模型。对于研究者和我这样的工程实践者来说,现在正是深入理解并尝试将KAN应用于特定领域(如计算金融、计算生物学、工程仿真)的好时机,它很可能为我们解决那些需要精度和可解释性双重保障的难题,提供一把新的钥匙。