PyTorch ImageFolder类别索引与子文件夹名称映射查询及自定义类别混淆矩阵构建问题
如何查询PyTorch ImageFolder中类别名称对应的索引
嗨,这个问题我太熟了!PyTorch的ImageFolder其实已经内置了类别和索引的映射关系,完全不用瞎猜,直接用它的属性就能搞定,而且还能轻松用到你的混淆矩阵里。
1. 直接查看类别到索引的映射
当你用ImageFolder加载数据集后,它会自动生成一个class_to_idx属性,这是一个字典,键是你的子文件夹名称(也就是类别标签,比如"banana"),值是对应的索引(0-9)。
举个代码例子:
from torchvision.datasets import ImageFolder from torchvision.transforms import ToTensor # 替换成你的数据集根路径 dataset = ImageFolder(root="./your_dataset_root", transform=ToTensor()) # 打印类别与索引的对应关系 print(dataset.class_to_idx)
运行后你会得到类似这样的输出:
{'banana': 0, 'cucumber': 1, 'orange': 2, ...}
注意:ImageFolder是按子文件夹名称的字典序分配索引的,所以如果有比"banana"字母顺序更早的文件夹,它的索引会排在前面,这也是为什么不要假设"banana"一定是0的原因,直接查这个属性最靠谱。
2. 反过来,从索引查类别名称
如果需要根据索引获取对应的类别名称(比如混淆矩阵里要把索引转成标签),可以用dataset.classes列表——这个列表的第i个元素,就是索引i对应的类别名称:
# 获取索引0对应的类别 print(dataset.classes[0]) # 输出比如'banana'
3. 把类别名称用到混淆矩阵上
结合你的需求,画混淆矩阵的时候直接用dataset.classes作为坐标轴标签就行,不用手动映射。这里用seaborn举个简单的例子:
import matplotlib.pyplot as plt import seaborn as sns import numpy as np # 假设你已经得到了模型输出的混淆矩阵(10x10的数组) confusion_matrix = np.random.randint(0, 100, (10, 10)) plt.figure(figsize=(10, 8)) # 用dataset.classes作为坐标轴标签 sns.heatmap(confusion_matrix, annot=True, fmt="d", xticklabels=dataset.classes, yticklabels=dataset.classes) plt.xlabel("Predicted Category") plt.ylabel("True Category") plt.title("Custom Confusion Matrix") plt.show()
这样你的混淆矩阵坐标轴就会显示"banana"、"orange"这些实际类别名称,完美符合你的需求。
内容的提问来源于stack exchange,提问作者Imahn
相关产品推荐
相关产品推荐

