KAN 柯尔莫哥洛夫-阿诺德网络

FreeGuideOnline 5阅读 2026-07-13

python import torch import numpy as np import matplotlib.pyplot as plt from kan import KAN

1. 创建训练数据

torch.manual_seed(42) n_samples = 1000 x = torch.linspace(-1, 1, n_samples).unsqueeze(1) # shape (1000,1) y = torch.linspace(-1, 1, n_samples).unsqueeze(1) X = torch.cat([x, y], dim=1) # 输入两维

目标函数:exp(sin(πx) + y^2)

target = torch.exp(torch.sin(np.pi * X[:,0]) + X[:,1]**2).unsqueeze(1)

2. 定义 KAN 模型

网络形状 [2, 5, 1]:2个输入,5个隐藏神经元,1个输出

grid=10 表示 B 样条初始网格数,k=3 表示三次样条

model = KAN(width=[2, 5, 1], grid=10, k=3, seed=42)

3. 训练模型(使用 LBFGS 优化器)

results = model.fit( {'train_input': X, 'train_label': target}, opt='LBFGS', steps=50, loss_fn='mse' )

4. 查看训练损失下降

plt.plot(results['train_loss']) plt.yscale('log') plt.title('Training Loss Curve') plt.xlabel('Step') plt.ylabel('MSE') plt.show()

5. 可视化学到的函数结构

model.plot()

绘制隐藏层中的激活函数(边函数)

model.fix_symbolic(0,0,0,'sin') # 可尝试固定符号,此处仅为示例 plt.show()

6. 测试并计算最终损失

with torch.no_grad(): pred = model(X) final_mse = torch.mean((pred - target) ** 2) print(f"Final MSE: {final_mse.item():.6f}")


**运行观察**:损失会迅速下降到极低水平(例如 \(10^{-5}\) 以下)。调用 `model.plot()` 会显示网络拓扑图,边的透明度代表其重要性(经过稀疏正则化后许多边接近零)。最重要的部分是,你可以单独查看每条边的函数形状,比如输入 x 到第一个隐藏神经元可能学到一个类似正弦的曲线。

## 7. 解读 KAN 的输出:从黑箱到透明公式

KAN 最具革命性的功能是**符号化 (symbolification)**。训练完成后,可以要求模型将学到的样条函数拟合为已知的符号函数(如 sin, cos, exp, x^2 等),并输出一个解析表达式。

```python
# 自动符号化(需先安装 sympy)
lib = ['x','x^2','x^3','x^4','exp','log','sqrt','tanh','sin','cos','abs']
model.auto_symbolic(lib=lib)
formula = model.symbolic_formula()[0][0]
print("Learned formula:", formula)

对于上面的例子,你可能会得到类似:

exp(sin(3.1415*x1) + x2^2)

KAN 自动发现了 π、sin 和 exp 的组合关系。这为科学研究提供了强大的假设生成工具。

7.1 边函数的可视化

即使不进行完全符号化,你也可以轻松绘制每条边的函数形状:

# 绘制 (0,0,0) 号边函数,即第0层第0个输入到第0个输出
model.plot_curve(0,0,0, num_pts=200)
plt.title('Edge function φ_{0,0} from x1')
plt.show()