Google Colab中torch.random_split()报错,未传Generator却提示参数错误
TypeError Traceback (most recent call last)
4 print("train_size: ", train_size)
5 print("validation_test_size: ", validation_test_size)
----> 6 train_dataset, validation_dataset, test_dataset = random_split(dataset,[train_size,validation_test_size/2,validation_test_size/2])
/usr/local/lib/python3.7/dist-packages/torch/utils/data/dataset.py in random_split(dataset, lengths, generator)
311 raise ValueError("Sum of input lengths does not equal the length of the input dataset!")
312
--> 313 indices = randperm(sum(lengths)).tolist()
314 return [Subset(dataset, indices[offset - length : offset]) for offset, length in zip(_accumulate(lengths), lengths)]
TypeError: randperm() received an invalid combination of arguments - got (float, generator=torch._C.Generator), but expected one of:
- (int n, *, torch.Generator generator, Tensor out, torch.dtype dtype, torch.layout layout, torch.device device, bool pin_memory, bool requires_grad)
- (int n, *, Tensor out, torch.dtype dtype, torch.layout layout, torch.device device, bool pin_memory, bool requires_grad)
# 问题解答 ## 1. torch._C.Generator是什么 `torch._C.Generator`是PyTorch随机数生成器的底层C++实现,是`torch.Generator`的后端支撑。`random_split`函数内部会默认调用这个生成器来生成随机划分的索引,报错中提到它只是因为参数类型错误触发了参数校验逻辑,并非你主动传入了该生成器导致的问题。 ## 2. 报错解决方法 报错的核心原因是传入`random_split`的长度列表包含**浮点数**:`validation_test_size/2`的计算结果是`2000.0`(浮点数),但`randperm`函数要求必须传入整数类型的长度参数。 只需将划分长度转为整数即可解决: ```python # 直接在传入时转成整数 train_dataset, validation_dataset, test_dataset = random_split( dataset, [train_size, int(validation_test_size/2), int(validation_test_size/2)] )
或者提前计算好整数形式的各集大小,代码更清晰:
train_size = int(0.8 * len(dataset)) val_size = int(0.1 * len(dataset)) test_size = len(dataset) - train_size - val_size train_dataset, validation_dataset, test_dataset = random_split(dataset, [train_size, val_size, test_size])
内容的提问来源于stack exchange,提问作者Ian Nathanael Hadinoto

