如何使用TensorFlow的ImageDataGenerator加载数据集指定类别的子集
解决方案
Keras的ImageDataGenerator的flow_from_directory方法内置了classes参数,可直接实现指定类别加载的需求,无需手动筛选数据集目录。
使用方法
- 提前定义本次实验需要加载的类别列表,列表元素要和数据集根目录下对应类别的子文件夹名称完全一致(大小写、符号都要匹配)
- 在调用
flow_from_directory时传入classes参数即可,生成器会自动忽略其他未被指定的类别
修改后的代码示例
from google.colab import drive drive.mount('/content/gdrive') import pathlib data = pathlib.Path('/content/gdrive/My Drive/Data_Clatech256/2_categ_caltech') from keras.preprocessing.image import ImageDataGenerator # 定义本次实验要加载的类别,可根据实验需求动态修改 TARGET_CLASSES = ["类别1文件夹名", "类别2文件夹名", "你需要的其他类别"] train_datagen = ImageDataGenerator(rescale=1./255) test_datagen = ImageDataGenerator(rescale=1./255) train_generator = train_datagen.flow_from_directory( data, target_size=(150, 150), batch_size=3, class_mode = "categorical", classes=TARGET_CLASSES # 新增指定类别参数 ) # 可打印验证加载的类别是否正确 print("加载的类别及对应索引:", train_generator.class_indices)
注意事项
- 测试集生成器也要传入和训练集完全一致的
classes参数,保证两者的类别索引映射关系相同,避免训练/评估时标签不匹配 - 如果需要动态调整实验类别,只需修改
TARGET_CLASSES列表的内容即可,无需改动数据集结构
内容的提问来源于stack exchange,提问作者Sadaf Shafi
相关产品推荐
相关产品推荐

