如何将PyTorch文件夹结构的ImageNet数据集加载到TensorFlow Datasets?
加载已提取的ImageNet数据集到TensorFlow Dataset
方法一:使用tf.keras.utils.image_dataset_from_directory(推荐)
这个API专门适配类别作为子文件夹的图像数据集结构,能快速完成加载。
加载训练集与验证集
import tensorflow as tf # 替换为你的数据集实际路径 train_dir = "path/to/imagenet/train" val_dir = "path/to/imagenet/val" # 加载训练集 train_ds = tf.keras.utils.image_dataset_from_directory( train_dir, image_size=(224, 224), # 根据模型需求调整尺寸,如ResNet系列常用224x224 batch_size=32, label_mode="int" # 自动将子文件夹映射为整数标签 ) # 加载验证集 val_ds = tf.keras.utils.image_dataset_from_directory( val_dir, image_size=(224, 224), batch_size=32, label_mode="int" )
可选:映射整数标签到类别名称
如果需要将整数标签转回ImageNet的类别名称,可单独下载synset_words.txt标签文件后构建映射:
# 假设synset_words.txt放在当前目录 label_map = {} with open("synset_words.txt", "r") as f: for line in f: synset, name = line.strip().split(" ", 1) label_map[synset] = name # 获取数据集的类别顺序(对应整数标签) class_names = train_ds.class_names # 构建整数标签到类别名称的映射表 int_to_name = {i: label_map[class_names[i]] for i in range(len(class_names))}
方法二:手动构建tf.data.Dataset(自定义场景)
如果需要更灵活的预处理逻辑,可手动遍历文件路径构建数据集:
步骤1:获取所有图像文件路径
import glob import os train_files = glob.glob(os.path.join(train_dir, "*", "*.JPEG")) val_files = glob.glob(os.path.join(val_dir, "*", "*.JPEG"))
步骤2:定义图像解析与预处理函数
# 先获取训练集的类别名称列表 class_names = sorted([name for name in os.listdir(train_dir) if os.path.isdir(os.path.join(train_dir, name))]) def parse_image(file_path): # 从文件路径提取类别标签(子文件夹名) label = tf.strings.split(file_path, os.sep)[-2] # 将类别名映射为整数标签 label = tf.argmax(tf.equal(class_names, label)) # 加载并预处理图像 img = tf.io.read_file(file_path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (224, 224)) # 可选:根据模型需求添加预处理,如ResNet的标准化 img = tf.keras.applications.resnet.preprocess_input(img) return img, label
步骤3:构建并优化数据集
# 构建训练集 train_ds = tf.data.Dataset.from_tensor_slices(train_files) train_ds = train_ds.map(parse_image, num_parallel_calls=tf.data.AUTOTUNE) train_ds = train_ds.shuffle(buffer_size=10000).batch(32).prefetch(tf.data.AUTOTUNE) # 构建验证集 val_ds = tf.data.Dataset.from_tensor_slices(val_files) val_ds = val_ds.map(parse_image, num_parallel_calls=tf.data.AUTOTUNE) val_ds = val_ds.batch(32).prefetch(tf.data.AUTOTUNE)
注意事项
image_size需与你使用的模型输入尺寸匹配(如EfficientNet可选用224/240/260等尺寸)。- 可根据需求扩展预处理逻辑,比如添加数据增强、自定义归一化规则。
- 若出现内存占用过高问题,可调小
batch_size,prefetch与AUTOTUNE参数已默认提升加载效率。
内容的提问来源于stack exchange,提问作者Anna Christine
相关产品推荐
相关产品推荐

