如何通过PyTorch DataLoader获取自定义数据集中图像的类别名称?
实现方法
首先说明:如果你的dataset_train是用torchvision.datasets.ImageFolder创建的(刚好适配你按文件夹分类的数据集结构),它自带类别与索引的映射关系,直接复用即可:
- 先构建索引到类别名称的映射字典:
# 反转ImageFolder自带的class_to_idx(结构为{类别名:索引}),得到索引到类别名的映射 idx_to_class = {v: k for k, v in dataset_train.class_to_idx.items()}
- 遍历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
相关产品推荐
相关产品推荐

