使用 pykan 求解二维 Poisson 方程并进行符号化解释:KAN 偏微分方程数值求解实战
2026/9/14 18:50:18 网站建设 项目流程

使用 pykan 求解二维 Poisson 方程并进行符号化解释:KAN 偏微分方程数值求解实战

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

导读

本文基于 pykan 仓库的官方示例(docs/Example/Example_6_PDE_interpretation.ipynb 及其 RST 版本 docs/Example/Example_6_PDE_interpretation.rst),完整演示如何用 KAN(Kolmogorov-Arnold Networks)求解带 Dirichlet 边界条件的二维 Poisson 方程:通过torch.autograd构造拉普拉斯算子残差损失与边界损失,使用 L-BFGS 优化器训练网络,再通过fix_symbolic将激活函数替换为线性函数与正弦函数,最终借助symbolic_formula输出闭合解析表达式,实现"数值求解 + 符号回归"一体化的 PDE 可解释求解流程。读完本文,你将掌握 KAN 求解 PDE 的完整代码骨架、自动微分求二阶导数的技巧,以及把训练好的网络翻译成数学公式的标准步骤。

一、问题设定:二维 Poisson 方程与 Dirichlet 边界条件

本示例求解的偏微分方程为:

$$ \nabla^2 f(x,y) = -2\pi^2\sin(\pi x)\sin(\pi y) $$

定义域为 $x, y \in [-1, 1]$,边界条件为 $f(-1,y)=f(1,y)=f(x,-1)=f(x,1)=0$,对应的解析真解为:

$$ f(x,y)=\sin(\pi x)\sin(\pi y) $$

该设定十分巧妙:真解本身是两个一元正弦函数的乘积,而 KAN 的核心假设(Kolmogorov-Arnold 表示定理)恰好表明多元函数可以表示为一元函数的有序叠加与复合。因此,用 KAN 求解该方程并在训练后把网络翻译回 $\sin(\pi x)\sin(\pi y)$,是验证"KAN 天然适合符号化解释 PDE 解"这一命题的典型实验。

二、环境准备与模型初始化

示例代码首先导入依赖并创建 KAN 模型:

from kan import * import matplotlib.pyplot as plt from torch import autograd from tqdm import tqdm device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(device) dim = 2 np_i = 21 # number of interior points (along each dimension) np_b = 21 # number of boundary points (along each dimension) ranges = [-1, 1] model = KAN(width=[2,2,1], grid=5, k=3, seed=1, device=device)

关键点说明:

  • from kan import *会从 kan/init.py 导入MultKAN(及KAN别名)和utils中的全部工具函数;
  • width=[2,2,1]表示网络结构为 2 个输入($x, y$)→ 2 个隐藏神经元 → 1 个输出($f(x,y)$);
  • grid=5表示每条样条激活函数的初始网格区间数为 5,k=3表示使用 3 阶(三次)B 样条。网格与阶数相关实现见 kan/KANLayer.py 中的KANLayer.__init__
  • seed=1固定随机种子保证实验可复现;
  • 设备优先使用 CUDA,无 GPU 时回退到 CPU。

三、通过自动微分构造 PDE 残差损失

KAN 作为可微网络,其输出对输入的导数可由torch.autograd直接获得。示例在 Notebook 内自定义了batch_jacobian,用于对批量输入逐样本计算 Jacobian:

def batch_jacobian(func, x, create_graph=False): # x in shape (Batch, Length) def _func_sum(x): return func(x).sum(dim=0) return autograd.functional.jacobian(_func_sum, x, create_graph=create_graph).permute(1,0,2)

该函数的思想是:先把批量输出按样本维求和,再对输入求 Jacobian,得到形状为(Batch, Length, Length)的张量,其中每个样本对应一个完整的 Jacobian 矩阵。这与仓库中 kan/utils.py 提供的batch_jacobianmode='vector'分支)实现完全一致,说明该写法是 pykan 的标准做法。

随后定义真解与源项(方程右端项):

# define solution sol_fun = lambda x: torch.sin(torch.pi*x[:,[0]])*torch.sin(torch.pi*x[:,[1]]) source_fun = lambda x: -2*torch.pi**2 * torch.sin(torch.pi*x[:,[0]])*torch.sin(torch.pi*x[:,[1]])

3.1 内部采样点

内部点在 $[-1,1]^2$ 上采样,支持两种模式:

sampling_mode = 'random' # 'random' or 'mesh' x_mesh = torch.linspace(ranges[0],ranges[1],steps=np_i) y_mesh = torch.linspace(ranges[0],ranges[1],steps=np_i) X, Y = torch.meshgrid(x_mesh, y_mesh, indexing="ij") if sampling_mode == 'mesh': #mesh x_i = torch.stack([X.reshape(-1,), Y.reshape(-1,)]).permute(1,0) else: #random x_i = torch.rand((np_i**2,2))*2-1 x_i = x_i.to(device)
  • mesh模式:在 $21\times21=441$ 个均匀网格点上求值;
  • random模式:在 $[-1,1]^2$ 内随机采样 441 个点(注意示例注释中'radnom'为原文笔误,实际判断逻辑为'mesh'之外一律走随机分支)。

3.2 边界采样点

边界点取自四条边($x=-1$、$x=1$、$y=-1$、$y=1$):

# boundary, 4 sides helper = lambda X, Y: torch.stack([X.reshape(-1,), Y.reshape(-1,)]).permute(1,0) xb1 = helper(X[0], Y[0]) xb2 = helper(X[-1], Y[0]) xb3 = helper(X[:,0], Y[:,0]) xb4 = helper(X[:,0], Y[:,-1]) x_b = torch.cat([xb1, xb2, xb3, xb4], dim=0) x_b = x_b.to(device)

这里X[0]Y[0]对应 $x=-1$ 边,X[-1]对应 $x=1$ 边,X[:,0]Y[:,-1]对应 $y$ 方向的两条边,四条边各 21 个点共 84 个边界点。

四、训练循环:L-BFGS + 自定义 closure

训练采用二阶优化器 L-BFGS,并启用 strong Wolfe 线搜索:

steps = 20 alpha = 0.01 log = 1 def train(): optimizer = LBFGS(model.parameters(), lr=1, history_size=10, line_search_fn="strong_wolfe", tolerance_grad=1e-32, tolerance_change=1e-32, tolerance_ys=1e-32) pbar = tqdm(range(steps), desc='description', ncols=100) for _ in pbar: def closure(): global pde_loss, bc_loss optimizer.zero_grad() # interior loss sol = sol_fun(x_i) sol_D1_fun = lambda x: batch_jacobian(model, x, create_graph=True)[:,0,:] sol_D1 = sol_D1_fun(x_i) sol_D2 = batch_jacobian(sol_D1_fun, x_i, create_graph=True)[:,:,:] lap = torch.sum(torch.diagonal(sol_D2, dim1=1, dim2=2), dim=1, keepdim=True) source = source_fun(x_i) pde_loss = torch.mean((lap - source)**2) # boundary loss bc_true = sol_fun(x_b) bc_pred = model(x_b) bc_loss = torch.mean((bc_pred-bc_true)**2) loss = alpha * pde_loss + bc_loss loss.backward() return loss if _ % 5 == 0 and _ < 50: model.update_grid_from_samples(x_i) optimizer.step(closure) sol = sol_fun(x_i) loss = alpha * pde_loss + bc_loss l2 = torch.mean((model(x_i) - sol)**2) if _ % log == 0: pbar.set_description("pde loss: %.2e | bc loss: %.2e | l2: %.2e " % (pde_loss.cpu().detach().numpy(), bc_loss.cpu().detach().numpy(), l2.cpu().detach().numpy())) train()

该训练循环包含几个值得深入理解的技术点:

1. 二阶导数的递推构造。先对model求一阶 Jacobian 得到梯度场 $(\partial f/\partial x, \partial f/\partial y)$,再对这个梯度场函数继续求 Jacobian 得到 Hessian 矩阵;取 Hessian 的对角线元素并求和,即得到拉普拉斯算子 $\nabla^2 f = f_{xx} + f_{yy}$:

lap = torch.sum(torch.diagonal(sol_D2, dim1=1, dim2=2), dim=1, keepdim=True)

注意create_graph=True必须保留,否则第二次求导无法对第一次求导结果继续反向传播。

2. 两项损失的加权组合。内部点损失约束方程残差(pde_loss = mean((lap - source)^2)),边界点损失约束边界条件(bc_loss = mean((bc_pred - bc_true)^2)),总损失为loss = alpha * pde_loss + bc_loss,其中alpha = 0.01用于平衡两项尺度差异。

3. 网格自适应更新。每 5 步调用一次model.update_grid_from_samples(x_i),让样条网格根据当前输入样本的分布自动重排,提升样条逼近精度。其底层实现在 kan/MultKAN.py:先对样本做一次前向得到各层激活self.acts,再逐层调用 kan/KANLayer.py 的update_grid_from_samples重排网格节点。

4. L-BFGS 需要 closure。优化器在每次optimizer.step(closure)时多次调用 closure 以进行线搜索,因此 closure 内必须完成"清零梯度 → 计算损失 → 反向传播 → 返回损失"的完整流程。该 L-BFGS 实现位于 kan/LBFGS.py,内部实现了 strong Wolfe 条件线搜索(_strong_wolfe),并支持history_sizetolerance_gradtolerance_changetolerance_ys等参数。

运行 20 步后的典型输出(原文档记录,在 CUDA 环境下):

cuda checkpoint directory created: ./model saving model version 0.0 pde loss: 2.23e+00 | bc loss: 5.99e-03 | l2: 3.78e-03 : 100%|███████| 20/20 [00:22<00:00, 1.11s/it]

训练日志同时显示,模型首次保存检查点到./model目录(saving model version 0.0),这是 pykan 自动保存机制(auto_save=True)在起作用,每个版本对应 model/ 目录中的0.0_config.yml0.0_state0.0_cache_data文件。

五、可视化训练结果

训练完成后直接调用model.plot()可视化网络结构:

model.plot(beta=10)

该图展示了训练后 KAN 的各层激活函数形状(样条曲线),beta=10控制激活函数曲线的颜色映射与线条粗细。plot的完整签名与参数说明见 kan/MultKAN.py。

六、符号化解释:fix_symbolic 与 symbolic_formula

这是本示例最核心的亮点:把数值网络"翻译"成解析公式。

6.1 将激活函数固定为符号函数

由于真解 $\sin(\pi x)\sin(\pi y)$ 是一元正弦函数与线性函数的组合,示例将第一层 4 个激活函数全部固定为线性函数'x',第二层(输出层)固定为正弦函数'sin'(示例代码中注释说明该步对超参数较敏感):

model.fix_symbolic(0,0,0,'x') model.fix_symbolic(0,0,1,'x') model.fix_symbolic(0,1,0,'x') model.fix_symbolic(0,1,1,'x')

fix_symbolic的完整签名与参数语义见 kan/MultKAN.py:

参数含义默认值
l层索引
i输入神经元索引
j输出神经元索引
fun_name符号函数名(如'x''sin''cos''exp'等)
fit_params_bool是否通过拟合确定仿射参数a, b, c, dTrue
a_range/b_range仿射参数ab的扫描范围(-10, 10)
verbose是否打印拟合信息True
random是否随机初始化仿射参数False
log_history是否记录历史True

调用时,如果fit_params_bool=True,会取该激活的输入样本x与样条输出y(即self.acts[l][:, i]self.spline_postacts[l][:, j, i]),通过 kan/utils.py 的fit_paramsa_range/b_range内网格扫描拟合最优仿射参数,并返回拟合优度r2。底层实现见 kan/Symbolic_KANLayer.py 的fix_symbolic

原文档记录的四次固定操作输出如下:

r2 is 0.8357976675033569 r2 is not very high, please double check if you are choosing the correct symbolic function. saving model version 0.1 r2 is 0.8300805687904358 r2 is not very high, please double check if you are choosing the correct symbolic function. saving model version 0.2 r2 is 0.8376883268356323 r2 is not very high, please double check if you are choosing the correct symbolic function. saving model version 0.3 r2 is 0.8372848629951477 r2 is not very high, please double check if you are choosing the correct symbolic function. saving model version 0.4

可以看到单个激活替换后的r2约为 0.83~0.84,并不算高——因为此时仿射参数尚未经过联合训练精调,且每次替换后都会保存一个新版本检查点(0.10.4)。原文档随后输出tensor(0.8373),对应符号化后的整体拟合优度。这解释了文档中"quite sensitive to hyperparams"的告诫:符号化环节需要后续训练配合才能收敛到机器精度。

6.2 符号化后继续训练达到机器精度

所有激活变为符号函数后,仿射参数仍然是可训练的,因此继续调用train()精调这些参数:

train()

原文档记录此时前 10 步日志为pde loss: 1.71e+01 | bc loss: 1.14e-02 | l2: 1.37e-01,并展示了一段KeyboardInterrupt的 Traceback(涉及 kan/LBFGS.py 中_strong_wolfe_directional_evaluateclosure的调用链)。这段 Traceback 并非错误,而是作者手动中断了训练过程——它恰好揭示了 L-BFGS 在step内部通过_directional_evaluate反复调用用户 closure 进行强 Wolfe 线搜索的执行路径。原文档指出,充分训练后模型可以达到机器精度(machine precision),即符号化后的 KAN 能精确复现真解。

6.3 输出闭合解析公式

最后打印符号化后的公式:

formula = model.symbolic_formula()[0][0] ex_round(formula,6)

symbolic_formula的实现见 kan/MultKAN.py:它遍历每一层的符号激活函数与仿射参数 $(a, b, c, d)$,用 sympy 表达式逐层组装出完整公式;ex_round则将表达式中的浮点数统一四舍五入到指定位数。原文档最终得到的公式为:

$$ \displaystyle - 0.5 \sin{\left(3.141592 x_{1} + 3.141593 x_{2} - 4.712389 \right)} + 0.5 \sin{\left(3.141593 x_{1} - 3.141592 x_{2} + 1.570797 \right)} $$

利用三角恒等式 $\sin(A)-\sin(B)=2\cos\frac{A+B}{2}\sin\frac{A-B}{2}$ 可以化简:令 $A = \pi x_1 + \pi x_2 - \frac{3\pi}{2}$,$B = \pi x_1 - \pi x_2 + \frac{\pi}{2}$,则 $\frac{A+B}{2} = \pi x_1 - \frac{\pi}{2}$,$\frac{A-B}{2} = \pi x_2 - \pi$,于是 $-0.5\sin A + 0.5\sin B = \cos(\pi x_1 - \frac{\pi}{2})\sin(\pi x_2 - \pi) = \sin(\pi x_1)\sin(\pi x_2)$。也就是说,KAN 通过符号回归精确恢复了真解 $f(x,y)=\sin(\pi x)\sin(\pi y)$——这正是本示例"PDE 解释"(interpretation)的含义所在。

七、完整流程总结与实验要点

阶段关键操作对应源码/文档位置
建模KAN(width=[2,2,1], grid=5, k=3, seed=1)kan/MultKAN.py
采样mesh / random 内部点 + 四边边界点docs/Example/Example_6_PDE_interpretation.rst
损失自动微分求拉普拉斯 + 边界 MSE 加权batch_jacobian,与 kan/utils.py 一致
优化L-BFGS + strong Wolfe 线搜索,closure 模式kan/LBFGS.py
网格每 5 步update_grid_from_sampleskan/MultKAN.py
符号化fix_symbolic替换激活为'x'/'sin'kan/MultKAN.py
解释symbolic_formula+ex_round输出公式kan/MultKAN.py

实操要点回顾:

  1. create_graph=True不能省略,否则无法对二阶导数继续反传;
  2. alpha平衡权重:PDE 残差与边界条件的量纲不同,alpha=0.01在本文设定下效果良好,实际问题中需按损失尺度调整;
  3. 符号化后必须继续训练fix_symbolic只是给出仿射参数初值,只有联合精调才能逼近机器精度;
  4. 检查点自动保存:每次fix_symbolic会以新版本号保存模型,可通过model.checkout(version)回溯历史版本;
  5. 公式化简验证symbolic_formula输出的表达式可能包含冗余项(如本例的相位偏移),可结合 sympy 化简并与真解对照,验证 KAN 是否学到了真实的物理规律。

对于更复杂的 PDE 或高维问题,可以复用本文的"自动微分残差 + L-BFGS + 符号化"三件套,仅需替换source_fun、边界条件与采样点生成逻辑;若希望自动挑选符号函数,还可参考model.auto_symbolic()model.suggest_symbolic()(同样位于 kan/MultKAN.py),它们基于r2与复杂度打分在预置函数库中自动搜索最佳符号候选,将"人工指定符号"升级为"自动符号发现"。

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

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

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

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

立即咨询