使用torch.utils.data.random_split按百分比分割数据集报错排查
问题原因与解决方法
错误根源
你把DataLoader实例传给了random_split,但这个函数仅支持传入Dataset类对象(或实现了__len__和__getitem__的可索引数据集),而DataLoader是用来加载Dataset的迭代器,并不满足random_split的输入要求。
虽然你传入的比例总和为1,但random_split在处理DataLoader时,无法正确完成数据集的分割逻辑(比如无法索引访问样本),最终触发长度不匹配的报错。
修正后的代码
import torch from torch.utils.data import DataLoader, random_split, TensorDataset # 原始数据 list_dataset = [1,2,3,4,5,6,7,8,9,10] # 方式1:直接用列表作为Dataset(PyTorch支持将可索引、有长度的序列当作Dataset) # dataset = list_dataset # 方式2:包装成规范的TensorDataset(更推荐) dataset = TensorDataset(torch.tensor(list_dataset)) # 先分割Dataset train_dataset, val_dataset, test_dataset = random_split( dataset, [0.8, 0.1, 0.1], generator=torch.Generator().manual_seed(123) ) # 再为每个分割后的子集创建DataLoader train_loader = DataLoader(train_dataset, batch_size=1, shuffle=False) val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False) test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)
关键说明
- Dataset与DataLoader的分工:Dataset负责存储和提供单个样本,DataLoader负责批量加载、打乱样本等操作。
random_split是针对Dataset的分割工具,不能直接作用于DataLoader。 - 比例分割的要求:当传入比例值时,PyTorch会自动根据Dataset的总长度计算各子集的样本数,前提是输入的是合法的Dataset对象。
内容的提问来源于stack exchange,提问作者Pro Q
相关产品推荐
相关产品推荐

