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

使用Keras ImageDataGenerator时出现MemoryError的解决方法咨询

解决Keras训练第4轮触发MemoryError的方案

看起来你遇到的是典型的内存泄漏积累导致的溢出问题——前3轮内存慢慢攒着没爆,到第4轮终于顶不住了。结合你的代码和运行环境,我给你几个针对性的解决办法:

1. 修复生成器的内存浪费问题

你的生成器代码里,每次迭代都在创建新的train_imgs_batch和train_masks_batch数组,这些临时数组如果没被垃圾回收,会在每一轮训练后堆积起来。改成提前初始化一次,循环复用就能解决这个问题:

def generator(train_imgs, train_masks=None, batch_size=None):
    # 提前初始化batch数组,只创建一次,指定dtype避免内存浪费
    train_imgs_batch = np.zeros((batch_size, y_to_res, x_to_res, bands), dtype=np.uint8)
    train_masks_batch = np.zeros((batch_size, y_to_res, x_to_res, 1), dtype=np.uint8) if train_masks is not None else None
    
    while True:
        # 随机抽取batch对应的索引
        idx = np.random.choice(len(train_imgs), batch_size, replace=False)
        
        # 填充batch数据
        for i in range(batch_size):
            train_imgs_batch[i] = train_imgs[idx[i]]
            if train_masks is not None:
                # 给掩码加一个通道维度,匹配模型输入要求
                train_masks_batch[i] = train_masks[idx[i]][..., np.newaxis]
        
        # 归一化(如果你的模型需要的话,记得转成float32)
        train_imgs_batch_normalized = train_imgs_batch.astype(np.float32) / 255.0
        
        if train_masks is not None:
            yield train_imgs_batch_normalized, train_masks_batch
        else:
            yield train_imgs_batch_normalized

2. 优化memmap的使用模式

你现在用的是mode='r+',如果不需要修改data.npy和mask.npy里的数据,改成只读模式mode='r',这样numpy不会把整个数组加载到内存,而是按需从磁盘读取,能节省大量主机内存:

training_imgs = np.lib.format.open_memmap(filename=os.path.join(data_path, 'data.npy'), mode='r')
training_masks = np.lib.format.open_memmap(filename=os.path.join(data_path, 'mask.npy'), mode='r')

3. 让TensorFlow按需分配GPU显存

g2.2xlarge的GPU只有4GB显存,TensorFlow默认会一次性占满所有显存,哪怕当前训练用不上。添加这段代码限制显存增长,避免显存溢出(有时候显存不足也会伪装成主机内存的MemoryError):

import tensorflow as tf
from keras.backend.tensorflow_backend import set_session

# 在创建模型前执行这段配置
config = tf.ConfigProto()
config.gpu_options.allow_growth = True  # 按需分配显存
set_session(tf.Session(config=config))

4. 手动触发垃圾回收

在每个epoch结束后,强制清理未被回收的临时变量,避免内存积累。可以写一个简单的回调函数:

import gc
import keras

class GarbageCollectCallback(keras.callbacks.Callback):
    def on_epoch_end(self, epoch, logs=None):
        gc.collect()
        print(f"Epoch {epoch+1} finished, garbage collected.")

然后在fit_generator里加入这个回调:

dl_model.fit_generator(
    generator(training_imgs, training_masks, batch_size),
    steps_per_epoch=(len(training_imgs)//batch_size),  # 用整数除法避免浮点型steps
    epochs=4,
    verbose=1,
    callbacks=[model_checkpoint, GarbageCollectCallback()]
)

为什么前3轮正常第4轮爆内存?

每一轮训练后,模型的历史梯度、生成器的临时数组、甚至TensorFlow的中间计算图都会占用一些内存,这些内存如果没被及时回收,会一轮轮积累。g2.2xlarge的主机内存是15GB,前3轮积累的内存还没到上限,第4轮就触发了溢出。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 04:03:29