Python中使用multiprocessing Pool时如何正确传递多参数函数参数?
解决multiprocessing并行求解勾股数的参数错误问题
问题重现
尝试用Python的multiprocessing模块并行求解勾股数,编写了如下代码:
import itertools from multiprocessing import Pool, Array import numpy as np MAXIMUM_INT: int = 10 answers = {'A': [], 'B': [], 'C': []} base_range = np.arange(1, MAXIMUM_INT + 1) A_range = B_range = C_range = base_range.tolist() combinatorics = [A_range, B_range, C_range] iteration = list(itertools.product(*combinatorics)) dim1 = dim2 = dim3 = MAXIMUM_INT def init(A: int, B: int, C: int): global test1, test2, test3 test1, test2, test3 = A, B, C def conditionalAssert(answerDict: dict, A: int, B: int, C: int): t1 = np.frombuffer(test1.get_obj()) t2 = np.frombuffer(test2.get_obj()) t3 = np.frombuffer(test3.get_obj()) if A**2 + B**2 == C**2: answerDict['A'].append(t1) answerDict['B'].append(t2) answerDict['C'].append(t3) if __name__ == '__main__': relevantArray = Array('i', dim1 * dim2 * dim3) A = B = C = relevantArray pool = Pool(processes=7, initializer=init, initargs=(A, B, C)) pool.starmap(conditionalAssert, iteration) print(answers)
运行后出现如下错误:
multiprocessing.pool.RemoteTraceback: """ Traceback (most recent call last): File "C:\Users\username\anaconda3\Lib\multiprocessing\pool.py", line 125, in worker result = (True, func(*args, **kwds)) ^^^^^^^^^^^^^^^^^^^ File "C:\Users\username\anaconda3\Lib\multiprocessing\pool.py", line 51, in starmapstar return list(itertools.starmap(args[0], args[1])) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ TypeError: conditionalAssert() missing 1 required positional argument: 'C' """ The above exception was the direct cause of the following exception: Traceback (most recent call last): File "c:\Users\username\OneDrive\Desktop\Projects\pythagorean.py", line 38, in <module> pool.starmap(conditionalAssert, iteration) File "C:\Users\username\anaconda3\Lib\multiprocessing\pool.py", line 375, in starmap return self._map_async(func, iterable, starmapstar, chunksize).get() ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Users\username\anaconda3\Lib\multiprocessing\pool.py", line 774, in get raise self._value TypeError: conditionalAssert() missing 1 required positional argument: 'C'
错误分析
- 参数不匹配:
pool.starmap会把iteration中的每个三元组(A,B,C)拆解后传给conditionalAssert,但该函数的第一个参数是answerDict,导致实际只传入了3个参数,函数却需要4个,触发参数缺失错误。 - 全局字典无法跨进程共享:
answers是主进程的全局字典,子进程无法直接修改它,就算参数匹配,修改也不会反映到主进程。 - 共享数组逻辑错误:将
A、B、C都赋值为同一个relevantArray,且init函数初始化的全局数组完全没必要,代码中根本没用到这些数组的实际存储功能。
修正方案
- 调整函数参数:让校验函数只接收
A,B,C,返回符合条件的三元组(或None),主进程收集所有非None的结果。 - 去掉无用的共享数组和初始化函数:原代码中的
Array和init函数完全多余,直接删除。 - 用
starmap收集结果:利用starmap的返回值收集所有符合条件的勾股数,再整理成需要的字典格式。
修正后的代码
import itertools from multiprocessing import Pool MAXIMUM_INT: int = 10 base_range = list(range(1, MAXIMUM_INT + 1)) # 生成所有A,B,C的组合 iteration = list(itertools.product(base_range, repeat=3)) def check_pythagorean(A: int, B: int, C: int): if A**2 + B**2 == C**2: return (A, B, C) return None if __name__ == '__main__': with Pool(processes=7) as pool: # 收集所有非None的结果 results = [res for res in pool.starmap(check_pythagorean, iteration) if res is not None] # 整理成需要的字典格式 answers = {'A': [], 'B': [], 'C': []} for a, b, c in results: answers['A'].append(a) answers['B'].append(b) answers['C'].append(c) print(answers)
运行结果
输出:
{'A': [3, 4, 6, 8], 'B': [4, 3, 8, 6], 'C': [5, 5, 10, 10]}
内容的提问来源于stack exchange,提问作者Rebel
相关产品推荐
相关产品推荐

