使用os.walk()存储图像像素为pickle时速度过慢及内存问题求助
图像数据集序列化问题分析与优化方案
你好呀!看了你的问题和代码,我来帮你拆解下核心问题,再给你几个实用的优化思路~
你的核心困惑&遇到的问题
- 第一次运行代码一切正常,但删除
data.pickle后再次执行速度大幅变慢 - 放大10倍的完整数据集跑了1天就被系统终止,疑惑Python会不会在执行结束后还留着内存数据
先解答你的内存疑惑
Python执行结束后会自动释放所有占用的内存,第二次运行变慢和内存残留没关系~大概率是第一次运行时系统缓存了图像文件,第二次读盘时没有缓存加持,所以速度慢了不少。
代码里的关键问题
1. 内存爆炸的根源:一次性加载所有图像到内存
你把所有图像都存在images和groundtruths两个列表里,最后才转成NumPy数组。10倍规模的话,光是解码后的图像数据就会占满内存——比如一张15KB的JPG解码成256×256的RGB图,就占了约192KB,10000张就是1.9GB,再加10000张PNG真值图,内存直接飙到好几GB,系统最后只能强制终止进程。
2. 嵌套os.walk太冗余,拖慢效率
你的代码嵌套了三层os.walk,但其实你已经知道文件夹的固定结构(每个category下有input和groundtruth),完全不用递归遍历所有目录,多余的IO操作会拖慢整体速度。
3. 图像读取的潜在隐患
- 用
cv2.imread时没做异常判断,如果某张图损坏会返回None,直接append会导致后续数组转换失败 - 文件名替换逻辑
file_.replace('gt','in').replace('png','jpg')不够健壮,万一文件名格式有变化(比如多了其他前缀)就会出错
优化方案
方案1:分批次处理+增量序列化(解决内存过载)
不要一次性把所有图像加载到内存,而是分批次读入、序列化,每处理完一批就释放内存:
import os import pickle import cv2 import numpy as np PATH_DATA = 'data' def process_single_category(category_path): # 直接获取固定路径,不用递归遍历 input_dir = os.path.join(category_path, 'input') gt_dir = os.path.join(category_path, 'groundtruth') # 排序文件名保证顺序一致 gt_files = sorted(os.listdir(gt_dir)) for gt_file in gt_files: # 生成对应输入图像的文件名,加后缀匹配更严谨 input_file = gt_file.replace('gt', 'in').replace('.png', '.jpg') input_path = os.path.join(input_dir, input_file) gt_path = os.path.join(gt_dir, gt_file) # 读取图像并做异常校验 img = cv2.imread(input_path) gt_img = cv2.imread(gt_path) if img is not None and gt_img is not None: yield img, gt_img def main(): batch_size = 1000 # 每1000张图像为一批 batch_idx = 1 temp_images = [] temp_gts = [] # 遍历所有类别文件夹 for category in os.listdir(PATH_DATA): category_path = os.path.join(PATH_DATA, category) if not os.path.isdir(category_path): continue # 逐张处理当前类别的图像 for img, gt_img in process_single_category(category_path): temp_images.append(img) temp_gts.append(gt_img) # 攒够一批就序列化并清空临时列表释放内存 if len(temp_images) >= batch_size: batch_imgs = np.array(temp_images) batch_gts = np.array(temp_gts) with open(f'{PATH_DATA}/batch_{batch_idx}.pickle', 'wb') as f: pickle.dump((batch_imgs, batch_gts), f) temp_images.clear() temp_gts.clear() batch_idx += 1 print(f"已完成第 {batch_idx-1} 批次存储") # 处理剩余不足一批的图像 if temp_images: batch_imgs = np.array(temp_images) batch_gts = np.array(temp_gts) with open(f'{PATH_DATA}/batch_final.pickle', 'wb') as f: pickle.dump((batch_imgs, batch_gts), f) print("所有图像处理完成!") if __name__ == '__main__': main()
方案2:改用HDF5格式存储(更适合大规模数据集)
Pickle其实不太适合存储超大NumPy数组,HDF5(用h5py库)是专门为大规模数值数据设计的,支持分块存储、按需读取,不会一次性占满内存:
import os import cv2 import numpy as np import h5py PATH_DATA = 'data' def main(): # 先统计总图像数量,方便创建HDF5数据集 total_img_count = 0 for category in os.listdir(PATH_DATA): category_path = os.path.join(PATH_DATA, category) if not os.path.isdir(category_path): continue gt_dir = os.path.join(category_path, 'groundtruth') total_img_count += len(os.listdir(gt_dir)) # 读取一张样本图获取图像尺寸(假设所有图像尺寸一致) sample_category = next(cat for cat in os.listdir(PATH_DATA) if os.path.isdir(os.path.join(PATH_DATA, cat))) sample_gt_file = os.listdir(os.path.join(PATH_DATA, sample_category, 'groundtruth'))[0] sample_img = cv2.imread(os.path.join(PATH_DATA, sample_category, 'groundtruth', sample_gt_file)) img_shape = sample_img.shape # 创建HDF5文件并初始化数据集 with h5py.File(f'{PATH_DATA}/data.h5', 'w') as hf: images_ds = hf.create_dataset('images', shape=(total_img_count, *img_shape), dtype=np.uint8) gts_ds = hf.create_dataset('groundtruths', shape=(total_img_count, *img_shape), dtype=np.uint8) current_idx = 0 for category in os.listdir(PATH_DATA): category_path = os.path.join(PATH_DATA, category) if not os.path.isdir(category_path): continue input_dir = os.path.join(category_path, 'input') gt_dir = os.path.join(category_path, 'groundtruth') gt_files = sorted(os.listdir(gt_dir)) for gt_file in gt_files: input_file = gt_file.replace('gt', 'in').replace('.png', '.jpg') input_path = os.path.join(input_dir, input_file) gt_path = os.path.join(gt_dir, gt_file) img = cv2.imread(input_path) gt_img = cv2.imread(gt_path) if img is not None and gt_img is not None: images_ds[current_idx] = img gts_ds[current_idx] = gt_img current_idx += 1 if current_idx % 100 == 0: print(f"已处理 {current_idx}/{total_img_count} 张图像") print("数据集存储完成!") if __name__ == '__main__': main()
其他小优化点
- 用
os.path.join代替你写的add_slash,它会自动适配Windows/Linux的路径分隔符,更可靠 - 读取图像后可以加尺寸校验,确保所有图像尺寸一致,避免转数组时出错
- 如果想加快第二次运行速度,可以考虑用系统级的文件缓存工具,或者提前把图像转成二进制格式,但本质还是读盘速度的问题,和Python内存无关
内容的提问来源于stack exchange,提问作者Will
相关产品推荐
相关产品推荐

