如何在Dask Distributed中存储数据到Zarr并复用中间结果优化计算?
问题
尝试将本地Dask Array计算迁移到Dask Distributed,但遇到适配问题。核心需求是:将分布式计算结果存储至内存Zarr数组,同时借助Dask缓存与图优化获取数组统计量。
预处理逻辑为减去参考行和行均值,本地实现代码如下:
# create toy data zarry = zarr.open('./example.zarr', mode='w') sample_data = zarry.create_group("sample_data", overwrite=True) s1 = sample_data.create_dataset('sample1', shape=(5, 1_000), chunks=(5, 100), dtype=int) s1[:,:] = np.random.randint(0, 100, (5, 1_000)) data = da.from_zarr(s1) analysed = zarry.create_group("analysed_data", overwrite=True) r1 = analysed.create_dataset('analysed1', shape=(5, 1_000), chunks=(5, 100), dtype=float) # computations data -= (data[2] + 1e-15) # add epsilon to avoid division by 0 data -= data.mean(axis=1, keepdims=True) # center data corrcoef_ = da.corrcoef(data) std_ = data.std(axis=1) # code in question pp_rec = da.to_zarr(data, r1, compute=False) _, corrcoef_, std_ = da.compute(pp_rec, corrcoef_, std_, optimized_graph=True)
(注:原代码中rec应为data,已修正)
在分布式环境运行时触发报错:
RuntimeError: Cannot store into in memory Zarr Array using the Distributed Scheduler.
目前可行方案是通过client.compute获取三个独立Future,但无法复用预处理中间结果,导致corrcoef和std重复计算,寻求更优方案。
解决方案
核心思路是复用预处理后的Dask Array,将存储操作转为分布式安全的任务,同时让统计计算共享中间结果,具体实现如下:
1. 分布式安全的存储方案(优先推荐)
分布式调度器禁止直接写入内存Zarr对象,优先改为写入Zarr路径(磁盘/云存储),同时基于同一预处理数组生成所有任务,让Dask自动优化依赖、共享计算:
from dask.distributed import Client import zarr import dask.array as da import numpy as np client = Client() # 初始化分布式客户端 # 加载数据(同本地逻辑) zarry = zarr.open('./example.zarr', mode='w') sample_data = zarry.create_group("sample_data", overwrite=True) s1 = sample_data.create_dataset('sample1', shape=(5, 1_000), chunks=(5, 100), dtype=int) s1[:,:] = np.random.randint(0, 100, (5, 1_000)) data = da.from_zarr(s1) analysed = zarry.create_group("analysed_data", overwrite=True) # 获取Zarr存储路径,而非内存对象 r1_path = f"./example.zarr/analysed_data/analysed1" analysed.create_dataset('analysed1', shape=(5, 1_000), chunks=(5, 100), dtype=float) # 定义预处理逻辑(核心复用对象) processed_data = data - (data[2] + 1e-15) processed_data = processed_data - processed_data.mean(axis=1, keepdims=True) # 生成所有任务 store_task = processed_data.to_zarr(r1_path, compute=False) corrcoef_task = da.corrcoef(processed_data) std_task = processed_data.std(axis=1) # 一次性提交任务,Dask自动共享中间计算 _, corrcoef_result, std_result = client.compute([store_task, corrcoef_task, std_task], optimize_graph=True) # 获取最终结果 corrcoef_result = corrcoef_result.result() std_result = std_result.result()
2. 必须使用内存Zarr的兼容方案
若业务场景要求写入内存Zarr数组,用dask.delayed包装写入操作,确保单任务执行,同时复用预处理结果:
from dask import delayed # 定义延迟写入函数 @delayed def write_to_memory_zarr(arr, zarr_ds): zarr_ds[:] = arr return None # 基于预处理数组生成任务 write_task = write_to_memory_zarr(processed_data.compute(), r1) corrcoef_task = da.corrcoef(processed_data) std_task = processed_data.std(axis=1) # 提交所有任务,共享预处理计算 _, corrcoef_result, std_result = client.compute([write_task, corrcoef_task, std_task], optimize_graph=True)
关键说明
- 所有任务基于同一个
processed_data生成,Dask图优化会自动复用中间计算步骤,彻底避免重复预处理 - 分布式环境下优先写入Zarr路径而非内存对象,这是调度器的安全操作模式
client.compute接受任务列表,会统一调度并优化依赖关系,效率远高于单独提交多个Future
内容的提问来源于stack exchange,提问作者Helmut
相关产品推荐
相关产品推荐

