如何用Keras image_dataset_from_directory筛选指定10个数据集类别?
我从Kaggle导入了包含15个类别的PlantVillage数据集,仅需保留其中10个番茄相关类别。当前使用以下代码加载全部数据:
image_size= 256 batch_size=8 channels=3 epochs=50 dataset = tf.keras.preprocessing.image_dataset_from_directory('/kaggle/input/plant-village/PlantVillage', seed=123, shuffle=True, image_size=(image_size,image_size), batch_size=batch_size) dataset.class_names
执行结果显示共找到20638个文件,分属15个类别:
['Pepper__bell___Bacterial_spot',
'Pepper__bell___healthy',
'Potato___Early_blight',
'Potato___Late_blight',
'Potato___healthy',
'Tomato_Bacterial_spot',
'Tomato_Early_blight',
'Tomato_Late_blight',
'Tomato_Leaf_Mold',
'Tomato_Septoria_leaf_spot',
'Tomato_Spider_mites_Two_spotted_spider_mite',
'Tomato__Target_Spot',
'Tomato__Tomato_YellowLeaf__Curl_Virus',
'Tomato__Tomato_mosaic_virus',
'Tomato_healthy']
期望仅保留以下10个类别:
desired_classes = [ 'Tomato_Bacterial_spot', 'Tomato_Early_blight', 'Tomato_Late_blight', 'Tomato_Leaf_Mold', 'Tomato_Septoria_leaf_spot', 'Tomato_Spider_mites_Two_spotted_spider_mite', 'Tomato__Target_Spot', 'Tomato__Tomato_YellowLeaf__Curl_Virus', 'Tomato__Tomato_mosaic_virus', 'Tomato_healthy' ]
有两种实用方法可以实现类别筛选:
方法一:加载数据时直接指定类别
利用image_dataset_from_directory的classes参数,在加载数据阶段就只读取目标类别,这种方法更高效,无需加载无关数据:
image_size= 256 batch_size=8 channels=3 epochs=50 desired_classes = [ 'Tomato_Bacterial_spot', 'Tomato_Early_blight', 'Tomato_Late_blight', 'Tomato_Leaf_Mold', 'Tomato_Septoria_leaf_spot', 'Tomato_Spider_mites_Two_spotted_spider_mite', 'Tomato__Target_Spot', 'Tomato__Tomato_YellowLeaf__Curl_Virus', 'Tomato__Tomato_mosaic_virus', 'Tomato_healthy' ] # 加载时指定目标类别 dataset = tf.keras.preprocessing.image_dataset_from_directory( '/kaggle/input/plant-village/PlantVillage', seed=123, shuffle=True, image_size=(image_size, image_size), batch_size=batch_size, classes=desired_classes # 新增参数指定要加载的类别 ) # 验证结果 print(dataset.class_names)
执行后,dataset仅包含你指定的10个番茄相关类别,类别顺序与desired_classes完全一致。
方法二:过滤已加载的完整数据集
如果已经加载了全部数据,可通过标签过滤保留目标类别:
# 获取目标类别的索引 class_indices = {name: idx for idx, name in enumerate(dataset.class_names)} desired_indices = [class_indices[name] for name in desired_classes] # 过滤数据集:仅保留标签属于目标索引的样本 filtered_dataset = dataset.filter( lambda x, y: tf.reduce_any(tf.equal(y, tf.constant(desired_indices))) ) # 更新类名(可选,保证后续使用时类名正确) filtered_dataset.class_names = desired_classes # 验证结果 print(filtered_dataset.class_names)
这种方法适合已加载数据后需要调整的场景,但效率低于方法一,因为会先加载所有数据再进行过滤。
内容的提问来源于stack exchange,提问作者basit khan

