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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 15:00:53