如何使用image_dataset_from_directory提取MINC数据集指定类别子集?
解决
image_dataset_from_directory获取指定类别子集的问题 问题原因
你之前的代码报错是因为class_names参数的作用是重命名所有类别,而非筛选类别——它要求传入的列表必须和目标目录下的所有子目录名称数量完全匹配,不能只传子集,所以直接传["wood", "foliage", "glass", "hair"]会触发不匹配的错误。
无需修改文件夹的解决方案
有两种可行的方法,不用删除或移动原数据集文件夹:
方法1:加载完整数据集后过滤
先加载整个MINC数据集,再通过filter方法筛选出目标类别的样本,步骤如下:
- 加载完整数据集
import tensorflow as tf folder_dir = "你的MINC数据集路径" full_dataset = tf.keras.preprocessing.image_dataset_from_directory( folder_dir, label_mode="categorical", # 可添加其他参数比如image_size、batch_size等 ) class_names = full_dataset.class_names
- 确定目标类别对应的索引
target_classes = ["wood", "foliage", "glass", "hair"] # 获取每个目标类别在完整类别列表中的索引 target_indices = [class_names.index(cls) for cls in target_classes]
- 过滤数据集
针对categorical格式的标签,定义过滤函数,保留目标索引对应的样本:
def filter_target_classes(image, label): # 将one-hot标签转换为类别索引 label_idx = tf.argmax(label) # 判断当前样本的类别是否在目标索引列表中 return tf.reduce_any(tf.equal(label_idx, target_indices)) # 应用过滤 filtered_dataset = full_dataset.filter(filter_target_classes)
- 更新数据集的class_names属性(可选)
方便后续查看类别名称:
filtered_dataset.class_names = target_classes
方法2:手动构建数据集(更灵活)
直接遍历目标类别的文件夹,收集文件路径和标签,再构建tf.data.Dataset,适合需要自定义加载逻辑的场景:
import tensorflow as tf import os folder_dir = "你的MINC数据集路径" target_classes = ["wood", "foliage", "glass", "hair"] # 收集目标类别的文件路径和标签 file_paths = [] labels = [] # 给目标类别分配新的索引(从0开始) class_to_idx = {cls: idx for idx, cls in enumerate(target_classes)} for cls_name in target_classes: cls_dir = os.path.join(folder_dir, cls_name) if not os.path.isdir(cls_dir): continue # 遍历当前类别下的所有图片文件 for fname in os.listdir(cls_dir): if fname.lower().endswith(('.png', '.jpg', '.jpeg')): file_paths.append(os.path.join(cls_dir, fname)) labels.append(class_to_idx[cls_name]) # 定义图片加载和预处理函数 def load_and_preprocess_image(file_path, label): # 读取图片 img = tf.io.read_file(file_path) # 根据实际图片格式解码(这里以JPEG为例) img = tf.image.decode_jpeg(img, channels=3) # 调整图片尺寸(根据你的需求修改) img = tf.image.resize(img, (224, 224)) # 可选:添加预处理步骤(比如归一化) img = tf.keras.applications.resnet50.preprocess_input(img) # 转换为categorical标签(如果需要) label = tf.one_hot(label, depth=len(target_classes)) return img, label # 构建数据集 dataset = tf.data.Dataset.from_tensor_slices((file_paths, labels)) dataset = dataset.map(load_and_preprocess_image, num_parallel_calls=tf.data.AUTOTUNE) # 打乱、分批、预取 dataset = dataset.shuffle(buffer_size=1000).batch(32).prefetch(tf.data.AUTOTUNE)
总结
- 方法1更简洁,基于
image_dataset_from_directory的结果直接过滤,适合快速实现; - 方法2灵活性更高,可自定义加载、预处理逻辑,适合复杂场景。
内容的提问来源于stack exchange,提问作者Barnacle Own
相关产品推荐
相关产品推荐

