)
pykan 教程利用 KAN Compiler 先验加速本构律 P12 的符号发现Physics 4B 实战【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan本文基于 pykan 仓库的官方教程 Physics_4B_constitutive_laws_P12_with_prior.ipynb及其 RST 版本 Physics_4B_constitutive_laws_P12_with_prior.rst讲解如何把已知的线性弹性本构关系 P12 μ(F12 F21) 作为符号先验通过kanpilerKAN Compiler编译成一个 MultKAN 初始网络再经宽度/深度扩展、扰动、训练、剪枝与自动符号化最终从数值数据中发现超弹性Neo-Hookean本构关系的闭合表达式。读完本文你将掌握使用kan.compiler.kanpiler从 SymPy 表达式构建 MultKAN 骨架、用expand_depth/expand_width/perturb改造先验网络、用create_dataset_from_data构造数据集并训练以及用auto_symbolicsymbolic_formula输出可读公式的完整工作流。1. 问题背景从线弹性先验出发识别超弹性本构律在连续介质力学中第一 Piola–Kirchhoff 应力张量P是变形梯度F3×3的函数。不同材料模型给出不同的解析关系线弹性Linear ElasticP11 2·μ·(F11 − 1) λ·(F11 F22 F33 − 3)P12 μ·(F12 F21)Neo-Hookean 超弹性P11 μ·(F11² F21² F31² − 1) λ·ln(|F|)P12 μ·(F12·F11 F22·F21 F32·F31)教程的目标是在**已知线性本构先验**的前提下让 KAN 通过数据驱动的方式把 P12 从线性的μ(F12 F21)推广到 Neo-Hookean 的非线性形式。这里的先验体现在两个层面用解析的线性表达式初始化网络结构符号先验同时生成线弹性与 Neo-Hookean 两类数据前者用于验证先验注入的正确性后者作为学习目标。1.1 数据生成与确定性设置教程首先固定随机种子以保证结果可复现并将默认张量类型设为float64random.seed(0) np.random.seed(0) torch.manual_seed(0) torch.use_deterministic_algorithms(True) torch.set_default_dtype(torch.float64) N 1000 sigma 0.2 # 变形扰动幅度 # 在单位阵附近生成 1000 个变形梯度 F保证 det(F) 0 F torch.eye(3,3)[None,:,:].expand(N,3,3) (torch.rand(N,3,3)*2-1)*sigma det torch.det(F) F * (det 0)[:,None,None] linear LinearElasticConstitutiveLaw(young_modulus1.0, poisson_ratio0.2) P_l linear(F) P11_l P_l[:,[0],[0]] P12_l P_l[:,[0],[1]] neo NeoHookeanConstitutiveLaw(young_modulus1.0, poisson_ratio0.2) P_n neo(F) P11_n P_n[:,[0],[0]] P12_n P_n[:,[0],[1]]LinearElasticConstitutiveLaw与NeoHookeanConstitutiveLaw来自教程自带的constitutive_laws_generator模块读者需按教程环境准备该生成器。两种材料均取杨氏模量young_modulus1.0、泊松比poisson_ratio0.2由此可解析地得到 Lame 参数 μ 与 λ见 2.1 节。2. 编译先验kanpiler 如何把公式变成 MultKAN2.1 构造输入变量与先验表达式KAN Compilerkanpiler即kan.compiler.expr2kan三者互为别名接受一组 SymPy 符号与一个 SymPy 表达式将其编译为具有求和节点 乘法节点 激活边结构的 MultKANmu, lambda_ linear.get_lame_parameters() input_vars F11, F12, F13, F21, F22, F23, F31, F32, F33 symbols(F11 F12 F13 F21 F22 F23 F31 F32 F33) # 编译真值中的更多项或让网络更大 P12_l_expr mu * (F12 F21) model kanpiler(input_vars, P12_l_expr, base_funidentity)input_vars是 9 个 SymPy 符号对应变形梯度 F 的 9 个分量P12_l_expr μ(F12 F21)即线弹性先验。base_funidentity表示样条激活的基础函数使用恒等函数即默认样条仅在恒等函数之上学习残差。从 compiler.py 的源码看expr2kan的核心步骤是用next_nontrivial_operation剥离表达式中的仿射部分系数/偏置把μ(F12F21)归一化为内部表达式F12 F21仿射参数(scale, bias)记录在边上递归解析表达式树Add产生求和节点、Mul产生乘法节点及其mult_arity乘法元数Pow/sin/exp等一元函数直接映射为符号激活边将节点/子节点/连接的三元组(depth, index)换算为 KAN 的(layer, width)坐标随后调用MultKAN(width..., mult_arity..., gridgrid, kk)构造网络并通过fix_symbolic把每条边固定为对应的符号函数fit_params_boolFalse即不训练仿射参数把树节点的 scale/bias 写入node_scale/node_bias与subnode_scale/subnode_bias。因此编译得到的模型天然是一个完全符号化的网络所有激活边已被fix_symbolic固定为x、x^2、sin等符号函数这正是符号先验的本质——后续训练只需要微调这些符号边的系数与结构。编译完成后调用model.get_act(F_flatten)F_flatten F.reshape(N, -1)形状为 (1000, 9)为各条边计算激活值供plot与后续剪枝使用然后用model.plot(...)可视化F_flatten F.reshape(N, -1) model.get_act(F_flatten) model.plot(in_varsinput_vars, out_vars[r$P_{12}$], varscale0.75, scale1.0, out_vars_offset0.08)控制台输出saving model version 0.1——pykan 会为每次结构变更自动保存模型版本详见第 5 节rewind。plot的可视化参数含义in_vars/out_vars用于标注输入/输出符号支持 LaTeX如r$P_{12}$varscale控制变量标签字号缩放scale控制整图缩放out_vars_offset控制输出标签的偏移量。3. 在先验骨架上升级容量深度扩展与宽度扩展初始编译网络直接实现P12 μ(F12 F21)。为了让它有能力表达 Neo-Hookean 的非线性项如F13·F23、F12·F11、平方项等需要先扩充网络容量再交给数据训练。3.1 深度扩展expand_depthmodel.expand_depth() model.plot(in_varsinput_vars, out_vars[r$P_{12}$], varscale0.75, scale1.0, out_vars_offset0.08)从 MultKAN.py 的expand_depth的实现看该方法在网络末端追加一个恒等层新增一个KANLayer(dim_out, dim_out)mask 置零即默认不激活并在symbolic_fun中追加一个Symbolic_KANLayer其对角边固定为x、非对角边固定为0——即新增层初始不改变前向输出只是把网络变深为后续乘法/复合结构留出空间。控制台显示saving model version 0.2。3.2 宽度扩展expand_widthmodel.expand_width(1,5,sum_boolFalse,mult_arity2) # 在第 1 层新增 5 个二元乘法节点 model.expand_width(1,4) # 在第 1 层再新增 4 个加法节点 model.plot(in_varsinput_vars, out_vars[r$P_{12}$], varscale0.75, scale1.0, out_vars_offset0.08)expand_width(layer_id, n_added_nodes, sum_bool, mult_arity)用于在指定层拓宽网络sum_boolTrue新增节点为加法节点sum_boolFalse新增节点为乘法节点mult_arity乘法节点的元数参与相乘的输入数此处mult_arity2表示二元乘法。新增乘法节点正是为了让网络能够表达 Neo-Hookean 中的乘积项如F12·F11、F22·F21、F13·F23这是从线性推广到非线性的关键结构升级。两次扩展后控制台依次输出saving model version 0.3、saving model version 0.4。3.3 扰动所有边perturb(modeall)model.perturb(modeall) # 输出 saving model version 0.5perturb用于复活网络中的非激活边让训练可以探索更多结构。从 perturb 源码 看modeall会将所有边的act_fun[l].mask置为扰动幅度mag1.0mode还支持non-intrusive仅扰动与已激活边直接相连的边与minimal。由于先验编译的初始网络里大量边被fix_symbolic(...,0)固定为零必须先perturb解除约束后续训练才能长出新项。model.plot展示了扰动后全部边都具备激活能力的状态。4. 数据驱动精调从先验网络拟合 Neo-Hookean 数据4.1 构造数据集并评估先验误差使用kan.utils.create_dataset_from_data从 (F, P12_n) 数据对构造数据集默认train_ratio0.8即 800 个训练样本、200 个测试样本随机划分dataset create_dataset_from_data(F_flatten, P12_n) torch.sqrt(torch.mean((model(dataset[train_input]) - dataset[train_label])**2)) # tensor(0.2937, grad_fnSqrtBackward0)从 utils.py 的create_dataset_from_data的实现可见它把输入、标签按train_ratio无放回抽样划分为train_input/train_label/test_input/test_label四个键并统一detach()到指定设备。上式的tensor(0.2937)表示未训练前模型在训练集上的 RMSE——由于先验是线性的而目标是 Neo-Hookean误差显著这正是需要训练优化的空间。4.2 第一阶段拟合from kan.utils import create_dataset_from_data model.fit(dataset, steps100, lamb1e-5) # | train_loss: 2.00e-03 | test_loss: 1.51e-03 | reg: 4.69e00 | 100%| 100/100 # saving model version 0.6fit是 MultKAN 的核心训练接口见 fit 定义默认使用 L-BFGS 优化器。关键参数steps100优化步数lamb1e-5正则化强度作用于 L1 稀疏化项控制网络稀疏度输出中的reg即稀疏化正则损失训练后 RMSE 从 0.2937 降至约 2.0e-03train/1.5e-03test。训练后绘图可以看到各条边的激活函数开始呈现非线性形态4.3 剪枝并回退prune 与 rewind训练后网络中大量边/节点接近零贡献可用prune清理model model.prune() model.plot(in_varsinput_vars, out_vars[r$P_{12}$], varscale0.75, scale1.0, out_vars_offset0.08) # saving model version 0.7从 prune 源码 看prune依次执行prune_node(node_th)按节点归因分数剪除节点默认阈值node_th1e-2、forward与attribute()重新计算归因、prune_edge(edge_th)按边归因剪除边默认edge_th3e-2并记录日志。剪枝后网络结构显著简化plot的轴数从 167 降到 68。剪枝可能误伤有用边因此教程回退到剪枝前的 0.7 版本再做精细训练model model.rewind(0.7) # rewind to model version 0.7, renamed as 1.7rewind(model_id)见 rewind 定义从模型的版本历史中恢复指定版本。注意输出显示回退后的版本被重命名为 1.7——版本号前缀从 0 变为 1表示进入了新的训练周期。model.fit(dataset, steps100) # | train_loss: 2.92e-04 | test_loss: 3.67e-04 | reg: 6.50e00 | 100%| 100/100 # saving model version 1.8此时不再传lamb默认lamb0.损失进一步降至 2.92e-04train/3.67e-04test几乎与数据噪声同量级。绘制最终网络model.plot()使用默认参数可以看到结构已经非常紧凑5. 自动符号化与公式提取5.1 二次剪枝与 auto_symbolicmodel model.prune() model.auto_symbolic() # saving model version 1.9 # fixing (0,0,0) with 0 ... fixing (0,1,1) with x^2, r20.9983, c2 ... # fixing (1,0,0) with x^2, r20.9993, c2 # fixing (1,1,0) with x, r20.9999, c1 # fixing (1,2,0) with x, r20.9985, c1 # fixing (1,3,0) with 0, r20.0, c0 # saving model version 1.10从 auto_symbolic 源码 可见该方法遍历所有边已符号化的边symbolic_fun[l].mask 0且激活 mask 为 0→ 跳过并打印skipping ...无贡献的边 →fix_symbolic(..., 0)打印fixing (...) with 0其余边 → 调用suggest_symbolic在候选函数库中搜索最佳符号函数打印fixing (...) with {name}, r2{r2}, c{c}其中r2为拟合优度、c为函数复杂度。从输出可以看到关键结果第 1 层的激活边被依次固定为x^2r2≈0.998、sinr2≈0.999、xr2≈0.999等——尤其是x^2与x正是 Neo-Hookean P12 中乘积/线性项的基础构件大量边被固定为 0说明稀疏化剪枝已把无关的 F 分量如 F13、F31 等排除。5.2 符号化后再拟合与公式输出model.fit(dataset, steps100) # | train_loss: 7.36e-03 | test_loss: 7.64e-03 | reg: 0.00e00 | 100%| 100/100 # saving model version 1.11符号化后reg归零所有边都已固定为符号函数、不再有样条仅需拟合符号函数中的仿射系数系数/偏置。训练损失略升至 7e-03 量级属于符号拟合 系数精调阶段的正常水平。最后用symbolic_formula导出公式并用ex_round把浮点数统一四舍五入到 2 位小数from kan.utils import ex_round ex_round(expand(ex_round(model.symbolic_formula(varinput_vars)[0][0],4)),2) # 0.02 F12² 0.42 F12 0.44 F13·F23 − 0.03 F21² 0.42 F21几点解读symbolic_formula(varinput_vars)把网络中所有符号边递归展开为 SymPy 表达式见 symbolic_formula[0][0]取第 1 个输出P12的完整表达式ex_round(expr, n)见 utils.py 的 ex_round遍历表达式中的所有sympy.Float并四舍五入到指定位数——外层expand(...)先展开乘积内层先保留 4 位、外层再规整为 2 位避免舍入误差累积最终表达式中的0.42与 Lame 参数中的 μ 高度吻合0.42 ≈ μ线性项0.42·F12 0.42·F21保留了线弹性先验μ(F12 F21)的核心结构而F13·F23乘积项与F12²、F21²平方项则是从数据中长出的非线性修正对应 Neo-Hookean P12 表达式μ(F12·F11 F22·F21 F32·F31)中的乘积机制具体保留哪些项取决于数据与稀疏化路径。这一结果完美展示了本教程的核心方法论先验线弹性提供骨架数据超弹性填充血肉——KAN 在训练中既继承了先验中的线性项又自主发现了需要非线性推广的项最终输出一个可读、可验证的符号公式。6. 方法总结与实验建议本教程Physics 4B的完整工作流可归纳为步骤关键 API作用编译先验kanpiler(input_vars, expr, base_funidentity)把已知解析公式变成符号化 MultKAN容量扩展expand_depth()/expand_width(layer, n, sum_bool, mult_arity)加深/加宽网络引入乘法节点解除约束perturb(modeall)激活全部边允许结构演化数据拟合fit(dataset, steps, lamb)用 L-BFGS 训练样条系数结构稀疏化prune()按归因分数剪除冗余节点/边版本管理rewind(model_id)回退到剪枝前的版本继续精调符号回归auto_symbolic()逐边从函数库挑选最佳符号函数公式输出symbolic_formula(var)ex_round导出并规整闭合表达式实际使用建议先验强度可控kanpiler编译出的网络边全部fix_symbolic固定训练前务必perturb解除部分边约束否则无法演化出先验之外的新项perturb(modeminimal)更适合希望强约束先验、只微调结构的场景正则项值得微调第一阶段lamb1e-5帮助稀疏化第二阶段去掉lamb追求精度若网络过于复杂可调大lamb若先验项被误删可先用rewind回退可复现性固定random/np/torch种子并开启torch.use_deterministic_algorithms(True)、torch.set_default_dtype(torch.float64)本实验才能逐字复现进一步探索可对照学习同一教程系列中无先验版本 Physics_4C_constitutive_laws_P12_without_prior.ipynb比较有/无符号先验时收敛速度与公式质量也可将本方法迁移到 P11含ln(|F|)对数项或其他本构关系的符号发现相关的编译器内部原理可参考 Interp_3_KAN_Compiler.ipynb 教程。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考