使用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
相关产品推荐
相关产品推荐

