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

如何用多进程(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 03:21:32