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 将MultKAN与utils全部导出,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);这段代码展示了设备管理的三个核心动作:
- 模型侧:
KAN(...).to(device)调用 PyTorch 标准的.to()方法,将 KAN 内部所有参数与缓冲区迁移到目标设备。 - 数据侧:
create_dataset(f, n_var=4, train_num=1000, device=device)将生成的训练/测试张量直接创建在目标设备上。 - 训练侧:
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 与熵正则的加权和),由lamb与lamb_entropy控制;it/s:每秒迭代步数,用于直观对比设备性能。
文档中的实测数据显示:CPU 上约 2.56 it/s(50 步耗时约 19 秒),CUDA 上约 26.45 it/s(50 步仅约 1 秒),同配置下 CUDA 吞吐量约为 CPU 的 10 倍。需要说明的是,该数据来自演示文档当时的运行环境,实际加速比取决于 GPU 型号、张量规模与 CPU 性能,仅作量级参考。同时可以看到:无论设备如何,损失曲线与正则值几乎一致(6.83e-01与7.21e-01),说明设备选择不影响训练结果的数值质量。
训练结束后,模型自动以版本号形式保存到./model目录(saving model version 0.0、saving 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):实现的关键点:
- 默认设备为 CPU(
device='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_input、train_label、test_input、test_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) |
lr | 1.0 | 学习率。Adam 下通常取1e-3量级 |
steps | 100 | 训练步数 |
lamb | 0.0 | 整体正则惩罚强度(1e-3) |
lamb_entropy | 2.0 | 熵正则惩罚强度,促使激活函数稀疏化(示例取5.) |
update_grid | True | 是否定期更新网格。示例取False(固定网格,加快训练) |
batch | -1 | 批大小,-1表示全量梯度 |
fit返回的results字典包含train_loss、test_loss、reg三个一维数组,即训练日志中逐项打印的指标来源。当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为准。
实操要点与常见问题
- 设备一致性:始终用同一个变量
device同时驱动model.to(device)与create_dataset(..., device=device),避免模型在 GPU、数据在 CPU(或相反)导致的Expected all tensors to be on the same device报错。 - 显存占用:
width=[4,100,100,100,1]这类宽网络在 GPU 上训练的显存开销远大于 CPU 内存,若显存不足可减小width或grid。 - CPU 回退:
torch.device('cuda' if torch.cuda.is_available() else 'cpu')是跨环境安全的写法,在无 GPU 的 CI 或服务器上自动回退 CPU,保证代码可移植。 - 确定性:
seed=0同时作用于create_dataset内部的np.random.seed与torch.manual_seed(见 kan/utils.py),保证同一设备上多次运行结果一致;跨设备之间数值可能有微小浮点差异,属于正常现象。 - 仅推理场景:如果只是加载已有 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),仅供参考