pykan 设备(Device)管理与 CUDA 加速实战:为 KAN 模型与数据集正确传递 device 参数
2026/9/14 18:53:16 网站建设 项目流程

pykan 设备(Device)管理与 CUDA 加速实战:为 KAN 模型与数据集正确传递 device 参数

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

导读

本文基于 pykan 官方 API 演示文档 API_10_device.rst 及其配套 Notebook API_10_device.ipynb,讲解 Kolmogorov-Arnold Networks 在 pykan 框架中如何完成 CPU / CUDA 设备选择与切换。核心要点是:使用 CUDA 时,必须同时把 device 参数传给模型(通过 PyTorch 标准的.to(device))和数据集(通过create_dataset(..., device=device)。读完本文,你将掌握设备探测的推荐写法、设备参数在模型与数据两侧的传递方式、以及一个可复现的 4 输入目标函数拟合示例,并能理解设备不一致可能带来的隐患。

为什么需要显式管理设备

在 pykan 的 API 系列演示中,其余示例默认都在 CPU 上运行(device = 'cpu')。但当模型规模增大、网格(grid)变密或训练步数变多时,CPU 训练会成为瓶颈。此时将模型与数据迁移到 NVIDIA GPU(CUDA)是最直接的加速手段。

关键点在于:pykan 的 KAN 模型是 PyTorch 的nn.Module子类,而数据是通过create_dataset生成的字典,二者都需要感知设备。仅把模型搬到 GPU 而数据仍在 CPU(或反之),在训练时就会触发张量设备不匹配的运行时错误(RuntimeError)。因此设备管理必须"双管齐下"。

从源码看,kan/init.py 将MultKANutils全部导出,KAN正是MultKAN的别名(见 kan/MultKAN.py 中的KAN = MultKAN),因此下面的示例中from kan import KAN, create_dataset即可同时拿到模型类与数据集构造工具。

第一步:探测可用设备

文档给出的推荐写法是利用torch.cuda.is_available()做运行时探测,优雅地回退到 CPU:

from kan import KAN, create_dataset import torch # 自动探测:有 CUDA 用 cuda,否则退回 cpu device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(device) # 也可以直接硬编码设备字符串(演示文档中展示了这种写法) device = 'cpu' print(device)

上述代码在无 GPU 的环境中输出cpu(演示文档的运行结果即为cpu)。

两种写法均合法:torch.device('cpu')与字符串'cpu'在 PyTorch API 中是可互换的。torch.device(...)返回的是torch.device对象,能提供更严格的类型检查,推荐在正式代码中使用;直接写字符串更简洁,适合在 Notebook 中快速切换实验。

第二步:把设备同时传给模型与数据集

设备探测完成后,模型和数据必须使用同一个device。文档给出了完整的训练流程:

model = KAN(width=[4,100,100,100,1], grid=3, k=3, seed=0).to(device) # 目标函数:四输入变量的正弦指数复合函数 f = lambda x: torch.exp((torch.sin(torch.pi*(x[:,[0]]**2+x[:,[1]]**2))+ torch.sin(torch.pi*(x[:,[2]]**2+x[:,[3]]**2)))/2) dataset = create_dataset(f, n_var=4, train_num=1000, device=device) # 训练模型 model.fit(dataset, opt="Adam", lr=1e-3, steps=50, lamb=1e-3, lamb_entropy=5., update_grid=False);

这段代码展示了设备管理的三个核心动作:

  1. 模型侧KAN(...).to(device)调用 PyTorch 标准的.to()方法,将 KAN 内部所有参数与缓冲区迁移到目标设备。
  2. 数据侧create_dataset(f, n_var=4, train_num=1000, device=device)将生成的训练/测试张量直接创建在目标设备上。
  3. 训练侧model.fit(dataset, ...)在训练时无需再指定设备,因为模型与数据已经位于同一设备。

关于模型构建参数:width=[4,100,100,100,1]表示 4 维输入、3 层各 100 个隐藏神经元、1 维输出;grid=3为 B 样条网格间隔数;k=3为样条阶数;seed=0固定随机种子,保证实验可复现。这些参数与设备无关,CPU 与 CUDA 下保持一致即可。

第三步:训练与性能对比

文档在 CPU 与 CUDA 两种设备上运行了完全相同的训练配置,训练过程打印如下信息:

checkpoint directory created: ./model saving model version 0.0 | train_loss: 6.83e-01 | test_loss: 7.21e-01 | reg: 1.04e+03 | : 100%|█| 50/50 [00:19<00:00, 2.56it/s] saving model version 0.1

训练日志中的关键信息解读:

  • train_loss: 6.83e-01/test_loss: 7.21e-01:训练与测试的 RMSE 损失,二者接近说明没有明显过拟合;
  • reg: 1.04e+03:正则化项(包含 L1 与熵正则的加权和),由lamblamb_entropy控制;
  • it/s:每秒迭代步数,用于直观对比设备性能。

文档中的实测数据显示:CPU 上约 2.56 it/s(50 步耗时约 19 秒),CUDA 上约 26.45 it/s(50 步仅约 1 秒),同配置下 CUDA 吞吐量约为 CPU 的 10 倍。需要说明的是,该数据来自演示文档当时的运行环境,实际加速比取决于 GPU 型号、张量规模与 CPU 性能,仅作量级参考。同时可以看到:无论设备如何,损失曲线与正则值几乎一致(6.83e-017.21e-01),说明设备选择不影响训练结果的数值质量。

训练结束后,模型自动以版本号形式保存到./model目录(saving model version 0.0saving model version 0.1),这与 API_12_checkpoint_save_load_model 演示的 checkpoint 机制一致。

源码级验证:device 参数在 create_dataset 中的传递

为了理解create_dataset如何把数据放到目标设备,可以查看 kan/utils.py 的实现。其函数签名如下:

def create_dataset(f, n_var=2, f_mode='col', ranges=[-1, 1], train_num=1000, test_num=1000, normalize_input=False, normalize_label=False, device='cpu', seed=0):

实现的关键点:

  • 默认设备为 CPUdevice='cpu'),与文档所述"所有其他 demo 默认使用 cpu"一致;
  • 输入与标签先以torch.zeros/torch.rand在 CPU 上生成,最后统一通过.to(device)迁移(见 kan/utils.py):
dataset['train_input'] = train_input.to(device) dataset['test_input'] = test_input.to(device) dataset['train_label'] = train_label.to(device) dataset['test_label'] = test_label.to(device)
  • 返回的dataset是一个字典,包含train_inputtrain_labeltest_inputtest_label四个键;
  • 其他常用参数:ranges控制输入采样区间(默认[-1, 1],既可以是 1D 也可以是(n_var, 2)的逐变量范围)、train_num/test_num控制样本量、normalize_input/normalize_label控制是否标准化、seed控制采样随机性(默认 0)。

由此可以确认一个事实:只要在create_dataset中传入device=device,返回的四个张量就已经位于目标设备,训练时无需再手动搬运。而model.to(device)则负责模型参数侧,两者结合才构成完整的设备配置。

fit 训练接口中的相关参数

文档中的训练调用使用了model.fit(...),其完整签名定义在 kan/MultKAN.py。与本文示例直接相关的参数含义如下:

参数默认值说明
opt"LBFGS"优化器,可选"LBFGS""Adam"。文档示例使用 Adam(配合lr=1e-3
lr1.0学习率。Adam 下通常取1e-3量级
steps100训练步数
lamb0.0整体正则惩罚强度(1e-3
lamb_entropy2.0熵正则惩罚强度,促使激活函数稀疏化(示例取5.
update_gridTrue是否定期更新网格。示例取False(固定网格,加快训练)
batch-1批大小,-1表示全量梯度

fit返回的results字典包含train_losstest_lossreg三个一维数组,即训练日志中逐项打印的指标来源。当lamb > 0时,源码要求save_act=True(见 kan/MultKAN.py),否则会打印警告并将lamb置 0,这一点在设置正则项时需要注意。

此外,演示文档中被注释掉的model.train(dataset, opt="LBFGS", steps=20, lamb=1e-3, lamb_entropy=2.)是历史 API 写法,当前仓库中 KAN 的训练入口统一为fit方法,建议以fit为准。

实操要点与常见问题

  1. 设备一致性:始终用同一个变量device同时驱动model.to(device)create_dataset(..., device=device),避免模型在 GPU、数据在 CPU(或相反)导致的Expected all tensors to be on the same device报错。
  2. 显存占用width=[4,100,100,100,1]这类宽网络在 GPU 上训练的显存开销远大于 CPU 内存,若显存不足可减小widthgrid
  3. CPU 回退torch.device('cuda' if torch.cuda.is_available() else 'cpu')是跨环境安全的写法,在无 GPU 的 CI 或服务器上自动回退 CPU,保证代码可移植。
  4. 确定性seed=0同时作用于create_dataset内部的np.random.seedtorch.manual_seed(见 kan/utils.py),保证同一设备上多次运行结果一致;跨设备之间数值可能有微小浮点差异,属于正常现象。
  5. 仅推理场景:如果只是加载已有 checkpoint 做推理(参考 API_12_checkpoint_save_load_model),同样需要先model.to(device)再喂入与模型同设备的输入张量。

小结

pykan 的设备管理遵循 PyTorch 标准约定:通过torch.device('cuda' if torch.cuda.is_available() else 'cpu')探测可用设备,再把得到的device同时传给model.to(device)create_dataset(..., device=device)。从 kan/utils.py 的实现可以看到,数据集四个张量最终都会.to(device),与模型保持同步。设备切换只影响训练吞吐(文档实测 CUDA 约为 CPU 的 10 倍),不改变损失曲线与正则化的数值质量。将本文示例中的device替换为'cpu''cuda',即可在任意 PyTorch 环境中快速复现这套完整的 KAN 训练流程。

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

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

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

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

立即咨询