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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 20:25:53