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

使用multiprocessing.Pool.map结合partial函数未提速的原因及优化方案

问题描述

执行刀切重采样任务时,尝试用multiprocessing加速代码:

  • 原串行代码通过内部函数_resample结合map实现重采样;
  • 改用多进程时,将_resample移到外部,用functools.partial传递value、weight、label参数后调用Pool.map,但代码未实现提速;
  • 改用全局变量传递参数时多进程有效,但代码不够优雅;
  • 尝试将重采样逻辑封装为类(Python 3.12支持实例方法序列化),但效果与使用partial一致。

原串行代码

import numpy as np
def jackknife_resampling(value, weight, label):
    # label是标记子样本的整数数组
    def _resample(i):
        _value = value[label != i]
        _weight = weight[label != i]
        return np.average(_value, weights=_weight)
    return list(map(_resample, np.unique(label)))

使用partial的多进程代码

from functools import partial
from multiprocessing import Pool
import time

def _resample(value, weight, label, i):
    _value = value[label != i]
    _weight = weight[label != i]
    return np.average(_value, weights=_weight)
def jackknife_resampling_mp(value, weight, label, Npro):
    my_pool = Pool(Npro)
    return my_pool.map(partial(_resample, value, weight, label), np.unique(label))

if __name__ == '__main__':
    Nsamp = 1000000
    value = np.random.uniform(0, 1, size=Nsamp)
    weight = np.random.uniform(1, 2, size=Nsamp)
    label = np.random.randint(0, 200, size=Nsamp)
    t = time.time()
    jackknife_resampling_mp(value, weight, label, 40)
    print('time', time.time()-t)

使用全局变量的多进程代码

from multiprocessing import Pool
import time

def _resample(i):
    global value, weight, label
    _value = value[label != i]
    _weight = weight[label != i]
    return np.average(_value, weights=_weight)

def jackknife_resampling_mp(_value, _weight, _label, Npro):
    global value, weight, label
    value = _value
    weight = _weight
    label = _label
    my_pool = Pool(Npro)
    return my_pool.map(_resample, np.unique(label))

if __name__ == '__main__':
    Nsamp = 1000000
    value = np.random.uniform(0, 1, size=Nsamp)
    weight = np.random.uniform(1, 2, size=Nsamp)
    label = np.random.randint(0, 200, size=Nsamp)
    t = time.time()
    jackknife_resampling_mp(value, weight, label, 40)
    print('time', time.time()-t)
问题原因
  1. partial传递大数组的开销问题:
    使用partial时,value、weight、label这些大numpy数组会被完整序列化(pickle)并复制到每个子进程。对于百万级元素的数组,序列化和跨进程传输的耗时远超过并行计算节省的时间,直接抵消了多进程的优势。

  2. 全局变量的高效性原理:
    全局变量方式下,Unix/Linux/macOS系统中multiprocessing.Pool默认用fork模式创建子进程,子进程会继承父进程的内存空间,且采用**写时复制(Copy-On-Write)**机制——只有当子进程修改数据时才会复制内存块,而这里我们只是读取数组,因此几乎没有数据复制开销,并行效率得以体现。

  3. 类实例方法的本质问题:
    Python 3.12支持实例方法序列化,但实例包含的大数组属性依然会被完整序列化并传递给子进程,本质和partial的问题一致,同样存在大数组复制的开销,所以无法提速。

更优实现方式

方案1:Unix/Linux/macOS平台(推荐)——用Pool的initializer传递共享数据

利用fork的写时复制特性,通过initializer和initargs将大数组传递给子进程的全局变量,既避免手动设置全局变量的不优雅,又能高效共享数据:

import numpy as np
from multiprocessing import Pool
import time

# 子进程共享的全局变量
_shared_value = None
_shared_weight = None
_shared_label = None

def _init_worker(value, weight, label):
    """初始化子进程,设置共享变量"""
    global _shared_value, _shared_weight, _shared_label
    _shared_value = value
    _shared_weight = weight
    _shared_label = label

def _resample(i):
    """重采样逻辑,直接访问共享变量"""
    mask = _shared_label != i
    return np.average(_shared_value[mask], weights=_shared_weight[mask])

def jackknife_resampling_mp(value, weight, label, Npro):
    with Pool(Npro, initializer=_init_worker, initargs=(value, weight, label)) as my_pool:
        return my_pool.map(_resample, np.unique(label))

if __name__ == '__main__':
    Nsamp = 1000000
    value = np.random.uniform(0, 1, size=Nsamp)
    weight = np.random.uniform(1, 2, size=Nsamp)
    label = np.random.randint(0, 200, size=Nsamp)
    t = time.time()
    jackknife_resampling_mp(value, weight, label, 40)
    print('time', time.time()-t)

方案2:跨平台(含Windows)——用shared_memory共享numpy数组

Windows系统不支持fork,需用multiprocessing.shared_memory创建共享内存块,让所有子进程访问同一块内存,避免大数组复制:

import numpy as np
from multiprocessing import Pool, shared_memory
import time

def _resample(args):
    """重采样逻辑,从共享内存读取数据"""
    i, shm_names, shapes, dtypes = args
    shm_name_val, shm_name_wgt, shm_name_lbl = shm_names
    shape_val, shape_wgt, shape_lbl = shapes
    dtype_val, dtype_wgt, dtype_lbl = dtypes

    # 连接到共享内存块
    shm_val = shared_memory.SharedMemory(name=shm_name_val)
    shm_wgt = shared_memory.SharedMemory(name=shm_name_wgt)
    shm_lbl = shared_memory.SharedMemory(name=shm_name_lbl)

    # 创建numpy数组视图(不复制数据)
    value = np.ndarray(shape_val, dtype=dtype_val, buffer=shm_val.buf)
    weight = np.ndarray(shape_wgt, dtype=dtype_wgt, buffer=shm_wgt.buf)
    label = np.ndarray(shape_lbl, dtype=dtype_lbl, buffer=shm_lbl.buf)

    # 计算重采样结果
    mask = label != i
    result = np.average(value[mask], weights=weight[mask])

    # 关闭共享内存连接
    shm_val.close()
    shm_wgt.close()
    shm_lbl.close()
    return result

def jackknife_resampling_mp(value, weight, label, Npro):
    # 创建共享内存块并写入数据
    shm_val = shared_memory.SharedMemory(create=True, size=value.nbytes)
    shm_wgt = shared_memory.SharedMemory(create=True, size=weight.nbytes)
    shm_lbl = shared_memory.SharedMemory(create=True, size=label.nbytes)

    np.ndarray(value.shape, dtype=value.dtype, buffer=shm_val.buf)[:] = value[:]
    np.ndarray(weight.shape, dtype=weight.dtype, buffer=shm_wgt.buf)[:] = weight[:]
    np.ndarray(label.shape, dtype=label.dtype, buffer=shm_lbl.buf)[:] = label[:]

    try:
        # 准备子进程参数
        unique_labels = np.unique(label)
        args_list = [
            (i,
             (shm_val.name, shm_wgt.name, shm_lbl.name),
             (value.shape, weight.shape, label.shape),
             (value.dtype, weight.dtype, label.dtype))
            for i in unique_labels
        ]
        with Pool(Npro) as my_pool:
            return my_pool.map(_resample, args_list)
    finally:
        # 释放共享内存
        shm_val.unlink()
        shm_wgt.unlink()
        shm_lbl.unlink()

if __name__ == '__main__':
    Nsamp = 1000000
    value = np.random.uniform(0, 1, size=Nsamp)
    weight = np.random.uniform(1, 2, size=Nsamp)
    label = np.random.randint(0, 200, size=Nsamp)
    t = time.time()
    jackknife_resampling_mp(value, weight, label, 40)
    print('time', time.time()-t)

内容的提问来源于stack exchange,提问作者Di He

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 05:23:12