使用PyTorch DataLoader时出现标签错误的问题排查
问题根源与修正方案
你的问题核心是混淆了「样本索引」和「类别标签」:
get_relevant_indicies函数收集的是每个样本的类别标签(dataset[i][1]),但SubsetRandomSampler需要的是数据集的样本位置索引(即0到len(trainset)-1的整数)。- 你遍历sampler时看到的0、1、2是类别标签值,当DataLoader用这些值作为索引取样本时,实际是在取数据集中第0、1、2个样本——这几个样本的标签刚好只有0和1,所以输出和预期不符。
修正步骤
- 移除错误的
get_relevant_indicies函数,不需要用标签当索引。 - 生成正确的样本索引列表,打乱后传入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
相关产品推荐
相关产品推荐

