如何让Dask的Client.map支持多参数函数,实现类似Pool.starmap效果
Dask Client 多参数函数映射实现方案
想要实现和Python标准库multiprocessing.Pool.starmap一致的效果,将dask.distributed.Client的映射方法直接作用于接收多个参数的函数,且不需要修改原有函数的参数结构。
问题复现
原有写法运行时报错,代码如下:
from contextlib import contextmanager from dask.distributed import Client @contextmanager def dask_client(**kwargs): """Dask客户端上下文管理器""" kwargs.setdefault("ip", "localhost:8786") client = Client(**kwargs) try: yield client except Exception: raise finally: client.close() def f(x,y,z): # 示例多参数函数 return x+y+z if __name__ == "__main__": with dask_client() as client: # 错误写法:仅传入两组可迭代参数,且map不会自动解包元组 client.map(f, (1,2,3), (1,2,3))
运行后报错显示函数f缺少第三个必填参数z,原因是Client.map默认按位置接收多组可迭代对象,每组对应函数的一个位置参数,上述写法仅传入了对应x和y的两组值,缺少对应z的参数序列。
可行解决方案(无需修改原函数定义)
方案1:使用Dask自带的Client.starmap(推荐)
Dask官方已经提供了完全对标标准库starmap的实现,会自动将传入的每个参数元组解包后传递给目标函数,用法和multiprocessing.Pool.starmap完全一致:
if __name__ == "__main__": with dask_client() as client: # 构造参数元组序列,每个元组对应一次f调用的全部位置参数 params = [(1,1,1), (2,2,2), (3,3,3)] # 调用starmap res = client.starmap(f, params) # 收集执行结果 print(client.gather(res)) # 输出 [3, 6, 9]
如果参数分别存储在不同的可迭代对象中,可通过zip打包后直接传入:
x_vals = (1,2,3) y_vals = (1,2,3) z_vals = (1,2,3) res = client.starmap(f, zip(x_vals, y_vals, z_vals))
方案2:lambda临时包装解包
如果受版本限制无法使用starmap,可以通过lambda对参数做临时解包,不需要修改原函数的逻辑:
if __name__ == "__main__": with dask_client() as client: params = [(1,1,1), (2,2,2), (3,3,3)] # 外层用lambda接收参数元组,解包后传给原函数f res = client.map(lambda args: f(*args), params) print(client.gather(res))
内容的提问来源于stack exchange,提问作者Andrex
相关产品推荐
相关产品推荐

