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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 00:52:52