如何实现多进程修改Numpy数组并正确获取最小值?
多进程填充数组并求最小值的问题解决
问题根源
多进程环境下,子进程会复制主进程的内存空间,你在limit函数里修改的min_array是子进程自己的副本,主进程的原数组根本没被改动,所以最后输出全1的最小值。
方案一:收集子进程返回结果(最简单)
不需要共享内存,让每个子进程计算完expensiveFunction()后返回结果,主进程收集所有结果再求最小值,甚至不用提前创建数组:
import multiprocessing import numpy as np import psutil import sys def expensiveFunction(): # 模拟耗时操作,返回示例随机数 import random return random.randint(1, 100) def limit(n): # 只返回计算结果,无需操作数组 return expensiveFunction() def run_parallel(function, nmax, nthreads, debug=False): pool = multiprocessing.Pool(nthreads) try: # 用map直接获取所有子进程的结果 results = pool.map(function, list(range(nmax))) except KeyboardInterrupt: print('Caught interrupt!') pool.terminate() exit(1) else: pool.close() pool.join() return results if __name__ == "__main__": nthreads = psutil.cpu_count() # 注意将命令行参数转为整数 number_expensive_calls = int(sys.argv[1]) results = run_parallel(limit, number_expensive_calls, nthreads, debug=False) # 直接计算结果的最小值 print(np.min(results))
方案二:使用共享内存(适合必须操作数组的场景)
之前用共享内存出现nan,大概率是共享内存的创建、关联步骤有误。下面是两种正确的实现方式:
用multiprocessing.Array实现
import multiprocessing import numpy as np import psutil import sys def expensiveFunction(): import random return random.randint(1, 100) def limit(args): n, shared_array = args entry = expensiveFunction() # 将共享内存数组转为numpy数组操作 np_array = np.frombuffer(shared_array, dtype=np.float64) np_array[n] = entry def run_parallel(function, nmax, nthreads, shared_array, debug=False): pool = multiprocessing.Pool(nthreads) # 把索引和共享数组作为参数传递给子进程 args_list = [(i, shared_array) for i in range(nmax)] try: pool.map(function, args_list) except KeyboardInterrupt: print('Caught interrupt!') pool.terminate() exit(1) else: pool.close() pool.join() if __name__ == "__main__": nthreads = psutil.cpu_count() number_expensive_calls = int(sys.argv[1]) # 创建共享内存数组,dtype按需选择,这里用float64 shared_array = multiprocessing.Array('d', number_expensive_calls) # 转为numpy数组用于后续操作 min_array = np.frombuffer(shared_array, dtype=np.float64) # 初始化数组(可选) min_array[:] = np.ones(number_expensive_calls) run_parallel(limit, number_expensive_calls, nthreads, shared_array, debug=False) print(np.min(min_array))
用shared_memory模块实现(Python 3.8+)
import multiprocessing import numpy as np import psutil import sys from multiprocessing import shared_memory def expensiveFunction(): import random return random.randint(1, 100) def limit(args): n, shm_name, shape, dtype = args # 关联主进程创建的共享内存块 existing_shm = shared_memory.SharedMemory(name=shm_name) # 创建numpy数组关联共享内存 np_array = np.ndarray(shape, dtype=dtype, buffer=existing_shm.buf) entry = expensiveFunction() np_array[n] = entry # 关闭共享内存(主进程最后统一释放) existing_shm.close() def run_parallel(function, nmax, nthreads, shm_args, debug=False): pool = multiprocessing.Pool(nthreads) shm_name, shape, dtype = shm_args args_list = [(i, shm_name, shape, dtype) for i in range(nmax)] try: pool.map(function, args_list) except KeyboardInterrupt: print('Caught interrupt!') pool.terminate() exit(1) else: pool.close() pool.join() if __name__ == "__main__": nthreads = psutil.cpu_count() number_expensive_calls = int(sys.argv[1]) shape = (number_expensive_calls,) dtype = np.float64 # 创建共享内存块 shm = shared_memory.SharedMemory(create=True, size=shape[0]*dtype.itemsize) # 创建numpy数组关联共享内存 min_array = np.ndarray(shape, dtype=dtype, buffer=shm.buf) min_array[:] = np.ones(shape) run_parallel(limit, number_expensive_calls, nthreads, (shm.name, shape, dtype), debug=False) print(np.min(min_array)) # 释放共享内存 shm.close() shm.unlink()
关键注意点
- 原代码中
sys.argv[1]是字符串,必须转为整数才能作为数组长度使用,否则会报错。 - 共享内存方案中,子进程必须正确关联主进程创建的共享内存块,否则会操作错误内存区域导致nan或异常。
- 如果仅需要最小值,方案一效率更高,无需额外内存同步开销,直接收集结果即可。
内容的提问来源于stack exchange,提问作者Cesar Diaz Blanco
相关产品推荐
相关产品推荐

