如何在不下载数据集的情况下获取torchvision.datasets.ImageNet类别列表?
获取ImageNet类别列表无需下载完整数据集
嘿,这个问题我太懂了!不用啃下几个G的ImageNet数据集,照样能拿到它的类别列表,给你两个简单好用的方法:
方法一:借助预训练模型权重的元信息
很多针对ImageNet训练的预训练模型(比如ResNet、VGG这些),它们的权重对象里已经内置了类别名称列表,完全不用碰数据集。举个ResNet50的例子:
from torchvision.models import resnet50, ResNet50_Weights # 获取默认的预训练权重 weights = ResNet50_Weights.DEFAULT # 从权重的元数据里提取类别列表 class_names = weights.meta["categories"] # 打印前5个类别看看效果 print(class_names[:5]) # 输出会像 ['tench', 'goldfish', 'great white shark', 'tiger shark', 'hammerhead']
如果你的环境里已经下载过预训练权重,连权重都不用重新下,直接就能拿到列表。
方法二:单独下载ImageNet元数据文件
ImageNet的类别对应关系其实存放在一个小巧的JSON文件里,我们可以直接下载这个文件,完全不用管庞大的图片数据集:
from torchvision.datasets.utils import download_url import json # 这个是torchvision内部使用的元数据文件地址 meta_file_url = "https://raw.githubusercontent.com/pytorch/vision/main/torchvision/datasets/imagenet_class_index.json" # 下载到当前目录(你也可以指定其他路径) download_url(meta_file_url, ".") # 读取并转换为类别列表 with open("imagenet_class_index.json", "r") as f: class_index = json.load(f) # 按顺序整理成和模型输出对应的类别名称 class_names = [class_index[str(i)][1] for i in range(len(class_index))] print(class_names[:5]) # 同样能得到正确的类别列表
确实,早期PyTorch版本里没有直接提供这个便捷方式,但现在torchvision已经把这些细节都处理好了,上面两种方法都能让你轻松拿到类别列表,完全不用下载整个ImageNet数据集。
内容的提问来源于stack exchange,提问作者irudyak
相关产品推荐
相关产品推荐

