使用Multiprocessing的Pool.map处理字典并保留键的简便方法
解决方法
要让并行处理后的结果和原字典的键对应,有两种可靠的方式,分别适用于不同场景:
方案一:利用Python 3.7+的字典有序性(简单直接)
从Python 3.7开始,字典会严格保留插入顺序,models.keys()和models.values()的顺序完全对应。结合Pool.map保证结果顺序与输入顺序一致的特性(官方文档明确说明),可以直接将键列表和结果列表配对成字典:
from multiprocessing import Pool from functools import partial from sklearn.model_selection import cross_validate from sklearn.dummy import DummyClassifier # 假设X、y已定义 models = { 'Dummy1': DummyClassifier(strategy="constant", constant=1), 'Dummy2': DummyClassifier(strategy="constant", constant=0), 'Dummy3': DummyClassifier(strategy="constant", constant=0.5) } threads = 3 with Pool(threads) as p: result_list = p.map(partial(cross_validate, X=X, y=y), models.values()) # 直接配对键与结果 results = dict(zip(models.keys(), result_list))
方案二:传递键值对(不依赖版本,更稳妥)
如果需要兼容Python 3.7以下版本,或者不想依赖字典有序性,可以把键和模型作为元组传入并行函数,让函数直接返回键与结果的配对:
方式1:全局变量传递X、y
from multiprocessing import Pool from sklearn.model_selection import cross_validate from sklearn.dummy import DummyClassifier X, y = ... # 你的数据集 models = { 'Dummy1': DummyClassifier(strategy="constant", constant=1), 'Dummy2': DummyClassifier(strategy="constant", constant=0), 'Dummy3': DummyClassifier(strategy="constant", constant=0.5) } threads = 3 def process_model(key_model_pair): key, model = key_model_pair return key, cross_validate(model, X=X, y=y) with Pool(threads) as p: result_pairs = p.map(process_model, models.items()) # 将元组列表转为字典 results = dict(result_pairs)
方式2:用partial传递X、y(更灵活,避免全局变量)
from multiprocessing import Pool from functools import partial from sklearn.model_selection import cross_validate from sklearn.dummy import DummyClassifier models = { 'Dummy1': DummyClassifier(strategy="constant", constant=1), 'Dummy2': DummyClassifier(strategy="constant", constant=0), 'Dummy3': DummyClassifier(strategy="constant", constant=0.5) } threads = 3 def process_model(key_model_pair, X, y): key, model = key_model_pair return key, cross_validate(model, X=X, y=y) with Pool(threads) as p: # 绑定X、y参数 task = partial(process_model, X=X, y=y) result_pairs = p.map(task, models.items()) results = dict(result_pairs)
关于顺序的补充说明
Pool.map的官方文档明确规定:返回的结果列表严格按照输入迭代器的顺序排列,即使并行任务的执行顺序是乱的,最终结果也会按输入顺序整理。所以方案一的顺序依赖是安全的,但方案二直接将键与结果绑定,完全避免了顺序相关的潜在问题,推荐用于生产环境。
内容的提问来源于stack exchange,提问作者steliosbl
相关产品推荐
相关产品推荐

