torch.utils.data.random_split()分割Dataset无效原因咨询
问题解答
random_split根本没失效,是你查看长度的方式错了。
这个函数不会复制原数据集生成新的独立数据集,它返回的是Subset对象——这类对象只是存了原数据集的引用和对应子集的索引列表。所以你访问train_dataset.dataset.imgs时,本质是在调用原full_dataset的imgs属性,长度自然和原数据集一模一样。
正确查看子集长度的方法是直接用len(train_dataset)和len(test_dataset),Subset类会根据内部的索引数量返回真实的子集长度。
给你贴个正确的代码示例:
from torchvision.datasets import ImageFolder from torch.utils.data import random_split full_dataset = ImageFolder('some_dir', transform=transform) print(len(full_dataset)) # 输出32854 # 按比例分割 train_set, test_set = random_split(full_dataset, [0.8, 0.2]) print(len(train_set)) # 约26283 print(len(test_set)) # 约6571 # 按固定长度分割 train_set, test_set = random_split(full_dataset, [len(full_dataset)-100, 100]) print(len(train_set)) # 32754 print(len(test_set)) # 100
要是你想确认分割后的样本确实是原数据集的子集,可以看train_set.indices,这里面存的就是子集中样本在原数据集中的索引位置,也可以直接通过train_set[0]取出对应样本验证。
内容的提问来源于stack exchange,提问作者let me down slowly
相关产品推荐
相关产品推荐

