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

Google Colab中torch.random_split()报错,未传Generator却提示参数错误

TypeError Traceback (most recent call last)
in
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 15:35:23