如何使用torchdata正确划分训练集与测试集?
TorchData v0.6.0 训练/测试集划分的官方推荐方法
首先说明你遇到的Sampler报错原因:torch.utils.data.SubsetRandomSampler是类实例,而TorchData的Sampler DataPipe要求第二个参数是可返回索引迭代器的可调用对象(比如函数),直接传实例会因不可调用报错。
下面是几种官方推荐的更简洁划分方式:
1. 使用torch.utils.data.random_split(适用于有明确长度的DataPipe)
如果你的DataPipe实现了__len__方法(比如SequenceWrapper),这个工具最简洁:
from torch.utils.data import random_split from torchdata.datapipes.iter import SequenceWrapper dp = SequenceWrapper(range(5)) # 可按比例划分,也可传入具体样本数如[4, 1] train_dp, test_dp = random_split(dp, [0.8, 0.2])
注:返回的train_dp和test_dp是Subset对象,若需继续使用TorchData的DataPipe操作,可通过IterDataPipeAdapter转换,一般直接用于训练即可。
2. 修正Sampler的用法(自定义索引划分)
若想继续用Sampler,需把SubsetRandomSampler包装成可调用函数(返回其迭代器),或直接生成索引迭代器:
from torchdata.datapipes.iter import SequenceWrapper, Sampler from torch.utils.data import SubsetRandomSampler dp = SequenceWrapper(range(5)) indices = list(range(len(dp))) # 手动划分索引(也可提前随机打乱) train_indices = indices[:4] test_indices = indices[4:] # 方式1:用SubsetRandomSampler包装成可调用函数 train_dp = Sampler(dp, sampler_fn=lambda: iter(SubsetRandomSampler(train_indices))) test_dp = Sampler(dp, sampler_fn=lambda: iter(SubsetRandomSampler(test_indices))) # 方式2:直接传入返回索引迭代器的函数 train_dp = Sampler(dp, sampler_fn=lambda: iter(train_indices)) test_dp = Sampler(dp, sampler_fn=lambda: iter(test_indices))
3. 简化版Demultiplexer用法
你之前的Demultiplexer方法可以简化为链式调用,避免单独拆包:
from torchdata.datapipes.iter import SequenceWrapper dp = SequenceWrapper(range(5)) train_len = int(len(dp) * 0.8) train_dp, test_dp = ( dp.enumerate() .demux(num_instances=2, classifier_fn=lambda x: 0 if x[0] < train_len else 1) .map(lambda x: x[1]) )
其中官方最推荐第一种random_split方法,当DataPipe有明确长度时,它最简洁且符合PyTorch生态的使用习惯;如果是无长度的流式DataPipe,则更适合用Demultiplexer或自定义Sampler的方式。
内容的提问来源于stack exchange,提问作者user3002473
相关产品推荐
相关产品推荐

