如何使用Python的mpire包同步更新多进程共享字典?
用Mpire实现多进程同步更新共享字典
我在配备8 vCPU、16GB内存的Amazon SageMaker Linux多核机器上,尝试用Python的mpire包实现多进程同步更新共享字典,但代码无法正常运行。我知道用multiprocessing的Process或map方法能完成,但想确认是否有办法用mpire实现。
错误代码
from typing import Dict import pandas as pd from datetime import datetime def myFunc(shared_objects, id_val): indata, output = shared_objects # Temporary store for model output for an input ID temp: Dict[str, int] = dict() # Filter data for input ID and store output in temp variable indata2 = indata.loc[indata['ID']==id_val] temp = indata2.groupby(['M_CODE'])['VALUE'].sum().to_dict() # store the result .. I want this to happen synchronously output[id_val] = temp #******************************************************************* if __name__ == '__main__': from mpire import WorkerPool from multiprocessing import Manager # This is just a sample data inputData = pd.DataFrame(dict({'ID':['A', 'B', 'A', 'C', 'A'], 'M_CODE':['AKQ1', 'ALM3', 'BLC4', 'ALM4', 'BLC4'], 'VALUE':[0.75, 1, 1.75, 0.67, 3], })) start_time = datetime.now() print(start_time, '>> Process started.') # Use a shared dict to store results from various workers manager = Manager() output: Dict[str, Dict[str, int]] = manager.dict() shared_objects = inputData, output with WorkerPool(n_jobs=7, shared_objects=shared_objects) as pool: results = pool.map_unordered(myFunc, inputData['ID'].unique(), progress_bar=True) print(datetime.now(), '>> Process completed -> total time taken:', datetime.now()-start_time)
报错信息
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) <ipython-input-10-df7d847398a1> in <module> 37 38 with WorkerPool(n_jobs=7, shared_objects=shared_objects) as pool: ---> 39 results = pool.map_unordered(myFunc, inputData['ID'].unique(), progress_bar=True) 40 41 print(datetime.now(), '>> Process completed -> total time taken:', datetime.now()-start_time) /opt/conda/lib/python3.7/site-packages/mpire/pool.py in map_unordered(self, func, iterable_of_args, iterable_len, max_tasks_active, chunk_size, n_splits, worker_lifespan, progress_bar, progress_bar_position, enable_insights, worker_init, worker_exit, task_timeout, worker_init_timeout, worker_exit_timeout) 418 n_splits, worker_lifespan, progress_bar, progress_bar_position, 419 enable_insights, worker_init, worker_exit, task_timeout, worker_init_timeout, ---> 420 worker_exit_timeout)) 421 422 def imap(self, func: Callable, iterable_of_args: Union[Sized, Iterable], iterable_len: Optional[int] = None, /opt/conda/lib/python3.7/site-packages/mpire/pool.py in imap_unordered(self, func, iterable_of_args, iterable_len, max_tasks_active, chunk_size, n_splits, worker_lifespan, progress_bar, progress_bar_position, enable_insights, worker_init, worker_exit, task_timeout, worker_init_timeout, worker_exit_timeout) 664 # Terminate if exception has been thrown at this point 665 if self._worker_comms.exception_thrown(): ---> 666 self._handle_exception(progress_bar_handler) 667 668 # All results are in: it's clean up time /opt/conda/lib/python3.7/site-packages/mpire/pool.py in _handle_exception(self, progress_bar_handler) 729 # Raise 730 logger.debug("Re-raising obtained exception") ---> 731 raise err(traceback_str) 732 733 def stop_and_join(self, progress_bar_handler: Optional[ProgressBarHandler] = None, ValueError: Exception occurred in Worker-0 with the following arguments: Arg 0: 'A' Traceback (most recent call last): File "/opt/conda/lib/python3.7/site-packages/mpire/worker.py", line 352, in _run_safely results = func() File "/opt/conda/lib/python3.7/site-packages/mpire/worker.py", line 288, in _func _results = func(args) File "/opt/conda/lib/python3.7/site-packages/mpire/worker.py", line 455, in _helper_func return self._call_func(func, args) File "/opt/conda/lib/python3.7/site-packages/mpire/worker.py", line 472, in _call_func return func(args) File "<ipython-input-10-df7d847398a1>", line 9, in myFunc indata2 = indata.loc[indata['ID']==id_val] File "/opt/conda/lib/python3.7/site-packages/pandas/core/ops/common.py", line 69, in new_method return method(self, other) File "/opt/conda/lib/python3.7/site-packages/pandas/core/arraylike.py", line 32, in __eq__ return self._cmp_method(other, operator.eq) File "/opt/conda/lib/python3.7/site-packages/pandas/core/series.py", line 5502, in _cmp_method res_values = ops.comparison_op(lvalues, rvalues, op) File "/opt/conda/lib/python3.7/site-packages/pandas/core/ops/array_ops.py", line 262, in comparison_op "Lengths must match to compare", lvalues.shape, rvalues.shape ValueError: ('Lengths must match to compare', (5,), (1,))
问题原因
报错核心是Mpire对shared_objects的参数传递逻辑与预期不符:当shared_objects为元组时,Mpire会将元组内的每个元素作为独立参数传递给函数,而非将整个元组作为单个参数。你的函数将shared_objects作为单个参数接收,导致参数顺序混乱——原本的id_val被错误解析为元组的一部分,最终引发DataFrame比较时的长度不匹配错误。
修正后的Mpire实现代码
from typing import Dict import pandas as pd from datetime import datetime from mpire import WorkerPool from multiprocessing import Manager def myFunc(indata, output, id_val): # 按ID筛选数据并计算分组汇总 indata2 = indata.loc[indata['ID'] == id_val] temp = indata2.groupby(['M_CODE'])['VALUE'].sum().to_dict() # 同步更新进程安全的共享字典 output[id_val] = temp if __name__ == '__main__': # 示例输入数据 inputData = pd.DataFrame({ 'ID': ['A', 'B', 'A', 'C', 'A'], 'M_CODE': ['AKQ1', 'ALM3', 'BLC4', 'ALM4', 'BLC4'], 'VALUE': [0.75, 1, 1.75, 0.67, 3] }) start_time = datetime.now() print(start_time, '>> Process started.') # 使用Manager创建进程安全的共享字典 with Manager() as manager: output: Dict[str, Dict[str, float]] = manager.dict() # 将DataFrame和共享字典打包为shared_objects元组 shared_objects = (inputData, output) # 初始化Mpire WorkerPool并行处理任务 with WorkerPool(n_jobs=7, shared_objects=shared_objects) as pool: pool.map_unordered(myFunc, inputData['ID'].unique(), progress_bar=True) # 将共享字典转换为普通字典(可选,便于后续操作) output = dict(output) print(datetime.now(), '>> Process completed -> total time taken:', datetime.now() - start_time) print("最终结果:", output)
关键修正点
- 调整函数参数结构:将
shared_objects拆分为indata和output两个独立参数,放在id_val之前,匹配Mpire的参数传递逻辑。 - 优化资源管理:使用
Manager上下文管理器自动释放资源,避免手动管理的繁琐。 - 修正类型注解:根据实际数据类型将
Dict[str, int]调整为Dict[str, float],确保注解准确。
对比可用的multiprocessing实现
def myFunc(id_val, output, indata): # Temporary store for model output for an input ID temp: Dict[str, int] = dict() # Filter data for input ID and store output in temp variable indata2 = indata.loc[indata['ID']==id_val] temp = indata2.groupby(['M_CODE'])['VALUE'].sum().to_dict() # store the result .. I want this to happen synchronously output[id_val] = temp #******************************************************************* if __name__ == '__main__': import pandas as pd from typing import Dict from itertools import repeat from multiprocessing import Manager from datetime import datetime # This is just a sample data inputData = pd.DataFrame(dict({'ID':['A', 'B', 'A', 'C', 'A'], 'M_CODE':['AKQ1', 'ALM3', 'BLC4', 'ALM4', 'BLC4'], 'VALUE':[0.75, 1, 1.75, 0.67, 3], })) start_time = datetime.now() print(start_time, '>> Process started.') # Use a shared dict to store results from various workers with Manager() as manager: output: Dict[str, Dict[str, int]] = manager.dict() # Start processes involving n workers with manager.Pool(processes=7, ) as pool: pool.starmap(myFunc, zip(inputData['ID'].unique(), repeat(output), repeat(inputData)), chunksize = max(inputData['ID'].nunique() // (7*4), 1)) output = dict(output) print(datetime.now(), '>> Process completed -> total time taken:', datetime.now()-start_time)
内容的提问来源于stack exchange,提问作者Sauvik De
相关产品推荐
相关产品推荐

