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

如何加快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)
      

内容的提问来源于stack exchange,提问作者Nathaniel Tan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 06:32:17