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

如何从磁盘向Keras流式传输输入与标签图像以避免内存溢出

解决方案:自定义Keras Sequence数据生成器

对于你的图像外扩任务,flow_from_directory确实不适用,它是为分类任务设计的。你可以通过自定义Keras Sequence类实现输入和标签图像的流式加载,这种方式会按需加载批次数据,彻底避免内存溢出问题。

实现思路

  1. 继承keras.utils.Sequence类,实现三个核心方法:

    • __init__:初始化参数,传入输入/标签目录、图像尺寸、批次大小等配置
    • __len__:计算每个epoch包含的批次数量
    • __getitem__:加载并返回单个批次的输入图像与对应标签图像
  2. 核心前提:确保输入目录和标签目录中的图像文件名一一对应(比如input/img001.jpg对应label/img001.jpg),保证输入与标签正确配对。

完整代码示例

import os
import numpy as np
from keras.models import Sequential
from keras.layers import Conv2D
from keras.utils import Sequence
from PIL import Image

# 定义图像基础参数
img_width = 400
img_height = 300
channels = 3

# 自定义流式数据生成器
class OutpaintDataGenerator(Sequence):
    def __init__(self, input_dir, label_dir, batch_size=32, img_size=(300,400), shuffle=True):
        self.input_dir = input_dir
        self.label_dir = label_dir
        self.batch_size = batch_size
        self.img_size = img_size
        self.shuffle = shuffle
        # 筛选目录下的文件(排除子文件夹)
        self.filenames = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))]
        self.on_epoch_end()

    def __len__(self):
        # 计算每个epoch的批次总数
        return int(np.ceil(len(self.filenames) / self.batch_size))

    def __getitem__(self, index):
        # 获取当前批次的文件名列表
        batch_filenames = self.filenames[index*self.batch_size : (index+1)*self.batch_size]
        
        # 初始化批次数据数组
        batch_input = np.zeros((len(batch_filenames), *self.img_size, channels), dtype=np.float32)
        batch_label = np.zeros((len(batch_filenames), *self.img_size, channels), dtype=np.float32)
        
        # 加载并预处理单批次图像
        for i, filename in enumerate(batch_filenames):
            input_path = os.path.join(self.input_dir, filename)
            label_path = os.path.join(self.label_dir, filename)
            
            # 加载图像并统一尺寸(resize参数为(width, height))
            input_img = Image.open(input_path).convert('RGB').resize(self.img_size[::-1])
            label_img = Image.open(label_path).convert('RGB').resize(self.img_size[::-1])
            
            # 归一化到[0,1]区间
            batch_input[i] = np.array(input_img) / 255.0
            batch_label[i] = np.array(label_img) / 255.0
        
        return batch_input, batch_label

    def on_epoch_end(self):
        # 每个epoch结束后打乱数据顺序,提升训练效果
        if self.shuffle:
            np.random.shuffle(self.filenames)

# 定义模型架构
model = Sequential()
model.add(Conv2D(32, (3, 3), activation='relu', padding='same', input_shape=(img_height, img_width, channels)))
model.add(Conv2D(64, (3, 3), activation='relu', padding='same'))
model.add(Conv2D(128, (3, 3), activation='relu', padding='same'))
model.add(Conv2D(64, (3, 3), activation='relu', padding='same'))
model.add(Conv2D(channels, (3, 3), activation='sigmoid', padding='same'))

# 编译模型
model.compile(optimizer='adam', loss='mse')

# 初始化训练数据生成器
train_generator = OutpaintDataGenerator(
    input_dir="input/",
    label_dir="label/",
    batch_size=32,
    img_size=(img_height, img_width),
    shuffle=True
)

# 启动训练
model.fit(
    train_generator,
    epochs=10
    # 若有验证集,可添加validation_data=val_generator
)

关键说明

  • 内存友好:仅加载当前训练批次的图像,不会一次性把2万张图存入内存,彻底解决内存溢出问题。
  • 灵活性强:可以在__getitem__方法中添加任意预处理逻辑,比如数据增强、格式转换、异常图像过滤等。
  • 文件名配对:如果输入与标签文件名不匹配,可通过维护CSV映射表的方式,在__init__中读取配对关系替代直接取文件名列表。
  • 多进程加速:在model.fit中设置workers参数,可启用多进程加载数据,提升训练效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 08:24:57