如何用TensorFlow便捷加载多标签图像数据集(无需CSV)
多标签图像分类的简便数据集加载方案
针对你当前的数据集结构(同一张图片存于多个类别目录),以下是几种比.csv文件更简便的加载方式:
方法1:基于tf.data.Dataset构建自定义加载逻辑
无需额外维护标签文件,直接通过遍历目录收集图像与对应多标签:
- 遍历所有类别目录,记录每个图像路径对应的类别标签
- 对重复的图像路径合并标签,生成多标签掩码(比如类别共N个,标签为长度N的二进制数组)
- 用
tf.data.Dataset封装图像路径与多标签,再映射加载图像的逻辑
示例代码:
import tensorflow as tf import os # 数据集根目录,子目录为类别名 data_root = "/path/to/your/dataset" class_names = sorted(os.listdir(data_root)) class_indices = {name: idx for idx, name in enumerate(class_names)} num_classes = len(class_names) # 收集图像-标签映射 image_label_map = {} for class_name in class_names: class_dir = os.path.join(data_root, class_name) for img_name in os.listdir(class_dir): img_path = os.path.join(class_dir, img_name) if img_path not in image_label_map: image_label_map[img_path] = [0] * num_classes image_label_map[img_path][class_indices[class_name]] = 1 # 转换为tf.data.Dataset image_paths = list(image_label_map.keys()) labels = list(image_label_map.values()) def load_image(img_path, label): img = tf.io.read_file(img_path) img = tf.image.decode_jpeg(img, channels=3) # 根据图像格式调整 img = tf.image.resize(img, (224, 224)) # 根据模型输入调整尺寸 img = tf.keras.applications.resnet.preprocess_input(img) # 可选,根据模型调整预处理 return img, label dataset = tf.data.Dataset.from_tensor_slices((image_paths, labels)) dataset = dataset.map(load_image, num_parallel_calls=tf.data.AUTOTUNE) # 后续可添加shuffle、batch等操作 dataset = dataset.shuffle(1000).batch(32).prefetch(tf.data.AUTOTUNE)
方法2:调整数据集结构(避免重复存储图像)
如果可以修改数据集存储方式,推荐只保留一份图像,为每个图像创建对应的标签文件(比.csv更轻量化):
- 所有图像放在同一个
images目录下 - 每个图像对应一个同名的
.txt文件(比如cat_dog.jpg对应cat_dog.txt),文件内写入所属类别名(每行一个) - 加载时通过图像文件名匹配标签文件,生成多标签掩码
示例代码:
import tensorflow as tf import os data_dir = "/path/to/your/dataset" images_dir = os.path.join(data_dir, "images") labels_dir = os.path.join(data_dir, "labels") class_names = sorted(os.listdir(labels_dir)) # 假设labels目录下的子目录是类别,或者直接定义类别列表 class_indices = {name: idx for idx, name in enumerate(class_names)} num_classes = len(class_names) def load_image_with_label(img_path): # 获取图像文件名,匹配标签文件 img_name = os.path.basename(img_path) label_path = os.path.join(labels_dir, f"{os.path.splitext(img_name)[0]}.txt") # 读取标签文件 with open(label_path, "r") as f: assigned_classes = f.read().splitlines() # 生成多标签掩码 label = [0] * num_classes for cls in assigned_classes: label[class_indices[cls]] = 1 # 加载并预处理图像 img = tf.io.read_file(img_path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (224, 224)) return img, tf.convert_to_tensor(label, dtype=tf.float32) # 构建数据集 image_paths = tf.data.Dataset.list_files(os.path.join(images_dir, "*.jpg")) # 根据图像格式调整 dataset = image_paths.map(lambda x: tf.py_function(load_image_with_label, [x], [tf.float32, tf.float32])) dataset = dataset.shuffle(1000).batch(32).prefetch(tf.data.AUTOTUNE)
方法3:复用image_dataset_from_directory并修正标签
如果你不想改动现有目录结构,可以先通过image_dataset_from_directory加载单标签数据集,再根据图像文件名重新匹配多标签:
- 先加载单标签数据集,获取每个样本的文件名与对应单标签
- 遍历所有类别目录,构建文件名到多标签的映射
- 在数据集的映射函数中替换原单标签为多标签
示例代码:
import tensorflow as tf import os data_root = "/path/to/your/dataset" class_names = sorted(os.listdir(data_root)) class_indices = {name: idx for idx, name in enumerate(class_names)} num_classes = len(class_names) # 先加载单标签数据集 single_label_ds = tf.keras.utils.image_dataset_from_directory( data_root, image_size=(224, 224), batch_size=32, label_mode='int' ) # 构建文件名到多标签的映射 image_label_map = {} for class_name in class_names: class_dir = os.path.join(data_root, class_name) for img_name in os.listdir(class_dir): base_name = os.path.basename(img_name) if base_name not in image_label_map: image_label_map[base_name] = [0] * num_classes image_label_map[base_name][class_indices[class_name]] = 1 # 修正标签 def fix_labels(img, single_label): # 获取图像文件名(从tf.Tensor中提取) img_path = single_label_ds.file_paths[tf.argmax(tf.cast(tf.equal(single_label_ds.file_paths, img), tf.int32))] base_name = os.path.basename(img_path.numpy().decode()) multi_label = image_label_map[base_name] return img, tf.convert_to_tensor(multi_label, dtype=tf.float32) # 注意:这里需要用tf.py_function处理numpy操作 multi_label_ds = single_label_ds.map(lambda img, lbl: tf.py_function(fix_labels, [img, lbl], [tf.float32, tf.float32])) multi_label_ds = multi_label_ds.prefetch(tf.data.AUTOTUNE)
以上三种方法都无需维护.csv文件,其中方法1最适配你当前的重复目录结构,方法2是长期更合理的多标签数据集存储方式,方法3则是最小改动现有结构的折中方案。
内容的提问来源于stack exchange,提问作者aeloan
相关产品推荐
相关产品推荐

