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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 17:44:06