pykan 实战:用 KAN 从二维哈密顿流场中无监督学习守恒律(Conservation Laws)
2026/9/14 18:38:11 网站建设 项目流程

pykan 实战:用 KAN 从二维哈密顿流场中无监督学习守恒律(Conservation Laws)

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

导读:本指南完整复现 pykan 仓库中 Physics 2B: Conservation Laws 的经典实验——在不提供任何能量标签的前提下,仅利用"守恒量的梯度场与相空间流场处处正交"这一物理先验,让 KAN 从二维简谐振子系统的流场数据中自行发现哈密顿量 H = 1/2(x² + p²)。读完本文,你将掌握无监督物理损失(正交损失 + 正则项)的构造方法、MultKAN 乘法节点的宽度语法、LBFGS 强 Wolfe 线搜索训练技巧,以及prune+auto_symbolic+symbolic_formula全链路符号提取与人工符号修正的完整流程。文中所有代码与输出均来自 docs/Physics/Physics_2B_conservation_law_2D.ipynb,并辅以仓库源码级解析。

一、任务背景:守恒律与"梯度正交"的物理先验

在哈密顿力学中,若存在一个守恒量 H(x, p)(例如能量),那么系统的相流(phase flow)必然沿着 H 的等值面演化,也就是说:守恒量的梯度场与相空间流场处处正交

数学上,对于二维简谐振子系统 H = 1/2(x² + p²),其哈密顿方程给出的流场为:

dx1/dt = x2, dx2/dt = -x1, dx3/dt = x4, dx4/dt = -x3

而梯度 ∇H = (x1, x2, x3, x4),可以验证 ∇H · flow = x1·x2 + x2·(−x1) + x3·x4 + x4·(−x3) = 0,严格正交。

因此,不必给出 H 的标签,只需约束"模型输出的某个标量场 φ(x) 的梯度 ∇φ 与给定流场 flow 正交",φ 的等高线就必然与流线一致,φ 即自动成为该系统的守恒量。这正是本实验无监督学习的核心思想。

二、问题设定:构造相空间样本与归一化流场

实验从均匀分布在超立方体 [−1, 1]⁴ 上的 1000 个四维样本出发,将输入视为 (x₁, x₂, x₃, x₄),其中 (x₁, x₂) 与 (x₃, x₄) 分别是两组共轭坐标-动量对,并按哈密顿方程构造流场向量,最后对每个样本的流场向量做 L2 归一化:

from kan import * from kan.utils import batch_jacobian, create_dataset_from_data import numpy as np torch.use_deterministic_algorithms(True) # model = KAN(width=[4,[0,2],1], seed=0, base_fun='identity') # model = KAN(width=[4,[0,2],1], seed=2, base_fun='identity') model = KAN(width=[4,[0,2],1], seed=12, base_fun='identity') # the model learns the Hamiltonian H = 1/2 * (x**2 + p**2) x = torch.rand(1000,4) * 2 - 1 flow = torch.cat([x[:,[1]], -x[:,[0]], x[:,[3]], -x[:,[2]]], dim=1) flow = flow/torch.linalg.norm(flow, dim=1, keepdim=True) loss_fn = lambda v1, v2: torch.mean(torch.sum(v1 * v2, dim=1)**2)

这里torch.use_deterministic_algorithms(True)保证了随机初始化与训练过程的可复现性。宽度语法width=[4,[0,2],1]是 MultKAN 独有的乘法节点写法:第 0 层为 4 个输入(加法节点),中间层配置为[0, 2]表示 0 个加法节点 + 2 个乘法节点,输出层为 1 个节点——这意味着网络可以在内部直接学习输入之间的乘法组合,这对恢复 H ∝ x² + p² 这样的平方和结构至关重要。从源码看,该参数对应 MultKAN.initwidth的双层列表约定:[[n_0,m_0=0], [n_1,m_1], ..., [n_{L-1},m_{L-1}]],每层分别给出加法/乘法节点数量。

base_fun='identity'将残差基函数从默认的 SiLU 改为恒等映射(默认值为'silu',见 MultKAN.init),使网络输出完全由可训练样条决定,有利于后续符号化。

三、无监督损失设计:正交损失 + 稀疏正则

3.1 核心损失 cq_loss:梯度与流场正交

def get_grad_normalized(model, x): grad = batch_jacobian(model, x, create_graph=True) grad_normalized = grad/torch.linalg.norm(grad, dim=1, keepdim=True) return grad_normalized

batch_jacobian是仓库在 kan/utils.py 中提供的便捷工具,其实现是:先把批量输出求和为标量(_func_sum),再调用torch.autograd.functional.jacobian(..., create_graph=...)得到批量雅可比。create_graph=True保证求得的梯度仍然保留计算图,从而允许对梯度继续求导(二阶信息参与反向传播),这是物理损失训练的标准要求。

损失函数loss_fn(v1, v2) = mean( (Σ v1·v2)² )对每个样本计算归一化梯度与归一化流场的内积并平方取均值:内积为 0 时损失为 0,即两者方向垂直;平方项确保损失非负且对方向符号不敏感。这就是"正交约束"的量化表达。

3.2 正则项 reg_loss:稀疏化

reg_loss = model.reg(lamb_l1=1., entropy_offset=1e-4, lamb_coef=1.)

model.reg(...)是本实验训练循环中使用的正则接口,用于在约束正交性的同时推动网络稀疏化。从底层实现看,MultKAN.reg 同时包含三类惩罚:按激活幅值统计的L1 惩罚lamb_l1)、按行列分布计算的熵惩罚entropy_row + entropy_col,内部以+1e-4平滑防止除零/对数零)、以及针对样条系数的系数 L1 与系数差分 L1lamb_coef/lamb_coefdiff)。稀疏化直接决定了后续剪枝与符号提取的成败——只有绝大多数边被正则压到接近 0,auto_symbolic才能把它们识别为无效连接并固定为 0。

3.3 训练循环与优化器

def closure(): global cq_loss, reg_loss optimizer.zero_grad() grads = [] grad = get_grad_normalized(model, x) cq_loss = loss_fn(grad, flow) reg_loss = model.reg(lamb_l1=1., entropy_offset=1e-4, lamb_coef=1.) lamb = 1e-2 objective = cq_loss + lamb * reg_loss objective.backward() return objective steps = 50 log = 1 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) # optimizer = torch.optim.Adam(params, lr=1e-2) pbar = tqdm(range(steps), desc='description', ncols=100) for _ in pbar: # update grid if _ < 5 and _ % 20 == 0: model.update_grid_from_samples(x) optimizer.step(closure) if _ % log == 0: pbar.set_description("| cq_loss: %.2e | reg_loss: %.2e |" % (cq_loss.cpu().detach().numpy(), reg_loss.cpu().detach().numpy()))

要点解读:

  • 总目标cq_loss + 1e-2 * reg_loss:正交损失主导学习方向,正则项作为轻微稀疏化约束(lamb = 1e-2),避免稀疏压力过强而破坏拟合。
  • 优化器:使用 pykan 自带的 LBFGS(继承自torch.optim.Optimizer,并实现了 strong Wolfe 条件的线搜索),配置lr=1history_size=10line_search_fn="strong_wolfe",并将各容差压到1e-32,强制进行充分精确的线搜索,这是样条类模型收敛的关键;文档同时注释了可选的 Adam 方案(lr=1e-2)作为对比。
  • 网格更新:训练前 5 步内调用model.update_grid_from_samples(x)(见 MultKAN.update_grid_from_samples),让 B 样条网格依据真实样本分布自适应调整,之后保持网格固定。

实际运行 50 步后的收敛输出为(示例输出):

checkpoint directory created: ./model saving model version 0.0 | cq_loss: 1.57e-03 | reg_loss: 1.01e+01 |: 100%|███████████████████| 50/50 [00:30<00:00, 1.63it/s]

正交损失被压到 1.57e-03 量级,说明模型的梯度场已基本与流场垂直;而正则项仍维持在 1.01e+01,符合"先拟合、后稀疏"的常规路径。训练过程中auto_save=True(默认开启,见 MultKAN.init)使模型在修改时自动落盘 checkpoint,输出中的saving model version x.y即由此而来。

四、可视化与结构剪枝

训练结束后直接调用model.plot()输出网络结构图(见上文第一张图),随后进入剪枝与符号化阶段:

# model = KAN(width=[4,[0,2],1], seed=12, base_fun='identity') model.plot() model = model.prune(edge_th=5e-2) model.auto_symbolic()

MultKAN.prune 先执行节点剪枝(node_th=1e-2默认阈值)再执行边剪枝(此处显式传入edge_th=5e-2),其依据是前向激活统计出的节点/边归因分数:归因分数低于阈值的单元被判定为"死亡"并置零,从而得到稀疏结构。剪枝后调用auto_symbolic对所有存活边做符号回归。

五、自动符号回归:从数值网络到显式公式

5.1 auto_symbolic 的逐边拟合日志

model.auto_symbolic()在 MultKAN.auto_symbolic 中实现:对每个(层 l, 输入 i, 输出 j)边,若归因分数为 0 则直接固定为常数 0,否则在符号库中搜索最优函数(综合 r² 与复杂度,weight_simple=0.8默认偏向简单函数),r² 达标才执行fix_symbolic。seed=12 实验的输出为:

saving model version 0.1 fixing (0,0,0) with 0 fixing (0,0,1) with 0 fixing (0,0,2) with 0 fixing (0,0,3) with 0 fixing (0,1,0) with 0 fixing (0,1,1) with 0 fixing (0,1,2) with 0 fixing (0,1,3) with 0 fixing (0,2,0) with 0 fixing (0,2,1) with 0 fixing (0,2,2) with x, r2=0.9983036518096924, c=1 fixing (0,2,3) with x, r2=0.9988861680030823, c=1 fixing (0,3,0) with x, r2=0.9961345195770264, c=1 fixing (0,3,1) with x, r2=0.9859936237335205, c=1 fixing (0,3,2) with 0 fixing (0,3,3) with 0 fixing (1,0,0) with x, r2=0.9999908804893494, c=1 fixing (1,1,0) with x, r2=0.9999944567680359, c=1 saving model version 0.2

日志中(l,i,j)三元组标注边位置,c为该符号函数的复杂度(此处均为 1,即最简单的线性函数)。可以看到:第一层多数输入边被判定为 0(无效连接),仅 (0,2,·)、(0,3,·) 保留线性映射;两个乘法节点的输出边则被拟合为 r² 接近 1 的线性函数 x。

5.2 用 symbolic_formula + ex_round 提取显式公式

from kan.utils import ex_round from sympy import * ex_round(expand(ex_round(model.symbolic_formula()[0][0],5)),3)

得到(数学表达式):

- 0.011 x_3² - 0.01 x_4² + 0.001 x_4 + 0.002

symbolic_formula(MultKAN.symbolic_formula)把每条已符号化的边组合为 sympy 表达式并逐层展开;ex_round(kan/utils.py)通过sympy.preorder_traversal遍历表达式树,将所有浮点常数四舍五入到指定位数,便于人类阅读。注意 sympy 变量下标从 1 开始,因此 x₃、x₄ 对应输入的第 3、4 列。该结果近似恢复了"平方和"形式的守恒量(缺少 x₁、x₂ 部分),说明无监督正交约束已成功捕获系统的部分能量结构。

六、人工符号修正:unfix_symbolic 与 fix_symbolic

自动符号化的结果不一定完美。文档随后演示了人工介入修正流程:当某条边被错误拟合为 exp 等复杂函数时,先解除符号化,再手动固定为期望的简单函数:

model = model.prune(edge_th=5e-2) model.auto_symbolic()

seed=0 变体的 auto_symbolic 日志显示 (1,0,0) 边被拟合为exp, r2=1.000000238418579, c=2(复杂度 2,指数函数)。虽然 r² 极高,但物理上不希望出现指数项,于是:

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

输出:

saving model version 0.3 Best value at boundary. r2 is 0.9992757439613342 saving model version 0.4 tensor(0.9993)

unfix_symbolic(MultKAN.unfix_symbolic)解除该边与符号函数的绑定,恢复为可训练样条;随后fix_symbolic(1,0,0,'x')(MultKAN.fix_symbolic)在a_range=(-10,10)b_range=(-10,10)内搜索最优仿射参数 a·x+b,最终以 r²=0.9993 的拟合优度将其固定为线性函数 x,复杂度从 2 降回 1。重新提取公式得到:

- 0.011 x_1² - 0.01 x_2² - 0.006

这次恢复的是 (x₁, x₂) 坐标子空间的平方和结构——与 seed=12 的结果互补,二者分别捕捉了不同共轭对的能量。

七、随机初始化敏感性:三次运行对比

文档通过注释行保留了三种随机种子(seed=0、seed=2、seed=12)的对照实验,完整演示了随机初始化对符号发现结果的影响:

随机种子auto_symbolic 亮点人工修正最终符号公式
seed=12(主实验)(0,2,2)/(0,2,3)/(0,3,0)/(0,3,1) 拟合为 x无需-0.011 x₃² - 0.01 x₄² + 0.001 x₄ + 0.002
seed=0(1,0,0) 拟合为 exp(r²≈1.0,复杂度 2)unfix_symbolic(1,0,0)fix_symbolic(1,0,0,'x'),r²=0.9993-0.011 x₁² - 0.01 x₂² - 0.006
seed=2(1,1,0) 拟合为 exp(复杂度 2)unfix_symbolic(1,1,0)fix_symbolic(1,1,0,'x'),r²=0.9832-0.003 x₁x₄ + 0.0031 x₂x₃ - 0.0819

seed=2 的完整流程为:

# model = KAN(width=[4,[0,2],1], seed=2, base_fun='identity') model.plot() model = model.prune() # 使用默认阈值 node_th=1e-2, edge_th=3e-2 model.auto_symbolic()

其日志中 (1,1,0) 边被拟合为exp, r2=1.0000001192092896, c=2,同样通过解绑再固定为 x 完成修正(r²=0.9832),最终符号公式为-0.003 x₁x₄ + 0.0031 x₂x₃ - 0.0819——一个包含交叉项的解。这一对比说明:正交约束下恢复守恒量本身是稳健的,但符号提取阶段收敛到的具体代数形式(平方和 / 交叉项 / 坐标子集)会随初始化不同而变化,因此"剪枝 → 符号化 → 人工审视与修正"的闭环流程是物理发现工作流中不可省略的一环。

八、方法总结与延伸

本实验展示了一条完整的"数据驱动的守恒律发现"技术路线,其关键环节可归纳为:

  1. 物理先验编码:用梯度-流场正交损失(loss_fn = mean((Σ v₁·v₂)²))替代有监督标签,配合batch_jacobian(model, x, create_graph=True)实现二阶可导的梯度计算;
  2. 稀疏正则model.reg(lamb_l1=1., entropy_offset=1e-4, lamb_coef=1.)融合 L1、熵与系数惩罚,为后续剪枝铺路;
  3. 二阶优化:pykan 自定义 LBFGS(strong Wolfe 线搜索)配合训练初期update_grid_from_samples自适应网格;
  4. 符号闭环prune(edge_th=...)依据归因分数裁剪 →auto_symbolic()逐边符号回归 →symbolic_formula()+ex_round提取人类可读公式 →unfix_symbolic/fix_symbolic人工修正异常拟合。

值得延伸的是:文档开头还引入了create_dataset_from_data(kan/utils.py)——它按train_ratio=0.8将原始数据划分为训练/测试集并封装为字典,当后续希望从轨迹采样数据(而非解析流场)学习守恒量时,该工具可将任意观测数据无缝接入上述无监督训练流程,从而把本方法推广到"仅有观测轨迹、未知运动方程"的真实物理场景。

相关资源:

  • 本文档源文件:docs/Physics/Physics_2B_conservation_law_2D.rst
  • 可执行 Notebook:docs/Physics/Physics_2B_conservation_law_2D.ipynb(另有 tutorials/Physics/Physics_2B_conservation_law_2D.ipynb 副本)
  • 核心实现:kan/MultKAN.py、kan/utils.py、kan/LBFGS.py
  • 姊妹篇(一维/标量守恒律):docs/Physics/Physics_2A_conservation_law.rst

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

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

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

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

立即咨询