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

如何向Keras模型输入大规模图像数据集?

大规模图像数据集训练:无需全量载入内存的解决方案

我有一个包含数十万张带标签图像的训练数据集,要训练ML模型,但没法像这样创建numpy数组存储所有图像:

all_images = np.zeros(shape=(500000, 256, 256, 3), dtype="uint8")

大型企业肯定不是靠超大内存处理大规模数据集训练的,所以想知道怎么不用把整个数据集载入内存就能调用model.fit()完成训练?

当前的图像加载函数如下(简化后,移除了部分np.transpose()调用,直接复制可能无法运行):

def load_images(images: list):
    # 创建空数组存储n张256x256的3通道RGB图像
    resized_images = np.zeros(shape=(len(images), 256, 256, 3), dtype="uint8")

    index = 0
    for image in images:
        print(index)

        # 用cv2加载图像(注:此处存在错误,应改为cv2.imread(image))
        img = cv2.imread(images)

        # 调整图像尺寸到256x256
        img = cv2.resize(img, dsize=(256, 256))

        # 将图像加入数组
        resized_images[index] = img

        index += 1
    return resized_images

这个函数的目标是调整训练图像尺寸并载入单个numpy数组,传入model.fit()。另外我之前尝试保存加载模型继续训练没成功(加载后模型属性不全),如果这个方案可行,也请分享具体方法。


一、流式数据加载方案(无需全量载入内存)

1. 使用ImageDataGenerator配合flow_from_directory

如果数据集按分类文件夹存储(如train/class1/xxx.jpg、train/class2/xxx.jpg),这是最便捷的方式:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 定义数据预处理(仅做像素缩放,可按需添加数据增强)
datagen = ImageDataGenerator(rescale=1./255)

# 从文件夹流式加载,每次返回一个batch的数据
train_generator = datagen.flow_from_directory(
    'path/to/train_dir',
    target_size=(256, 256),  # 统一图像尺寸
    batch_size=32,  # 每次加载的图像数量
    class_mode='categorical'  # 分类任务选categorical/binary,回归任务选None
)

# 直接用generator训练模型
model.fit(
    train_generator,
    epochs=10,
    steps_per_epoch=train_generator.samples // train_generator.batch_size
)

2. 使用tf.data.Dataset自定义加载逻辑

如果数据集有单独的路径列表和标签列表,用这个方法灵活性更高:

import tensorflow as tf

def load_and_preprocess_image(image_path, label):
    # 读取图像文件
    img = tf.io.read_file(image_path)
    # 解码为RGB图像(灰度图用tf.image.decode_grayscale)
    img = tf.image.decode_jpeg(img, channels=3)
    # 调整尺寸
    img = tf.image.resize(img, (256, 256))
    # 像素值缩放到0-1
    img = tf.cast(img, tf.float32) / 255.0
    return img, label

# 假设你有图像路径列表和对应标签列表
image_paths = ['path/to/img1.jpg', 'path/to/img2.jpg', ...]
labels = [0, 1, 0, ...]

# 构建流式数据集
dataset = tf.data.Dataset.from_tensor_slices((image_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)

# 训练模型
model.fit(dataset, epochs=10)

3. 自定义Python生成器

若不依赖TensorFlow工具,可手动实现符合Keras要求的生成器:

import numpy as np
import cv2

def image_generator(image_paths, labels, batch_size=32):
    num_samples = len(image_paths)
    while True:  # 生成器需无限循环,供Keras重复调用
        # 每个epoch打乱数据顺序
        indices = np.random.permutation(num_samples)
        for i in range(0, num_samples, batch_size):
            batch_indices = indices[i:i+batch_size]
            batch_images = []
            batch_labels = []
            for idx in batch_indices:
                # 加载并预处理图像
                img = cv2.imread(image_paths[idx])
                img = cv2.resize(img, (256, 256))
                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)  # cv2默认BGR,转RGB
                img = img.astype(np.float32) / 255.0
                batch_images.append(img)
                batch_labels.append(labels[idx])
            yield np.array(batch_images), np.array(batch_labels)

# 使用生成器训练
train_gen = image_generator(image_paths, labels, batch_size=32)
steps_per_epoch = len(image_paths) // 32
model.fit(train_gen, epochs=10, steps_per_epoch=steps_per_epoch)

二、模型保存与继续训练的正确方法

1. 保存整个模型(推荐,保留所有训练状态)

# 训练中保存模型(HDF5格式)
model.save('my_model.h5')
# 或保存为TensorFlow推荐的SavedModel格式
model.save('my_saved_model')

# 加载模型继续训练
from tensorflow.keras.models import load_model
loaded_model = load_model('my_model.h5')  # 或load_model('my_saved_model')

# 直接继续训练,指定上次训练的结束epoch
loaded_model.fit(dataset, epochs=10, initial_epoch=5)

此方法会保存模型结构、权重、优化器状态、损失函数等所有属性,加载后可无缝继续训练。

2. 仅保存权重(需保留模型结构代码)

# 保存权重
model.save_weights('my_model_weights.h5')

# 加载权重:需先重建与原模型完全一致的结构
model = build_my_model()  # 此处调用你定义模型的函数
model.load_weights('my_model_weights.h5')

# 重新编译模型(权重不包含优化器状态)
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
model.fit(dataset, epochs=10)

内容的提问来源于stack exchange,提问作者AR Osornio

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 09:46:05