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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 19:23:10