如何用多进程(joblib/multiprocessing)优化双层循环的Python数据预处理函数?
优化数据预处理函数的多进程实现
核心思路
原函数的耗时瓶颈在PIL图像转扁平化numpy数组的CPU密集型操作,我们可以通过拆分独立任务、并行执行来提速。下面提供两种基于joblib的实现方案,适配你的双层循环逻辑。
方案一:单样本粒度并行(简单易实现)
步骤1:拆分单样本处理逻辑
把单个样本的图像转换和标签提取抽成独立函数,方便并行调用:
import numpy as np from joblib import Parallel, delayed import random import os def process_single_sample(img, label): return np.array(img).flatten(), label
步骤2:重构并行版主函数
保留原有的样本索引采样逻辑,先收集所有待处理样本,再批量并行处理:
def dataSizeSelection_parallel(data, data_path, labels, sample_sz, n_jobs=-1): sz_set = [] # 计算每个标签下的样本数量 for label in labels.values(): dir_path = os.path.join(data_path, label) # 替换硬拼接,避免路径问题 sz_set.append(len(os.listdir(dir_path))) idx_start = 0 pending_samples = [] # 收集所有选中的样本(图像+标签) for sz in sz_set: idx_end = idx_start + sz selected_idx = random.sample(range(idx_start, idx_end), sample_sz) idx_start += sz for j in selected_idx: pending_samples.append((data[j][0], data[j][1])) # 并行处理所有样本 results = Parallel(n_jobs=n_jobs)( delayed(process_single_sample)(img, label) for img, label in pending_samples ) # 拆分结果为数据和标签集合 data_set = [res[0] for res in results] label_set = [res[1] for res in results] return data_set, label_set
方案二:按标签分组并行(适合超大数据集)
如果数据集极大,按标签分组并行能减少任务调度开销,每个进程单独处理一个标签下的样本:
步骤1:定义单标签处理函数
def process_label_group(data, idx_start, group_size, sample_sz): idx_end = idx_start + group_size selected_idx = random.sample(range(idx_start, idx_end), sample_sz) group_data = [] group_label = [] for j in selected_idx: group_data.append(np.array(data[j][0]).flatten()) group_label.append(data[j][1]) return group_data, group_label
步骤2:并行处理各标签组
def dataSizeSelection_parallel_by_group(data, data_path, labels, sample_sz, n_jobs=-1): sz_set = [] for label in labels.values(): dir_path = os.path.join(data_path, label) sz_set.append(len(os.listdir(dir_path))) # 准备每个标签组的任务参数 task_args = [] idx_start = 0 for sz in sz_set: task_args.append((data, idx_start, sz, sample_sz)) idx_start += sz # 并行处理所有标签组 results = Parallel(n_jobs=n_jobs)( delayed(process_label_group)(*args) for args in task_args ) # 合并所有组的结果 data_set = [] label_set = [] for group_d, group_l in results: data_set.extend(group_d) label_set.extend(group_l) return data_set, label_set
关键注意事项
n_jobs=-1表示使用全部CPU核心,可根据机器性能调整为具体数值。- 多进程有进程间通信开销,仅适合样本量较大的场景,小数据集可能比串行更慢。
- 原代码中
dir是Python内置函数名,建议改用dir_path避免命名冲突。
内容的提问来源于stack exchange,提问作者wdy0308
相关产品推荐
相关产品推荐

