如何加快Dask数组转Numpy数组的compute运算速度?
我正在参与2023年EY开放数据科学挑战赛,需提取Sentinel-1-RTC卫星数据作为Keras CNN或SKLearn模型的输入。直接加载像素数据耗时过长,因此选择将VV和VH波段数据加载为Dask数组。以下是单个坐标点的示例代码:
import pystac_client import planetary_computer as pc from odc.stac import stac_load latlong = (10.323727047081501, 105.2516346045924) box_size_deg = 0.002 min_lon = float(latlong[1])-box_size_deg/2 min_lat = float(latlong[0])-box_size_deg/2 max_lon = float(latlong[1])+box_size_deg/2 max_lat = float(latlong[0])+box_size_deg/2 bbox = (min_lon , min_lat, max_lon, max_lat) time_slice = "2022-01-01/2022-12-31" scale = 10/111320.0 catalog = pystac_client.Client.open( "https://planetarycomputer.microsoft.com/api/stac/v1") search = catalog.search( collections=["sentinel-1-rtc"], bbox=bbox, datetime=time_slice) items = search.get_all_items() scale = 10/111320.0 test = stac_load(items, patch_url=pc.sign, bbox=bbox, bands=assets, chunks={}, crs="EPSG:4326", resolution=scale) print(test)
输出结果:
<xarray.Dataset> Dimensions: (latitude: 23, longitude: 23, time: 2) Coordinates: * latitude (latitude) float64 10.32 10.32 10.32 ... 10.32 10.32 10.32 * longitude (longitude) float64 105.3 105.3 105.3 ... 105.3 105.3 105.3 spatial_ref int32 4326 * time (time) datetime64[ns] 2022-01-09T22:46:06.347730 2022-01-10T... Data variables: vh (time, latitude, longitude) float32 dask.array<chunksize=(1, 23, 23), meta=np.ndarray> vv (time, latitude, longitude) float32 dask.array<chunksize=(1, 23, 23), meta=np.ndarray>
其中vh和vv变量的Dask数组仅约118kiB,但使用test.compute()将其转换为Numpy数组时,本地机器耗时超40秒。我需要处理600个坐标点,当前效率无法满足需求。Dask数组test.vv.data的任务图包含大量细碎任务,已尝试重新分块但无效果,同时也接受直接将Dask数组作为模型输入的建议。请问如何加快Dask转Numpy数组的运算速度?
调整分块策略,减少细碎任务
当前分块为(1,23,23),每个时间片单独成块导致任务数量过多。加载时指定合并时间维度的分块,直接将所有时间片整合为一个块,大幅降低任务数:test = stac_load(items, patch_url=pc.sign, bbox=bbox, bands=assets, chunks={"time": -1, "latitude":23, "longitude":23}, crs="EPSG:4326", resolution=scale)批量处理多个坐标点,减少重复请求
不要逐个处理600个坐标点,而是将相邻坐标点的bbox合并为更大区域,一次性加载后再切割成单个样本。这能减少STAC搜索、认证和数据请求的重复开销,降低网络延迟影响。启用Dask本地集群,利用多核资源
默认Dask用单线程执行,启动本地集群可调动全部CPU核心并行计算:from dask.distributed import Client, LocalCluster cluster = LocalCluster(n_workers=4, threads_per_worker=2) # 根据本地CPU核心数调整参数 client = Client(cluster)之后执行
test.compute()会自动分配任务到多个核心并行处理。预缓存数据到本地,避免重复下载
将加载后的数据集保存为Zarr格式,后续读取直接从本地获取,无需重复从Planetary Computer下载:# 保存数据到本地Zarr文件 test.to_zarr("sentinel_local_cache.zarr", mode="w") # 后续加载缓存数据 import xarray as xr test = xr.open_zarr("sentinel_local_cache.zarr")直接用Dask数组作为模型输入,跳过转Numpy步骤
- Scikit-learn场景:使用
dask_ml库,它支持直接传入Dask数组,自动并行处理数据,无需提前调用compute()。 - Keras/TensorFlow场景:将Dask数组转为TensorFlow的
tf.data.Dataset,实现边加载边训练,避免一次性占用大量内存:import tensorflow as tf from dask_tensorflow import from_dask # 将Dask数组转为TF数据集 tf_dataset = from_dask(test.vv).batch(1)
- Scikit-learn场景:使用
内容的提问来源于stack exchange,提问作者Nathaniel Tan

