Python循环保存Keras生成的CNN预测图像迭代逐渐变慢该如何解决
循环调用matplotlib savefig保存图像迭代速度变慢的原因及解决方案
问题原因
你遇到的速度衰减问题核心来自matplotlib的全局画布状态管理机制:每次调用plt.imshow()时,程序都会在当前持有的全局画布上叠加新的绘图对象,不会主动清理上一轮循环生成的内容。随着迭代次数增加,内存中堆积的绘图对象量线性上升,每一轮的绘制、保存操作需要处理的资源越来越多,最终出现单轮耗时暴涨的情况,严重时还会触发内存溢出。
解决方案
下面提供三种不同适配场景的修复方案:
方案1:最小改动适配,每次保存后清空画布
仅需在原有代码的plt.savefig后新增一行plt.clf()(清空当前画布所有内容)即可解决对象堆积问题,修改后代码如下:results = model.predict_generator(test_img_gen,len(os.listdir(child)),verbose=1) for i,img in tqdm(enumerate(results)): plt.imshow(np.reshape(img,(512,512)), interpolation='nearest') resultDir = '{}_{}_{}.png'.format(resdir,filenames[i],str(i)) plt.savefig(resultDir) plt.clf() # 清空当前画布,清除上一轮的绘图对象方案2:独立画布管理,资源释放更彻底
避免使用matplotlib的全局画布,每次循环独立创建画布实例,用完直接关闭,适配绘图逻辑更复杂的场景,不会出现全局状态污染问题:results = model.predict_generator(test_img_gen,len(os.listdir(child)),verbose=1) for i,img in tqdm(enumerate(results)): fig, ax = plt.subplots() # 新建独立画布 ax.imshow(np.reshape(img,(512,512)), interpolation='nearest') resultDir = '{}_{}_{}.png'.format(resdir,filenames[i],str(i)) fig.savefig(resultDir) plt.close(fig) # 关闭画布,释放所有关联资源方案3:轻量工具替代,性能提升最明显
如果你不需要matplotlib提供的坐标轴、画布留白等元素,仅需要把预测数组保存为图像,直接用PIL或者OpenCV等更轻量的图像处理库写入文件,整体性能会提升数倍,完全不会出现速度衰减问题,PIL实现示例如下:from PIL import Image import numpy as np results = model.predict_generator(test_img_gen,len(os.listdir(child)),verbose=1) for i,img in tqdm(enumerate(results)): img_arr = np.reshape(img,(512,512)) # 若预测输出已经是0-255范围的整数,可删除下面两行归一化代码 img_arr = (img_arr - img_arr.min()) / (img_arr.max() - img_arr.min()) * 255 img_arr = img_arr.astype(np.uint8) resultDir = '{}_{}_{}.png'.format(resdir,filenames[i],str(i)) Image.fromarray(img_arr).save(resultDir)
内容的提问来源于stack exchange,提问作者llllllllllx
相关产品推荐
相关产品推荐

