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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 19:22:49