如何实现可修改参数的Python函数并行处理?
并行修改共享数组的可行方案
你想要并行化修改数组的函数,但多进程默认会触发写时复制机制,子进程修改的是数组副本,导致原数组未受影响。以下几种方案可以解决这个问题:
1. 使用multiprocessing共享内存数组
利用multiprocessing.RawArray或Array创建共享内存区域,再将其包装为numpy数组,让所有进程直接操作同一块内存,修改会直接反映到原数组。
import numpy as np import multiprocessing def Test(i, shared_arr, length): # 将共享内存转为numpy数组 a = np.frombuffer(shared_arr, dtype=np.uint8) if i % 2 == 0: for x in range(0, length, 2): a[x] = 2 else: for x in range(1, length, 2): a[x] = 1 if __name__ == "__main__": length = 10 # 创建共享内存数组,'B'对应numpy.uint8类型 shared_arr = multiprocessing.RawArray('B', length) # 绑定到numpy数组 a = np.frombuffer(shared_arr, dtype=np.uint8) with multiprocessing.Pool(processes=2) as pool: data = [(0, shared_arr, length), (1, shared_arr, length)] pool.starmap(Test, data) print(a)
2. 使用multiprocessing.Manager管理共享数组
Manager会启动一个专门的管理进程来维护共享对象,子进程通过代理访问。这种方式适合复杂共享结构,但性能略低于直接共享内存。
import numpy as np import multiprocessing def Test(i, shared_list, length): if i % 2 == 0: for x in range(0, length, 2): shared_list[x] = 2 else: for x in range(1, length, 2): shared_list[x] = 1 if __name__ == "__main__": length = 10 with multiprocessing.Manager() as manager: # 创建共享list shared_list = manager.list([0]*length) with multiprocessing.Pool(processes=2) as pool: data = [(0, shared_list, length), (1, shared_list, length)] pool.starmap(Test, data) # 转为numpy数组查看结果 a_result = np.array(shared_list, dtype=np.uint8) print(a_result)
3. 使用线程并行(threading模块)
线程属于同一进程,共享内存空间,修改的直接是原数组。注意:CPU密集型任务可能受GIL限制,但numpy底层操作通常会释放GIL,这类数组修改任务能获得并行效率。
import numpy as np import threading def Test(i, a, length): if i % 2 == 0: for x in range(0, length, 2): a[x] = 2 else: for x in range(1, length, 2): a[x] = 1 if __name__ == "__main__": a = np.ndarray(shape=10, dtype=np.uint8) t1 = threading.Thread(target=Test, args=(0, a, a.shape[0])) t2 = threading.Thread(target=Test, args=(1, a, a.shape[0])) t1.start() t2.start() t1.join() t2.join() print(a)
4. 拆分任务+结果合并
避免共享内存,让每个子进程处理自己负责的数组片段,返回索引和对应值,最后在主进程合并到原数组。这是多进程编程的常用模式,无需处理共享状态问题。
import numpy as np import multiprocessing def process_chunk(i, length): if i % 2 == 0: indices = range(0, length, 2) values = [2]*len(indices) else: indices = range(1, length, 2) values = [1]*len(indices) return indices, values if __name__ == "__main__": length = 10 a = np.ndarray(shape=length, dtype=np.uint8) with multiprocessing.Pool(processes=2) as pool: results = pool.starmap(process_chunk, [(0, length), (1, length)]) # 合并结果到原数组 for indices, values in results: a[list(indices)] = values print(a)
内容的提问来源于stack exchange,提问作者FiReTiTi
相关产品推荐
相关产品推荐

