尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

pykan 训练过程可视化视频生成实战:用 `fit(save_fig=True)` 与 MoviePy 记录 KAN 结构演化

pykan 训练过程可视化视频生成实战:用 `fit(save_fig=True)` 与 MoviePy 记录 KAN 结构演化 pykan 训练过程可视化视频生成实战用fit(save_figTrue)与 MoviePy 记录 KAN 结构演化【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan导读本篇指南聚焦 pykanKolmogorov-Arnold Networks教程 API 9: Videos讲解如何把 KAN 模型在训练过程中每个 step 的结构图由plot()生成自动保存为图片序列再用 MoviePy 合成 MP4 视频从而完整记录模型结构的演化动态。读完本文你将掌握fit()中save_fig、img_folder、save_fig_freq、in_vars/out_vars等关键参数的正确用法并能够把任意 KAN 训练过程导出为可回放、可分享的视频。本文所有代码与参数均以当前仓库 kan/MultKAN.py 的实现为准。背景为什么需要记录 KAN 的训练动态plot()方法可以在任意时刻绘制 KAN 的结构图包括各层激活函数曲线、边的透明度反映激活强度以及输入输出变量标注。单张静态图只能展示某一时刻的网络状态而 KAN 的训练过程尤其是样条网格自适应更新、符号化切换、稀疏化与剪枝会显著改变网络结构。把这些快照按训练 step 串联成视频是理解模型内部演化、调试正则化强度与诊断过拟合的直观手段。本教程展示的正是这样一条完整链路训练时按 step 落盘结构图 → 按文件名排序读取图片序列 → 用 MoviePy 编码为 MP4。第一步构造 KAN 模型与合成数据集视频功能本身不依赖特殊数据教程使用create_dataset生成合成数据来演示。from kan import * import torch device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) # create a KAN: 2D inputs, 1D output, and 5 hidden neurons. cubic spline (k3), 5 grid intervals (grid5). model KAN(width[4,2,1,1], grid3, k3, seed1, devicedevice) f lambda x: torch.exp((torch.sin(torch.pi*(x[:,[0]]**2x[:,[1]]**2))torch.sin(torch.pi*(x[:,[2]]**2x[:,[3]]**2)))/2) dataset create_dataset(f, n_var4, train_num3000, devicedevice)要点说明KAN(width[4,2,1,1], ...)构建一个 4 维输入、2 个隐藏神经元、1 个输出的 4 层 KANwidth定义各层神经元数量。grid3, k3样条网格间隔数为 3样条阶数 k3三次样条。grid 越大样条表达能力越强。seed1固定随机种子保证可复现。create_dataset在 kan/utils.py 中定义签名含n_var、ranges默认 [-1,1]、train_num、test_num默认各 1000等参数本示例取n_var4、train_num3000返回的dataset字典包含train_input/train_label/test_input/test_label四组张量。目标函数f由两个sin(pi*(x_i^2x_j^2))项之和取指数构成便于观察 KAN 是否逼近出对称结构。第二步在训练时开启结构图快照视频的关键开关在fit()方法中。从 kan/MultKAN.py 的签名可见相关参数model.fit(dataset, optLBFGS, steps5, lamb0.001, lamb_entropy2., save_figTrue, beta10, in_vars[r$x_1$, r$x_2$, r$x_3$, r$x_4$], out_vars[r${\rm exp}({\rm sin}(x_1^2x_2^2){\rm sin}(x_3^2x_4^2))$], img_folderimage_folder);其中image_folder video_img是图片输出目录。关键参数详解参数默认值作用save_figFalse置为True开启训练过程结构图保存img_folder./video图片输出目录会自动创建见 kan/MultKAN.pysave_fig_freq1每save_fig_freq个 step 保存一次图见 kan/MultKAN.pybeta3边透明度控制transparency tanh(beta*l1)beta越大弱激活边越透明见 kan/MultKAN.pyin_vars/out_varsNone输入/输出变量的 LaTeX 名称用于标注图的输入输出节点optLBFGS优化器支持LBFGS或Adam见 kan/MultKAN.pylamb0.总体正则强度注意若lamb 0而self.save_actFalse代码会打印提示并把lamb置 0见 kan/MultKAN.pylamb_entropy2.熵正则强度促使激活稀疏化steps100训练步数源码视角快照如何落盘在 kan/MultKAN.py 的训练循环中核心逻辑为若save_figTrue且img_folder不存在则os.makedirs(img_folder)自动建目录每个 step 开头若save_fig且_ % save_fig_freq 0临时把self.save_act置为True确保前向传播会缓存激活值model.acts否则plot()会提示cannot plot since data are not saved见 kan/MultKAN.py完成该 step 的优化与评估后调用self.plot(folderimg_folder, in_varsin_vars, out_varsout_vars, titleStep {}.format(_), betabeta)绘制结构图并通过plt.savefig(img_folder / str(_) .jpg, bbox_inchestight, dpi200)以 200 dpi 写入{step}.jpg随后plt.close()释放画布并恢复save_act原值。因此video_img/目录下会生成0.jpg、1.jpg、2.jpg、3.jpg、4.jpg等以 step 编号命名的 JPG 文件这些正是后续合成视频的帧。运行输出参考教程在 CUDA 环境下运行会输出cuda checkpoint directory created: ./model saving model version 0.0 | train_loss: 2.89e-01 | test_loss: 2.96e-01 | reg: 1.31e01 | : 100%|█| 5/5 [00:0900:00, 1.94s/it saving model version 0.1其中train_loss/test_loss为 RMSEreg为当前正则项数值saving model version 0.x说明训练过程同时触发了检查点保存机制参考 model/0.0_config.yml 等文件。第三步把图片序列合成为 MP4 视频训练结束后教程使用 MoviePy版本要求moviepy 1.0.3把图片序列编码为视频import os import numpy as np import moviepy.video.io.ImageSequenceClip # moviepy 1.0.3 video_namevideo fps5 files os.listdir(image_folder) train_index [] for file in files: if file[0].isdigit() and file.endswith(.jpg): train_index.append(int(file[:-4])) train_index np.sort(train_index) image_files [image_folder/str(train_index[index]).jpg for index in train_index] clip moviepy.video.io.ImageSequenceClip.ImageSequenceClip(image_files, fpsfps) clip.write_videofile(video_name.mp4)这段代码的关键点files列出video_img/下所有文件筛选出以数字开头且以.jpg结尾的文件取int(file[:-4])得到 step 编号对应 kan/MultKAN.py 的命名规则str(_) .jpgnp.sort保证帧按训练步数升序排列避免字符串排序导致10.jpg排在2.jpg之前ImageSequenceClip(image_files, fpsfps)用 5 fps 把帧序列封装成视频片段clip.write_videofile(video.mp4)调用 FFmpeg 编码输出。运行日志参考Moviepy - Building video video.mp4. Moviepy - Writing video video.mp4 Moviepy - Done ! Moviepy - video ready video.mp4参数调优建议fps帧率教程使用 5。steps较大时可适当提高如 1015以加快播放帧数少时建议降低避免一闪而过。save_fig_freq默认 1每个 step 都存图。若训练步数多且图尺寸大可设为 2 或 5 以减小磁盘占用与编码耗时它同时控制图片数量从而影响视频时长。beta控制边透明度对比度。beta10时弱激活边几乎完全透明视频中能清晰看到激活结构的收缩与稀疏化调小则整体更实。in_vars/out_vars传入 LaTeX 字符串如r$x_1$可在图上标注变量名便于对照目标函数理解每一路激活不传则显示默认x_1...x_n/y_1...。optLBFGS适合小批量全量优化本项目对 LBFGS 启用了 strong_wolfe 线搜索等配置见 kan/MultKAN.pyAdam适合大批量场景不同优化器的收敛轨迹不同视频内容也会有差异。常见问题cannot plot since data are not savedfit(save_figTrue)会临时开启save_act不会出现该问题若单独调用plot()需先保证self.save_actTrue并做过前向传播见 kan/MultKAN.py。setting lamb0提示当lamb 0但save_actFalse时正则会被静默关闭请在需要正则时保持save_figTrue或手动设置self.save_actTrue见 kan/MultKAN.py。视频画面顺序错乱确保按数字排序用np.sort而非默认的字符串排序。MoviePy 报错本项目教程基于moviepy 1.0.3版本不匹配可能引起编码器接口变化建议按此版本安装。图片目录已存在旧文件img_folder不会被清空重跑前建议清理旧 JPG否则os.listdir会把历史帧一并纳入视频。延伸阅读完整教程源码见 tutorials/API_demo/API_9_video.ipynb对应 RST 文档见 docs/API_demo/API_9_video.rstplot()的完整参数metric、scale、tick、sample、varscale等见 kan/MultKAN.pyfit()的完整参数与返回值说明见 kan/MultKAN.py数据集构造工具create_dataset/create_dataset_from_data见 kan/utils.py 与 kan/utils.py更多训练可视化用例可参考 docs/API_demo/API_2_plotting.rst。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表