使用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)
问题原因
partial传递大数组的开销问题:
使用partial时,value、weight、label这些大numpy数组会被完整序列化(pickle)并复制到每个子进程。对于百万级元素的数组,序列化和跨进程传输的耗时远超过并行计算节省的时间,直接抵消了多进程的优势。全局变量的高效性原理:
全局变量方式下,Unix/Linux/macOS系统中multiprocessing.Pool默认用fork模式创建子进程,子进程会继承父进程的内存空间,且采用**写时复制(Copy-On-Write)**机制——只有当子进程修改数据时才会复制内存块,而这里我们只是读取数组,因此几乎没有数据复制开销,并行效率得以体现。类实例方法的本质问题:
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
相关产品推荐
相关产品推荐

