You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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)

关键修正点

  1. 调整函数参数结构:将shared_objects拆分为indata和output两个独立参数,放在id_val之前,匹配Mpire的参数传递逻辑。
  2. 优化资源管理:使用Manager上下文管理器自动释放资源,避免手动管理的繁琐。
  3. 修正类型注解:根据实际数据类型将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.22 20:42:19