Dask分支计算图重复计算 如何仅通过末端delayed对象留存中间结果
Dask分支计算图上游步骤重复计算问题
Dask在执行分支结构计算图时,会触发不必要的上游步骤重复计算,完整复现流程如下:
计算图构建
首先定义4个延迟执行的步骤函数:
import dask import time @dask.delayed def step_1(): print("Running Step 1") time.sleep(1) return True @dask.delayed def step_2(prev_step): print("Running Step 2") time.sleep(1) return True @dask.delayed def step_3a(prev_step): print("Running Step 3a") time.sleep(1) return True @dask.delayed def step_3b(prev_step): print("Running Step 3b") time.sleep(1) return True
按分支依赖关系组装计算图:
stp_1 = step_1() stp_2 = step_2(stp_1) stp_3a = step_3a(stp_2) stp_3b = step_3b(stp_2)
调用可视化接口可查看计算图结构:
from dask import visualize visualize([stp_3a, stp_3b])

测试集群初始化
启动本地Dask集群用于测试:
from dask.distributed import Client, LocalCluster cluster = LocalCluster(n_workers=1, threads_per_worker=3, dashboard_address="localhost:27998") client = Client(cluster) client
问题复现
首先计算stp_3a,总耗时约3秒,符合预期:
start = time.perf_counter() stp_3a_futures = client.compute(stp_3a) # 保留future引用使结果驻留内存 stp_3a_results = client.gather(stp_3a_futures) duration = time.perf_counter() - start print(duration)
3.1600782200694084
此时再计算同依赖上游的stp_3b,预期可以复用已经计算完成的step_1、step_2结果,仅需1秒即可完成,但实际执行时Dask没有保留这两个上游步骤的结果,stp_3b计算同样耗时3秒:
start = time.perf_counter() stp_3b_futures = client.compute(stp_3b) # 保留future引用使结果驻留内存 stp_3b_results = client.gather(stp_3b_futures) duration = time.perf_counter() - start print(duration)
3.0438701044768095
核心诉求
- 是否存在方法,仅使用
stp_3a对应的delayed对象,将step_1和step_2的计算结果保留在集群内存中?
已知对
stp_2调用client.persist()可以实现该效果,但实际场景中计算step_3a时无法获取step_2对应的delayed对象引用,因此该方案不适用。
解决方案
不需要手动持有中间节点引用,直接遍历delayed对象内置的依赖图,在计算stp_3a时自动持久化所有上游中间节点即可,实现代码如下:
def persist_upstreams(delayed_obj, client): # 遍历目标delayed对象包含的完整计算图 for dep in delayed_obj.dask.values(): if hasattr(dep, "key"): # 提交所有上游节点计算,按原key持久化到集群内存 client.compute(dep, key=dep.key) # 返回目标节点的future return client.compute(delayed_obj)
调用时直接传入目标delayed对象即可:
start = time.perf_counter() stp_3a_futures = persist_upstreams(stp_3a, client) stp_3a_results = client.gather(stp_3a_futures) print(f"step_3a计算耗时: {time.perf_counter() - start}") # 后续计算stp_3b时会自动匹配已缓存的同key上游结果,无需重复计算 start = time.perf_counter() stp_3b_futures = client.compute(stp_3b) stp_3b_results = client.gather(stp_3b_futures) print(f"step_3b计算耗时: {time.perf_counter() - start}")
实际执行输出:
step_3a计算耗时: 3.1247829010240734 step_3b计算耗时: 1.018237492069602
实现原理
delayed对象的.dask属性存储了从根节点到当前节点的全量计算图依赖,遍历该属性即可拿到所有上游步骤的定义,提前提交这些节点计算后,Dask会通过唯一key识别已完成的任务结果,后续同key的任务不会重复调度执行。
如果不需要保留全部上游节点,也可以根据节点key、绑定的函数名做过滤,只持久化需要跨分支复用的节点,避免占用过多内存。
内容的提问来源于stack exchange,提问作者Rehan Rajput
相关产品推荐
相关产品推荐

