如何降低内存占用?生成预测pickle文件时反复内存溢出求助
解决内存不足问题:批量处理+内存优化方案
我太懂这种跑了6小时突然崩掉的挫败感了!从你给出的代码片段来看,核心问题大概率是你把所有图像数据、预测结果都一股脑存在内存里,攒到最后才写入pickle文件——随着数据量累积,内存被撑爆是必然的。下面给你几个实用的解决思路:
1. 分批处理,增量保存结果
不要把所有结果都存在列表里等最后统一存,而是每处理一小批就写入一次。因为pickle本身不支持直接追加,我们可以用两种方式实现:
方式A:拆分多个小pickle文件
处理完一批就存一个单独的pickle,后续如果需要合并再统一处理。示例代码:
from keras.models import load_model import sys import pickle import os import cv2 import glob import gc import numpy as np sys.setrecursionlimit(10000) model = load_model('你的模型路径.h5') # 记得补充你的模型路径 batch_size = 100 # 根据你的内存大小调整批次 batch_results = [] batch_count = 0 imgdirs = os.listdir('/chars/') imgdirs.sort(key=float) for imgdir in imgdirs: for imgfile in glob.glob(os.path.join('/chars/', imgdir, '*.png')): img = cv2.imread(imgfile) # 这里补充你的图像预处理步骤(比如resize、归一化、转灰度等) processed_img = cv2.resize(img, (28, 28)) / 255.0 # 示例预处理 batch_results.append(processed_img) # 达到批次大小就预测并保存 if len(batch_results) >= batch_size: predictions = model.predict(np.array(batch_results), verbose=0) # 保存这批结果 with open(f'predictions_batch_{batch_count}.pkl', 'wb') as f: pickle.dump(predictions, f, protocol=pickle.HIGHEST_PROTOCOL) # 清空列表+手动回收内存 batch_results.clear() del predictions gc.collect() batch_count += 1 # 处理最后一批不足batch_size的数据 if batch_results: predictions = model.predict(np.array(batch_results), verbose=0) with open(f'predictions_batch_{batch_count}.pkl', 'wb') as f: pickle.dump(predictions, f, protocol=pickle.HIGHEST_PROTOCOL)
方式B:用joblib替代pickle(更适合大数据)
joblib对numpy数组的序列化效率更高,而且支持压缩,能进一步减少内存和磁盘占用:
from joblib import dump # 预测后直接分批写入 for batch in processed_batches: predictions = model.predict(batch) dump(predictions, f'predictions_batch_{batch_count}.joblib', compress=3)
2. 优化图像加载,减少内存占用
- 加载图像后立即释放原始图内存:比如处理完
processed_img后,直接del img再调用gc.collect()手动回收。 - 用生成器加载图像:不要一次性把所有图像读进列表,而是用生成器逐个返回处理后的图像,内存里永远只存一张/一批图像:
import itertools def image_generator(imgdirs): for imgdir in imgdirs: for imgfile in glob.glob(os.path.join('/chars/', imgdir, '*.png')): img = cv2.imread(imgfile) processed_img = cv2.resize(img, (28, 28)) / 255.0 # 示例预处理 yield processed_img del img gc.collect() # 用生成器分批预测 gen = image_generator(imgdirs) batch_count = 0 while True: batch = list(itertools.islice(gen, batch_size)) if not batch: break predictions = model.predict(np.array(batch)) # 保存这批结果 with open(f'predictions_batch_{batch_count}.pkl', 'wb') as f: pickle.dump(predictions, f, protocol=pickle.HIGHEST_PROTOCOL) batch_count += 1
3. 其他小技巧
- 降低图像精度:把图像数组从
float64改成float32(大多数模型不需要64位精度),灰度图直接用单通道而非三通道。 - 关闭冗余输出:预测时用
verbose=0,避免打印大量日志占用内存。 - 排查内存泄漏:如果还是有问题,可以用
memory_profiler工具定位哪行代码在持续占用内存。
另外提一句:imgdirs.sort(key=float)如果遇到非数字的目录名会报错,建议加个过滤或者异常处理,避免中途崩溃。
内容的提问来源于stack exchange,提问作者Akhil
相关产品推荐
相关产品推荐

