PyTorch TensorDataset随机分割报错:randperm参数无效求助
问题原因
random_split()的第二个参数需要传入整数类型的子集长度列表,你用0.6*len(dataset)和0.4*len(dataset)计算得到的是浮点数(哪怕数值是整数,类型仍为float),导致randperm()接收到无效的浮点参数,触发类型错误。
解决方法
将子集长度转换为整数,同时确保两个子集长度之和与原数据集总长度一致:
修正后的代码:
import numpy as np import torch from torch.utils.data import TensorDataset, random_split x_numpy = # (20640x8) matrix of floats y_numpy = # (20640x1) vector of floats x = torch.from_numpy(x_numpy.astype(np.float32)) y = torch.from_numpy(y_numpy.astype(np.float32)) dataset = TensorDataset(x, y) # 计算整数形式的训练集、测试集长度 train_size = int(0.6 * len(dataset)) test_size = len(dataset) - train_size # 避免四舍五入导致的长度偏差,保证总数匹配 trainSet, testSet = random_split(dataset, [train_size, test_size])
补充说明
若你的PyTorch版本在v1.10及以上,也可直接传入比例列表(需确保比例和为1):
trainSet, testSet = random_split(dataset, [0.6, 0.4])
但使用整数长度的方式更稳妥,能避免浮点数精度问题引发的长度不匹配。
内容的提问来源于stack exchange,提问作者Thunder
相关产品推荐
相关产品推荐

