PyTorch中如何获取CIFAR100的粗标签与细标签?默认target对应哪种标签?
PyTorch中CIFAR100数据集的细标签与粗标签获取方法
一、默认target的类型
PyTorch官方的torchvision.datasets.CIFAR100默认返回的target是细标签(fine label),对应数据集里的100个具体子类。
二、同时获取细标签和粗标签的方法
CIFAR100数据集实例本身内置了两个核心属性:
fine_labels:存储所有样本的细标签列表coarse_labels:存储所有样本的粗标签列表(对应20个超类,每个超类包含5个细类)
方法1:直接访问数据集属性
无需修改数据集结构,直接通过索引获取对应标签:
import torchvision.datasets as datasets # 加载训练集 train_data = datasets.CIFAR100(root="./data", train=True, download=True) # 加载测试集 test_data = datasets.CIFAR100(root="./data", train=False, download=True) # 获取第0个样本的标签 sample_idx = 0 fine_label = train_data.fine_labels[sample_idx] coarse_label = train_data.coarse_labels[sample_idx] print(f"细标签ID: {fine_label}, 粗标签ID: {coarse_label}")
方法2:自定义数据集包装类
如果需要在DataLoader迭代时直接拿到细、粗标签,可以自定义子类重写__getitem__方法:
import torchvision.datasets as datasets from torch.utils.data import DataLoader class CIFAR100WithCoarse(datasets.CIFAR100): def __getitem__(self, index): img, fine_label = super().__getitem__(index) coarse_label = self.coarse_labels[index] return img, fine_label, coarse_label # 初始化自定义数据集 train_dataset = CIFAR100WithCoarse(root="./data", train=True, download=True) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True) # 迭代获取批量数据 for imgs, fine_labels, coarse_labels in train_loader: print(f"批量细标签示例: {fine_labels[:3]}") print(f"批量粗标签示例: {coarse_labels[:3]}") break
内容的提问来源于stack exchange,提问作者JobHunter69
相关产品推荐
相关产品推荐

