PyKAN分类示例代码运行报错:数据类型与设备不匹配求助
PyKAN分类任务中的类型与设备不匹配问题解决
问题背景
严格按照PyKAN文档示例3编写分类代码时,先后遇到两个错误:
- 首次报错:
RuntimeError: expected scalar type Double but found Float - 修改模型为double类型并指定设备后,新报错:
RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cpu!
首次报错栈
description: 0%| | 0/20 [00:00<?, ?it/s] --------------------------------------------------------------------------- RuntimeError Traceback (most recent call last) Cell In[44], line 9 6 def test_acc(): 7 return torch.mean((torch.argmax(model(dataset['test_input']), dim=1) == dataset['test_label']).float()) ----> 9 results = model.train(dataset, opt="LBFGS", steps=20, metrics=(train_acc, test_acc), loss_fn=torch.nn.CrossEntropyLoss()) File ~\anaconda3\Lib\site-packages\kan\KAN.py:898, in KAN.train(self, dataset, opt, steps, log, lamb, lamb_l1, lamb_entropy, lamb_coef, lamb_coefdiff, update_grid, grid_update_num, loss_fn, lr, stop_grid_update_step, batch, small_mag_threshold, small_reg_factor, metrics, sglr_avoid, save_fig, in_vars, out_vars, beta, save_fig_freq, img_folder, device) 895 test_id = np.random.choice(dataset['test_input'].shape[0], batch_size_test, replace=False) 897 if _ % grid_update_freq == 0 and _ < stop_grid_update_step and update_grid: --> 898 self.update_grid_from_samples(dataset['train_input'][train_id].to(device)) 900 if opt == "LBFGS": 901 optimizer.step(closure) File ~\anaconda3\Lib\site-packages\kan\KAN.py:243, in KAN.update_grid_from_samples(self, x) 220 ''' 221 update grid from samples 222 (...) 240 tensor([0.0128, 1.0064, 2.0000, 2.9937, 3.9873, 4.9809]) 241 ''' 242 for l in range(self.depth): --> 243 self.forward(x) 244 self.act_fun[l].update_grid_from_samples(self.acts[l]) File ~\anaconda3\Lib\site-packages\kan\KAN.py:311, in KAN.forward(self, x) 307 self.acts.append(x) # acts shape: (batch, width[l]) 309 for l in range(self.depth): --> 311 x_numerical, preacts, postacts_numerical, postspline = self.act_fun[l](x) 313 if self.symbolic_enabled == True: 314 x_symbolic, postacts_symbolic = self.symbolic_fun[l](x) File ~\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1532, in Module._wrapped_call_impl(self, *args, **kwargs) 1530 return self._compiled_call_impl(*args, **kwargs) # type: ignore[misc] 1531 else: -> 1532 return self._call_impl(*args, **kwargs) File ~\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1541, in Module._call_impl(self, *args, **kwargs) 1536 # If we don't have any hooks, we want to skip the rest of the logic in 1537 # this function, and just call forward. 1538 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks 1539 or _global_backward_pre_hooks or _global_backward_hooks 1540 or _global_forward_hooks or _global_forward_pre_hooks): -> 1541 return forward_call(*args, **kwargs) 1543 try: 1544 result = None File ~\anaconda3\Lib\site-packages\kan\KANLayer.py:173, in KANLayer.forward(self, x) 171 preacts = x.permute(1, 0).clone().reshape(batch, self.out_dim, self.in_dim) 172 base = self.base_fun(x).permute(1, 0) # shape (batch, size) --> 173 y = coef2curve(x_eval=x, grid=self.grid[self.weight_sharing], coef=self.coef[self.weight_sharing], k=self.k, device=self.device) # shape (size, batch) 174 y = y.permute(1, 0) # shape (batch, size) 175 postspline = y.clone().reshape(batch, self.out_dim, self.in_dim) File ~\anaconda3\Lib\site-packages\kan\spline.py:100, in coef2curve(x_eval, grid, coef, k, device) 65 ''' 66 converting B-spline coefficients to B-spline curves. Evaluate x on B-spline curves (summing up B_batch results over B-spline basis). 67 (...) 96 torch.Size([5, 100]) 97 ''' 98 # x_eval: (size, batch), grid: (size, grid), coef: (size, coef) 99 # coef: (size, coef), B_batch: (size, coef, batch), summer over coef --> 100 y_eval = torch.einsum('ij,ijk->ik', coef, B_batch(x_eval, grid, k, device=device)) 101 return y_eval File ~\anaconda3\Lib\site-packages\torch\functional.py:385, in einsum(*args) 380 return einsum(equation, *_operands) 382 if len(operands) <= 2 or not opt_einsum.enabled: 383 # the path for contracting 0 or 1 time(s) is already optimized 384 # or the user has disabled using opt_einsum -> 385 return _VF.einsum(equation, operands) # type: ignore[attr-defined] 387 path = None 388 if opt_einsum.is_available(): RuntimeError: expected scalar type Double but found Float
修改后的代码
from kan import KAN import matplotlib.pyplot as plt from sklearn.datasets import make_moons import torch import numpy as np device = "cuda:0" if torch.cuda.is_available() else "cpu" dataset = {} train_input, train_label = make_moons( n_samples=1000, shuffle=True, noise=0.1, random_state=None) test_input, test_label = make_moons( n_samples=1000, shuffle=True, noise=0.1, random_state=None) dataset['train_input'] = torch.from_numpy(train_input) dataset['test_input'] = torch.from_numpy(test_input) dataset['train_label'] = torch.from_numpy(train_label) dataset['test_label'] = torch.from_numpy(test_label) dataset['train_input'] = dataset['train_input'].to(device) dataset['test_input'] = dataset['train_input'].to(device) dataset['train_label'] = dataset['train_input'].to(device) dataset['test_label'] = dataset['train_input'].to(device) X = dataset['train_input'] y = dataset['train_label'] model = KAN(width=[2, 2], grid=3, k=3, device=device).double() print(model.device) print(dataset['train_input'].device) print(dataset['test_input'].device) print(dataset['train_label'].device) print(dataset['test_label'].device) def train_acc(): return torch.mean((torch.argmax(model(dataset['train_input']), dim=1) == dataset['train_label'])) def test_acc(): return torch.mean((torch.argmax(model(dataset['test_input']), dim=1) == dataset['test_label'])) results = model.train(dataset, opt="LBFGS", steps=20, metrics=( train_acc, test_acc), loss_fn=torch.nn.CrossEntropyLoss())
修改后报错栈
cuda:0 cuda:0 cuda:0 cuda:0 cuda:0 description: 0%| | 0/20 [00:00<?, ?it/s] Traceback (most recent call last): File "c:\Users\kshit\OneDrive\Documents\IRIS\KANforIRIS.py", line 45, in <module> results = model.train(dataset, opt="LBFGS", steps=20, metrics=( ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Users\kshit\AppData\Local\Programs\Python\Python311\Lib\site-packages\kan\KAN.py", line 898, in train self.update_grid_from_samples(dataset['train_input'][train_id].to(device)) File "C:\Users\kshit\AppData\Local\Programs\Python311\Lib\site-packages\kan\KAN.py", line 243, in update_grid_from_samples self.forward(x) File "C:\Users\kshit\AppData\Local\Programs\Python311\Lib\site-packages\kan\KAN.py", line 311, in forward x_numerical, preacts, postacts_numerical, postspline = self.act_fun[l](x) ^^^^^^^^^^^^^^^^^^ File "C:\Users\kshit\AppData\Local\Programs\Python311\Lib\site-packages\torch\nn\modules\module.py", line 1511, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Users\kshit\AppData\Local\Programs\Python311\Lib\site-packages\torch\nn\modules\module.py", line 1520, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Users\kshit\AppData\Local\Programs\Python311\Lib\site-packages\kan\KANLayer.py", line 170, in forward x = torch.einsum('ij,k->ikj', x, torch.ones(self.out_dim, device=self.device)).reshape(batch, self.size).permute(1, 0) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Users\kshit\AppData\Local\Programs\Python311\Lib\site-packages\torch\functional.py", line 380, in einsum return _VF.einsum(equation, operands) # type: ignore[attr-defined] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cpu!
问题分析与解决方案
错误根源
- 首次类型不匹配:PyKAN默认使用Double类型,但
make_moons生成的numpy数组转成torch张量后是Float类型,导致计算时类型冲突。 - 修改后的设备不匹配:代码中存在明显错误——将
test_input、train_label、test_label都错误赋值为train_input的拷贝;同时损失函数未同步到指定设备,部分内部张量可能存在设备不一致问题。
修复步骤
步骤1:修正数据集赋值与类型/设备同步
将数据集的创建代码修正为:
# 修正数据集类型与设备,标签转long类型(符合CrossEntropyLoss要求) dataset['train_input'] = torch.from_numpy(train_input).to(device).double() dataset['test_input'] = torch.from_numpy(test_input).to(device).double() dataset['train_label'] = torch.from_numpy(train_label).to(device).long() dataset['test_label'] = torch.from_numpy(test_label).to(device).long()
- 直接在转张量时指定
double()类型,与PyKAN默认类型保持一致; - 标签需转为
long()类型,因为CrossEntropyLoss要求目标为整数类型; - 修正之前错误的赋值逻辑,确保每个数据集元素对应正确的输入/标签。
步骤2:同步损失函数到设备
将损失函数也移动到指定设备:
loss_fn = torch.nn.CrossEntropyLoss().to(device)
步骤3:移除冗余的模型类型设置(可选)
如果数据集已经统一为double类型,无需手动设置模型为.double(),PyKAN默认的Double类型会与数据集自动匹配;若保留.double(),需确保所有输入张量类型一致。
完整修复后的代码
from kan import KAN import matplotlib.pyplot as plt from sklearn.datasets import make_moons import torch import numpy as np device = "cuda:0" if torch.cuda.is_available() else "cpu" dataset = {} train_input, train_label = make_moons( n_samples=1000, shuffle=True, noise=0.1, random_state=None) test_input, test_label = make_moons( n_samples=1000, shuffle=True, noise=0.1, random_state=None) # 修正数据集类型与设备,标签转long类型 dataset['train_input'] = torch.from_numpy(train_input).to(device).double() dataset['test_input'] = torch.from_numpy(test_input).to(device).double() dataset['train_label'] = torch.from_numpy(train_label).to(device).long() dataset['test_label'] = torch.from_numpy(test_label).to(device).long() # 初始化模型,无需额外.double()也可,因为数据集已匹配默认类型 model = KAN(width=[2, 2], grid=3, k=3, device=device) print(model.device) print(dataset['train_input'].device, dataset['train_input'].dtype) print(dataset['test_input'].device, dataset['test_input'].dtype) print(dataset['train_label'].device, dataset['train_label'].dtype) print(dataset['test_label'].device, dataset['test_label'].dtype) def train_acc(): return torch.mean((torch.argmax(model(dataset['train_input']), dim=1) == dataset['train_label']).float()) def test_acc(): return torch.mean((torch.argmax(model(dataset['test_input']), dim=1) == dataset['test_label']).float()) # 损失函数移到设备 loss_fn = torch.nn.CrossEntropyLoss().to(device) results = model.train(dataset, opt="LBFGS", steps=20, metrics=( train_acc, test_acc), loss_fn=loss_fn)
验证点
- 所有输入张量的
device和dtype与模型一致; - 标签类型为
long,符合CrossEntropyLoss的要求; - 损失函数同步到对应设备。
内容的提问来源于stack exchange,提问作者SidKT
相关产品推荐
相关产品推荐

