PyTorch中移除MNIST数据集中数字9的方法求助
解决MNIST数据集移除指定类别(数字9)的问题
因为MNIST类的targets(旧版本可能为train_labels)和data属性是只读的,直接赋值修改会报错,这里提供两种无需手动遍历生成元组列表的解决方案:
方案一:使用Subset快速创建过滤子集
torch.utils.data.Subset可以基于索引直接包装原数据集生成子集,完全不需要修改原数据集的只读属性,同时保留原数据集的所有配置(比如transform):
import torch from torchvision import datasets, transforms from torch.utils.data import Subset # 按需定义数据预处理transform transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载原始MNIST训练集 full_dataset = datasets.MNIST( root='./data', train=True, download=True, transform=transform ) # 生成所有标签不等于9的样本索引 valid_indices = torch.where(full_dataset.targets != 9)[0] # 创建过滤后的数据集 filtered_dataset = Subset(full_dataset, valid_indices) # 直接传入DataLoader使用 train_loader = torch.utils.data.DataLoader( filtered_dataset, batch_size=64, shuffle=True )
方案二:自定义Dataset实现灵活过滤
如果需要更复杂的过滤逻辑(比如排除多个类别),可以自定义一个简单的Dataset类包装原数据集,通过索引映射实现按需取数:
import torch from torchvision import datasets, transforms from torch.utils.data import Dataset, DataLoader class FilteredMNIST(Dataset): def __init__(self, original_dataset, exclude_label): self.original_dataset = original_dataset # 预计算符合条件的样本索引 self.valid_indices = torch.where(original_dataset.targets != exclude_label)[0] def __len__(self): return len(self.valid_indices) def __getitem__(self, idx): # 将当前索引映射到原数据集的有效索引 original_idx = self.valid_indices[idx] return self.original_dataset[original_idx] # 使用示例 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) full_dataset = datasets.MNIST( root='./data', train=True, download=True, transform=transform ) # 排除数字9的数据集 filtered_dataset = FilteredMNIST(full_dataset, exclude_label=9) train_loader = DataLoader(filtered_dataset, batch_size=64, shuffle=True)
两种方案都不需要手动遍历整个数据集生成元组列表,都是基于索引映射获取过滤后的样本,既规避了只读属性的限制,又保留了原数据集的所有特性。
内容的提问来源于stack exchange,提问作者smartstix
相关产品推荐
相关产品推荐

