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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 05:15:29