如何在PyTorch中从Cifar-100按类别抽取10%样本构建自定义训练集
PyTorch下CIFAR-100每类按比例抽取训练样本的高效实现
你完全不需要手动为每个类别单独定义变量存储索引,以下两种方法都可以自动适配任意类别数的数据集,代码量不会随类别数量增长:
方法1:通用字典分组实现(兼容性最好)
该方法不需要额外依赖,兼容所有PyTorch内置数据集以及自定义数据集:
import random from torch.utils.data import Subset from torchvision import datasets # 加载CIFAR-100训练集(你原有加载逻辑不变) trainset = datasets.CIFAR100(root='./data', train=True, download=True, transform=None) class_indices = {} # 直接读取数据集内置的标签列表,无需遍历每个样本读取标签 for idx, label in enumerate(trainset.targets): if label not in class_indices: class_indices[label] = [] class_indices[label].append(idx) selected_indices = [] sample_ratio = 0.1 for indices in class_indices.values(): # 计算单类采样数量,max避免样本数过少时采样数为0 take_count = max(1, int(len(indices) * sample_ratio)) # 若需要随机采样而非取前10%,打开下面这行打乱索引 # random.shuffle(indices) selected_indices.extend(indices[:take_count]) # 生成最终的小子集 small_trainset = Subset(trainset, selected_indices)
方法2:Numpy向量化实现(性能更高)
如果数据集规模较大,用Numpy的向量化操作可以进一步降低遍历耗时:
import numpy as np from torch.utils.data import Subset from torchvision import datasets trainset = datasets.CIFAR100(root='./data', train=True, download=True, transform=None) targets = np.array(trainset.targets) sample_ratio = 0.1 selected_indices = [] for label in np.unique(targets): # 批量获取该类别所有样本的索引 label_indices = np.where(targets == label)[0] # 若需要随机采样,用下面这行代替切片逻辑 # selected = np.random.choice(label_indices, size=max(1, int(len(label_indices)*sample_ratio)), replace=False) selected = label_indices[:max(1, int(len(label_indices)*sample_ratio))] selected_indices.extend(selected.tolist()) small_trainset = Subset(trainset, selected_indices)
注意事项
- 两种方法都可以直接复用在CIFAR-10、ImageNet等其他PyTorch数据集上,不需要修改核心逻辑
- 如果要保证采样无偏差,建议打开随机打乱索引的逻辑,避免取到的样本都是原始数据集同一分区的内容
内容的提问来源于stack exchange,提问作者Zixuan Liu
相关产品推荐
相关产品推荐

