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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 00:06:23