使用xarray.apply_ufunc调用np.linalg.pinv处理dask.array报错
xarray apply_ufunc 调用 pinv 报错,但 dask map_blocks 正常运行
问题场景
运行以下最小可复现代码时触发错误:
import xarray as xr import numpy as np from numpy.linalg import pinv import dask data = np.random.randn(4, 4, 3, 2) da = xr.DataArray(data=data, dims=("x", "y", "i", "j"),) da = da.chunk(x=1, y=1) da_inv = xr.apply_ufunc(pinv, da, input_core_dims=[["i", "j"]], output_core_dims=[["i", "j"]], exclude_dims=set(("i", "j")), dask = "parallelized", )
错误信息
Traceback (most recent call last): File "/glade/scratch/tomasc/tracer_inversion2/mwe.py", line 14, in <module> da_inv = xr.apply_ufunc(pinv, da, File "/glade/u/home/tomasc/miniconda3/envs/py310/lib/python3.10/site-packages/xarray/core/computation.py", line 1204, in apply_ufunc return apply_dataarray_vfunc( File "/glade/u/home/tomasc/miniconda3/envs/py310/lib/python3.10/site-packages/xarray/core/computation.py", line 315, in apply_dataarray_vfunc result_var = func(*data_vars) File "/glade/u/home/tomasc/miniconda3/envs/py310/lib/python3.10/site-packages/xarray/core/computation.py", line 771, in apply_variable_ufunc result_data = func(*input_data) File "/glade/u/home/tomasc/miniconda3/envs/py310/lib/python3.10/site-packages/xarray/core/computation.py", line 747, in func res = da.apply_gufunc( File "/glade/u/home/tomasc/miniconda3/envs/py310/lib/python3.10/site-packages/dask/array/gufunc.py", line 489, in apply_gufunc core_output_shape = tuple(core_shapes[d] for d in ocd) File "/glade/u/home/tomasc/miniconda3/envs/py310/lib/python3.10/site-packages/dask/array/gufunc.py", line 489, in <genexpr> core_output_shape = tuple(core_shapes[d] for d in ocd) KeyError: 'dim0'
可行对比代码
直接使用dask.array.map_blocks可以正常运行:
data_inv = dask.array.map_blocks(pinv, da.data).compute() # works!
问题原因
numpy.linalg.pinv的输入输出维度存在转置关系:输入是形状为(i,j)的矩阵(此处i=3、j=2),输出是形状为(j,i)的伪逆矩阵(即2,3)。
你在apply_ufunc中指定output_core_dims=[["i", "j"]],要求输出的核心维度顺序和输入一致,但实际pinv返回的维度顺序是j,i,导致xarray在映射维度名称时出现不匹配,触发KeyError。
而dask.map_blocks仅直接对块数据做运算,不处理维度名称的映射,因此不会出现该问题。
解决方案
方案1:修正输出核心维度顺序
显式指定output_core_dims为["j", "i"],匹配pinv的输出形状:
da_inv = xr.apply_ufunc( pinv, da, input_core_dims=[["i", "j"]], output_core_dims=[["j", "i"]], # 对应pinv输出的维度顺序 exclude_dims=set(("i", "j")), dask="parallelized", output_dtypes=[da.dtype] # 显式指定输出类型,避免推断错误 )
方案2:保持输入输出维度顺序一致
如果需要输出的维度顺序仍为i,j,可以包装pinv函数,对结果做转置:
def pinv_transposed(arr): return pinv(arr).T da_inv = xr.apply_ufunc( pinv_transposed, da, input_core_dims=[["i", "j"]], output_core_dims=[["i", "j"]], exclude_dims=set(("i", "j")), dask="parallelized", output_dtypes=[da.dtype] )
内容的提问来源于stack exchange,提问作者TomCho
相关产品推荐
相关产品推荐

