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

如何用Dask并行化基于Xarray与Numba的多维数组逐元计算

问题:基于Xarray/Dask实现多维数据的并行计算

背景与需求

现有维度为(nt, nb, nx, ny)的Xarray DataArray数据,需针对第0维度的不同取值,对nx和ny维度的每个单元计算相关量。该计算可在nt、nx、ny维度独立执行,希望利用Dask实现并行化并充分利用数据分块结构,避免串行执行的低效问题。

以下是串行执行的示例代码(实际计算逻辑更复杂):

import numpy as np
import xarray as xr
import xarray.tutorial
from numba import njit, float32
from itertools import product

@njit('Tuple((float32[:, :],float32[:,:]))(float32[:, :, :], float32[:, :,:])')
def do_smthg(ar1, ar2):
    n1, n2, n3 = ar1.shape
    outa = np.zeros((n2, n3), dtype=np.float32)
    outb = np.zeros((n2, n3), dtype=np.float32)
    for i in range(n1):
        for j in range(n2):
            outa[i,j] = np.sum(ar1[:, i,j] - ar2[:, i,j])
            outb[i,j] = np.sum(ar1[:, i,j] + ar2[:, i,j])
    return outa, outb
    
da = xr.tutorial.load_dataset("era5-2mt-2019-03-uk.grib")
da = da.chunk("auto")
F = {}
for (t1,tt1), (t2, tt2) in product(da.t2m.groupby("time.day"),
                           da.t2m.groupby("time.day")):
    if t2 > t1:
        F[(t1, t2)] = do_smthg(tt1.values, tt2.values)

现有并行尝试的问题

方案1:直接使用Dask Client提交任务

该方案可运行,但存在大量数据传输开销,无法充分利用Xarray/Dask的分块优化,不适合大型集群与超大分块数据集:

from distributed import LocalCluster, Client
cluster = LocalCluster()
client = Client(cluster)
F = {}
for (t1,tt1), (t2, tt2) in product(da.t2m.groupby("time.day"),
                           da.t2m.groupby("time.day")):
    if t2 > t1:
        F[(t1, t2)] = client.submit(do_smthg, tt1.values, tt2.values)
F = {k:v.result() for k,v in F.items()}

方案2:使用Xarray map_blocks(报错)

尝试用map_blocks实现并行,但遇到两类错误:

# 模板输出数据集
out = xr.Dataset(
    data_vars={"outa":(["lat", "lon"], np.random.rand(33, 49)),
               "outb":(["lat", "lon"], np.random.rand(33, 49))})
out.coords["lat"] = da.coords["latitude"].values
out.coords["lon"] = da.coords["longitude"].values
out = out.chunk("auto")

F = {}
for (t1,tt1), (t2, tt2) in product(da.t2m.groupby("time.day"),
                           da.t2m.groupby("time.day")):
    if t2 > t1:
        F[(t1, t2)] = tt1.drop("time").map_blocks(do_smthg, args=[tt2.drop("time")], template=out)
F[(1,5)].outb.values
  • 带Numba装饰器时的错误:

TypeError: No matching definition for argument type(s) pyobject, pyobject

  • 移除Numba装饰器后的错误:

~/mambaforge/lib/python3.9/site-packages/dask/core.py in _execute_task(arg, cache, dsk)
117 # temporaries by their reference count and can execute certain
118 # operations in-place.
--> 119 return func(*(_execute_task(a, cache) for a in args))
120 elif not ishashable(arg):
121 return arg

~/mambaforge/lib/python3.9/site-packages/xarray/core/parallel.py in _wrapper(func, args, kwargs, arg_is_array, expected)
286
287 # check all dims are present
--> 288 missing_dimensions = set(expected["shapes"]) - set(result.sizes)
289 if missing_dimensions:
290 raise ValueError(

AttributeError: 'numpy.ndarray' object has no attribute 'sizes'

可行解决方案

问题核心在于:map_blocks更适合单数组块处理,而多输入场景下apply_ufunc更适配;同时需让函数返回Xarray对象而非numpy数组,并修正Numba函数的逻辑问题。

步骤1:修正Numba计算函数

原函数存在索引越界问题(outa维度为(n2,n3),却用i循环索引),先修正逻辑:

@njit('Tuple((float32[:, :],float32[:,:]))(float32[:, :, :], float32[:, :,:])')
def do_smthg(ar1, ar2):
    # ar1/ar2 shape: (n_time, n_lat, n_lon)
    n_time, n_lat, n_lon = ar1.shape
    outa = np.zeros((n_lat, n_lon), dtype=np.float32)
    outb = np.zeros((n_lat, n_lon), dtype=np.float32)
    # 遍历每个经纬度单元
    for j in range(n_lat):
        for k in range(n_lon):
            outa[j, k] = np.sum(ar1[:, j, k] - ar2[:, j, k])
            outb[j, k] = np.sum(ar1[:, j, k] + ar2[:, j, k])
    return outa, outb

步骤2:包装函数适配Xarray

编写包装函数,将Xarray对象转为numpy数组计算,再返回Xarray Dataset以适配并行接口:

def xr_do_smthg(da1, da2):
    ar1 = da1.data
    ar2 = da2.data
    outa_np, outb_np = do_smthg(ar1, ar2)
    # 保留原始坐标
    return xr.Dataset(
        {
            "outa": (["latitude", "longitude"], outa_np),
            "outb": (["latitude", "longitude"], outb_np)
        },
        coords={"latitude": da1.latitude, "longitude": da1.longitude}
    )

步骤3:用apply_ufunc实现并行

apply_ufunc支持多输入,可自动处理Dask分块,无需手动管理数据传输:

# 预分组日数据
daily_groups = list(da.t2m.groupby("time.day"))
F = {}

for (t1, tt1), (t2, tt2) in product(daily_groups, daily_groups):
    if t2 > t1:
        result = xr.apply_ufunc(
            xr_do_smthg,
            tt1,
            tt2,
            input_core_dims=[["time"], ["time"]],  # 指定需要计算的核心维度
            output_core_dims=[["latitude", "longitude"], ["latitude", "longitude"]],
            dask="parallelized",  # 启用Dask并行
            output_dtypes=[np.float32, np.float32],
        )
        F[(t1, t2)] = result

# 触发计算(按需执行)
# F[(1,5)].compute()

该方案会自动利用Dask的分块结构,在集群上高效并行计算,避免不必要的数据传输。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 12:10:37