You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyKAN分类示例代码运行报错:数据类型与设备不匹配求助

PyKAN分类任务中的类型与设备不匹配问题解决

问题背景

严格按照PyKAN文档示例3编写分类代码时,先后遇到两个错误:

  1. 首次报错:RuntimeError: expected scalar type Double but found Float
  2. 修改模型为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!

问题分析与解决方案

错误根源

  1. 首次类型不匹配:PyKAN默认使用Double类型,但make_moons生成的numpy数组转成torch张量后是Float类型,导致计算时类型冲突。
  2. 修改后的设备不匹配:代码中存在明显错误——将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 14:27:02