如何用multiprocessing并行化类方法内的numpy数组生成
多进程并行化类内数组填充的实现方案
首先修正原代码中的语法错误:
DummyCalculator.compute方法定义缺少冒号DummyKeeper.myfunc调用时需要传入x参数
核心实现方案(进程池版)
利用multiprocessing.Pool分配进程,让每个DummyCalculator实例在独立进程中执行计算。当前代码结构下,DummyCalculator实例可被pickle序列化,满足多进程数据传递要求。
修改后的完整代码:
import multiprocessing as mp import numpy as np from itertools import count import time class DummyCalculator: _ids = count(0) def __init__(self): self._id = next(self._ids) def compute(self, x: np.array) -> np.array: time.sleep(10) return x + self._id class DummyKeeper: def __init__(self): self._calculators = [] def addCalculator(self, calculator: DummyCalculator): self._calculators.append(calculator) def myfunc(self, x: np.array) -> np.array: # 初始化进程池,进程数与计算器实例数一致 with mp.Pool(processes=len(self._calculators)) as pool: # 用starmap传递多参数,并行调用每个计算器的compute方法 results = pool.starmap(lambda calc, arr: calc.compute(arr), [(c, x) for c in self._calculators]) # 合并结果并扁平化 out = np.array(results).flatten() return out if __name__ == '__main__': keeper = DummyKeeper() # 添加4个计算器实例 for _ in range(4): keeper.addCalculator(DummyCalculator()) x = np.zeros(100000) start_time = time.time() result = keeper.myfunc(x) print(f"总耗时: {time.time() - start_time:.2f}秒") print(f"结果前5个元素: {result[:5]}")
关键说明
- 进程池通过
with语句自动管理生命周期,无需手动关闭进程 - 串行版本需耗时约40秒(4个实例×10秒),多进程版本耗时约10秒,符合并行加速预期
starmap方法适配多参数函数调用,这里用lambda包装compute方法的调用逻辑
替代方案(手动创建进程)
如果需要更精细的进程控制(如超时、自定义生命周期),可手动创建进程并通过队列收集结果:
def myfunc(self, x: np.array) -> np.array: queue = mp.Queue() processes = [] # 定义子进程执行的工作函数 def worker(calc, arr, q): q.put(calc.compute(arr)) # 创建并启动所有进程 for c in self._calculators: p = mp.Process(target=worker, args=(c, x, queue)) processes.append(p) p.start() # 从队列收集所有结果 results = [] for _ in processes: results.append(queue.get()) # 等待所有进程执行完毕 for p in processes: p.join() out = np.array(results).flatten() return out
内容的提问来源于stack exchange,提问作者user20679
相关产品推荐
相关产品推荐

