如何将Dask集群作为dask.compute调度器?序列化报错求解
我有一个类,通过如下上下文管理器创建Dask Client与集群:
class some_class(): def __init__(self,engine_kwargs: dict = None): self.distributed = engine_kwargs.get("distributed", False) self.dask_client = None self.n_workers = engine_kwargs.get( "n_workers", int(os.getenv("SLURM_CPUS_PER_TASK", os.cpu_count())) ) @contextmanager def dask_context(self): """Dask context manager to set up and close down client""" if self.distributed: if self.distributed_mode == "processes": processes = True self.dask_cluster = LocalCluster(n_workers=self.n_workers, processes=processes) self.dask_client = Client(self.dask_cluster) try: yield finally: if self.dask_client is not None: self.dask_client.close() self.local_cluster.close()
该类中有一个使用delayed的方法,旨在将任务分配到集群执行:
def some_class_method( self, ): min_ind = segy_container.trace_headers["SOME_GROUPER"].values.min() max_ind = segy_container.trace_headers["SOME_GROUPER"].values.max() tasks = [ delayed(self._process_group)(index,some,other,method,arguments,here ) for index in range(min_ind, max_ind + 1) ] with ProgressBar(): with self.dask_context(): results = compute(*tasks, scheduler=self.dask_client) #scheduler="processes"
当最后一行使用上下文管理器创建的dask_client作为调度器时,出现错误:
TypeError: cannot pickle '_asyncio.Task' object
若移除上下文管理器并使用scheduler="processes"则可正常运行。我推测scheduler="processes"不会序列化任务,因此可正常执行,但使用LocalCluster时会触发序列化并报错。请问是否可以将LocalCluster与delayed、compute配合使用,或有其他解决该问题的方法?
核心原因
报错根源是:用delayed装饰实例方法self._process_group时,会尝试序列化整个类实例self,而实例中持有dask_client(包含无法被pickle序列化的_asyncio.Task对象),导致分布式调度时序列化失败。而scheduler="processes"属于本地进程调度,序列化逻辑与分布式集群不同,因此未触发该问题。
具体解决方法
将
_process_group改为静态方法或独立函数
避免任务依赖整个类实例,把_process_group定义为类的静态方法(用@staticmethod装饰),或者提取为类外的独立函数,这样序列化时仅需传递必要参数,无需处理整个self对象。
示例:class some_class(): # ... 其他代码 ... @staticmethod def _process_group(index, some, other, arguments): # 处理逻辑,无需访问self pass移除
compute中的scheduler参数
创建Client后,Dask会自动将其设为全局默认客户端,直接调用compute(*tasks)即可,集群会自动接管任务调度,无需手动指定scheduler=self.dask_client。
修改后的代码:with ProgressBar(): with self.dask_context(): results = compute(*tasks) # 去掉scheduler参数修复上下文管理器的变量名错误
原代码中关闭集群时使用了未定义的self.local_cluster,应改为与创建时一致的self.dask_cluster:finally: if self.dask_client is not None: self.dask_client.close() self.dask_cluster.close() # 修正变量名将客户端/集群设为上下文局部变量
若无需在上下文之外访问dask_client或dask_cluster,可将它们定义为上下文管理器内的局部变量,避免被包含在要序列化的self实例中:@contextmanager def dask_context(self): """Dask context manager to set up and close down client""" dask_client = None dask_cluster = None if self.distributed: processes = self.distributed_mode == "processes" dask_cluster = LocalCluster(n_workers=self.n_workers, processes=processes) dask_client = Client(dask_cluster) try: yield finally: if dask_client is not None: dask_client.close() dask_cluster.close()
关于LocalCluster与delayed的兼容性
LocalCluster完全可以和delayed、compute配合使用,只要确保任务及依赖对象可序列化,并正确使用客户端,上述修改后即可正常运行。
内容的提问来源于stack exchange,提问作者abinitio

