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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.31 20:48:47