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

Keras训练CNN时,单文件夹加JSON标签映射的数据集高效使用方法问询

高效处理单文件夹+JSON标签映射的Keras图像分类数据集

绝对有高效的方案!不用手动移动或复制文件(既费时间又占空间),下面几种方法都能直接适配你的数据集结构,完美对接Keras的训练流程:

方法1:用tf.data.Dataset自定义加载(最灵活,推荐)

tf.data是TensorFlow生态里最灵活的数据加载工具,能直接从文件名列表和标签映射构建数据集,还能轻松集成预处理、增强等操作。

步骤很清晰:

  • 加载JSON文件,获取文件名到标签的映射
  • 生成所有图像文件的完整路径
  • 用tf.data.Dataset构建数据集,映射图像读取和预处理逻辑

示例代码:

import tensorflow as tf
import json
import os

# 1. 加载标签映射
with open('labels.json', 'r') as f:
    label_map = json.load(f)

# 2. 生成图像文件路径和对应标签
image_dir = 'path/to/your/single/image/folder'
file_names = list(label_map.keys())
file_paths = [os.path.join(image_dir, fname) for fname in file_names]
labels = list(label_map.values())

# 3. 构建tf.data.Dataset
def load_and_preprocess_image(file_path, label):
    # 读取图像
    img = tf.io.read_file(file_path)
    img = tf.image.decode_jpeg(img, channels=3)  # 图像是png就用decode_png
    # 预处理:调整尺寸、归一化
    img = tf.image.resize(img, (224, 224))  # 替换成你的模型输入尺寸
    img = img / 255.0  # 归一化到[0,1]区间
    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)

# 4. 配置数据集:打乱、批处理、预取(提升训练速度)
batch_size = 32
dataset = dataset.shuffle(buffer_size=len(file_names)).batch(batch_size).prefetch(tf.data.AUTOTUNE)

# 之后直接传入model.fit()即可
# model.fit(dataset, epochs=10)

优点:速度快,支持并行处理,能无缝集成TensorFlow的图像增强API(比如tf.image.random_flip_left_right),适合大规模数据集。

方法2:自定义Keras Sequence类(内存友好,适配Keras习惯)

如果习惯Keras的生成器模式,或者数据集太大没法一次性加载到内存,自定义Sequence类是个好选择——它会按批次加载数据,内存占用低,还支持多线程。

示例代码:

from tensorflow.keras.utils import Sequence
import json
import os
import cv2
import numpy as np

class CustomImageSequence(Sequence):
    def __init__(self, image_dir, label_json, batch_size=32, img_size=(224,224)):
        self.image_dir = image_dir
        self.batch_size = batch_size
        self.img_size = img_size
        
        # 加载标签映射
        with open(label_json, 'r') as f:
            self.label_map = json.load(f)
        self.file_names = list(self.label_map.keys())
        self.labels = list(self.label_map.values())

    def __len__(self):
        # 返回总批次数
        return np.ceil(len(self.file_names) / self.batch_size).astype(int)

    def __getitem__(self, idx):
        # 获取第idx批次的数据
        batch_files = self.file_names[idx*self.batch_size : (idx+1)*self.batch_size]
        batch_labels = self.labels[idx*self.batch_size : (idx+1)*self.batch_size]
        
        # 加载并预处理图像
        batch_images = []
        for fname in batch_files:
            img_path = os.path.join(self.image_dir, fname)
            img = cv2.imread(img_path)
            img = cv2.resize(img, self.img_size)
            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)  # OpenCV默认BGR,转成RGB适配模型
            img = img / 255.0
            batch_images.append(img)
        
        return np.array(batch_images), np.array(batch_labels)

    def on_epoch_end(self):
        # 每个epoch结束后打乱数据,提升训练效果
        idx = np.random.permutation(len(self.file_names))
        self.file_names = [self.file_names[i] for i in idx]
        self.labels = [self.labels[i] for i in idx]

# 使用示例
train_sequence = CustomImageSequence(
    image_dir='path/to/your/image/folder',
    label_json='labels.json',
    batch_size=32,
    img_size=(224,224)
)

# 传入model.fit()
# model.fit(train_sequence, epochs=10)

优点:符合Keras用户的使用习惯,内存占用可控,适合中小规模数据集或者内存有限的环境。

方法3:创建软链接适配flow_from_directory(快速复用现有代码)

如果你已经写好了基于flow_from_directory的代码,不想改逻辑,可以用软链接(符号链接)快速构建符合要求的目录结构——不用复制文件,只是创建指向原文件的快捷方式,节省空间。

示例代码(Linux/macOS):

import json
import os

# 配置路径
source_dir = 'path/to/your/single/image/folder'
target_dir = 'path/to/training_directory'
label_json = 'labels.json'

# 加载标签映射
with open(label_json, 'r') as f:
    label_map = json.load(f)

# 创建目标目录和分类子目录
os.makedirs(target_dir, exist_ok=True)
class_ids = set(label_map.values())
for cid in class_ids:
    class_dir = os.path.join(target_dir, f'class_{cid}')
    os.makedirs(class_dir, exist_ok=True)

# 创建软链接
for fname, cid in label_map.items():
    src_path = os.path.join(source_dir, fname)
    dst_path = os.path.join(target_dir, f'class_{cid}', fname)
    # 避免重复创建链接
    if not os.path.exists(dst_path):
        os.symlink(src_path, dst_path)

Windows系统需要用mklink命令(需管理员权限),可把创建链接的代码改成:

import subprocess
# ... 前面代码相同 ...
for fname, cid in label_map.items():
    src_path = os.path.join(source_dir, fname)
    dst_path = os.path.join(target_dir, f'class_{cid}', fname)
    if not os.path.exists(dst_path):
        subprocess.run(['mklink', dst_path, src_path], shell=True, check=True)

之后你就可以直接用flow_from_directory加载数据了:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

datagen = ImageDataGenerator(rescale=1./255)
train_generator = datagen.flow_from_directory(
    target_dir,
    target_size=(224,224),
    batch_size=32,
    class_mode='categorical'  # 二分类就用'binary'
)

优点:完全复用现有flow_from_directory的代码,零修改训练逻辑,适合快速验证。


内容的提问来源于stack exchange,提问作者Davide Maran

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 20:27:36