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

使用PyTorch DataLoader时出现标签错误的问题排查

问题根源与修正方案

你的问题核心是混淆了「样本索引」和「类别标签」:

  • get_relevant_indicies函数收集的是每个样本的类别标签(dataset[i][1]),但SubsetRandomSampler需要的是数据集的样本位置索引(即0到len(trainset)-1的整数)。
  • 你遍历sampler时看到的0、1、2是类别标签值,当DataLoader用这些值作为索引取样本时,实际是在取数据集中第0、1、2个样本——这几个样本的标签刚好只有0和1,所以输出和预期不符。

修正步骤

  1. 移除错误的get_relevant_indicies函数,不需要用标签当索引。
  2. 生成正确的样本索引列表,打乱后传入Sampler即可。

修正后的完整代码

def get_data(batch_size, folder):
    """Takes a batch_size and the name of the folder (name of folder most likely called dataset)
    Example:
    get_data(1, "~/aps360-proj/dataset")
    
    """
    classes = ("testing1", "testing2", "testing3")
    
    transform = transforms.Compose(
        [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]
    )
    # Load images
    trainset = torchvision.datasets.ImageFolder(folder, transform=transform)    
    # 生成正确的样本索引列表:0到len(trainset)-1
    relevant_train_indicies = list(range(len(trainset)))

    np.random.seed(1)
    np.random.shuffle(relevant_train_indicies)
    random_sampler = SubsetRandomSampler(relevant_train_indicies)
    
    # 验证:遍历sampler会输出打乱后的样本索引
    for i in random_sampler:
        print(i)
    
    train_loader = torch.utils.data.DataLoader(trainset, sampler=random_sampler, batch_size=batch_size)
    for images, labels in train_loader:
        print(labels)

额外说明

如果你的需求是按类别筛选样本(比如只保留某些类的样本),正确做法是先筛选符合条件的样本索引,再传给Sampler,示例如下:

# 比如只保留类别标签为0和2的样本
filtered_indices = [i for i, (_, label) in enumerate(trainset) if label in (0, 2)]
np.random.shuffle(filtered_indices)
random_sampler = SubsetRandomSampler(filtered_indices)

内容的提问来源于stack exchange,提问作者danielliuce

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 06:55:25