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

如何通过PyTorch DataLoader获取自定义数据集中图像的类别名称?

实现方法

首先说明:如果你的dataset_train是用torchvision.datasets.ImageFolder创建的(刚好适配你按文件夹分类的数据集结构),它自带类别与索引的映射关系,直接复用即可:

  1. 先构建索引到类别名称的映射字典:
# 反转ImageFolder自带的class_to_idx(结构为{类别名:索引}),得到索引到类别名的映射
idx_to_class = {v: k for k, v in dataset_train.class_to_idx.items()}
  1. 遍历DataLoader,将每个样本的标签索引转换为类别名称并打印:
for images_batch, labels_batch in data_loader_train:
    # 逐个处理当前batch中的标签
    for label in labels_batch:
        # 把张量类型的标签转为整数,再映射到对应类别名称
        class_name = idx_to_class[label.item()]
        print(f"该图像所属类别: {class_name}")

如果你的dataset_train是自定义Dataset类,需要手动维护类别映射:

  • 在自定义Dataset的初始化方法中,保存类别名称列表或idx_to_class字典(比如读取根目录下的文件夹名作为类别名)
  • 遍历DataLoader时,用同样的方式将标签索引对应到类别名称即可。

自定义Dataset示例片段:

from torch.utils.data import Dataset
import os

class CustomDataset(Dataset):
    def __init__(self, root_dir):
        self.root_dir = root_dir
        # 获取所有类别文件夹名称并排序
        self.class_names = sorted(os.listdir(root_dir))
        # 构建索引到类别名称的映射
        self.idx_to_class = {i: name for i, name in enumerate(self.class_names)}
        # 其他初始化逻辑(比如收集所有图像路径)...
    
    def __getitem__(self, idx):
        # 图像加载与预处理逻辑...
        image = ... 
        # 获取当前图像对应的标签索引
        label_idx = ...
        return image, label_idx

之后在遍历DataLoader时,直接用自定义Dataset实例的idx_to_class属性转换标签即可。

内容的提问来源于stack exchange,提问作者Anik Chaudhuri

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 18:01:18