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

如何用TensorFlow便捷加载多标签图像数据集(无需CSV)

多标签图像分类的简便数据集加载方案

针对你当前的数据集结构(同一张图片存于多个类别目录),以下是几种比.csv文件更简便的加载方式:

方法1:基于tf.data.Dataset构建自定义加载逻辑

无需额外维护标签文件,直接通过遍历目录收集图像与对应多标签:

  1. 遍历所有类别目录,记录每个图像路径对应的类别标签
  2. 对重复的图像路径合并标签,生成多标签掩码(比如类别共N个,标签为长度N的二进制数组)
  3. 用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加载单标签数据集,再根据图像文件名重新匹配多标签:

  1. 先加载单标签数据集,获取每个样本的文件名与对应单标签
  2. 遍历所有类别目录,构建文件名到多标签的映射
  3. 在数据集的映射函数中替换原单标签为多标签

示例代码:

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 14:38:17