使用xarray.apply_ufunc处理分块DataArray时出现维度错误
问题
尝试用xarray结合NumPy的digitize函数对多个netCDF4文件的降水数据进行分箱处理,使用xarray.open_mfdataset读取大量文件,但调用xarray.apply_ufunc时出现维度不匹配的错误。
代码示例
import numpy as np import xarray as xr precip_bins = np.arange(0, 120, .1) # mm/hr, 0.1mm/hr bins ds = xr.open_mfdataset(['/home/data/2022_rain.nc', '/home/data/2023_rain.nc'], chunks={'time': -1, 'latitude': 500, 'longitude': 700}) hourly_precip_ds = ds['precipitation'] print(hourly_precip_ds) # Digitize the data into bins bin_idxs = xr.apply_ufunc(np.digitize, hourly_precip_ds, precip_bins, input_core_dims=[['time'], []], output_core_dims=[['time']], dask="parallelized", dask_gufunc_kwargs={'allow_rechunk': True}) - 1 # Subtract 1 to make bins 0-indexed
运行报错输出
<xarray.DataArray 'precipitation' (time: 17481, latitude: 500, longitude: 700)> Size: 24GB dask.array<concatenate, shape=(17481, 500, 700), dtype=float32, chunksize=(8760, 500, 700), chunktype=numpy.ndarray> Coordinates: * time (time) datetime64[ns] 140kB 2022-01-01 ... 2023-12-31T23:00:00 * latitude (latitude) float32 2kB 69.95 69.85 69.75 ... 20.25 20.15 20.05 * longitude (longitude) float32 3kB -24.95 -24.85 -24.75 ... 44.85 44.95 Attributes: units: mm/hr --------------------------------------------------------------------------- ValueError Traceback (most recent call last) Cell In[9], line 2 1 # Digitize the data into bins ----> 2 bin_idxs = xr.apply_ufunc(np.digitize, 3 hourly_precip_ds, 4 precip_bins, 5 input_core_dims=[['time'], []], 6 output_core_dims=[['time']], 7 dask="parallelized", 8 dask_gufunc_kwargs={'allow_rechunk': True}) - 1 # Subtract 1 to make bins 0-indexed 9 print("++++++++++++++++++++++++++++++++++++++++++++++++") 10 print(bin_idxs) File ~/.conda/envs/geo_jupyter_env/lib/python3.12/site-packages/xarray/core/computation.py:1278, in apply_ufunc(func, input_core_dims, output_core_dims, exclude_dims, vectorize, join, dataset_join, dataset_fill_value, keep_attrs, kwargs, dask, output_dtypes, output_sizes, meta, dask_gufunc_kwargs, on_missing_core_dim, *args) 1276 # feed DataArray apply_variable_ufunc through apply_dataarray_vfunc 1277 elif any(isinstance(a, DataArray) for a in args): -> 1278 return apply_dataarray_vfunc( 1279 variables_vfunc, 1280 *args, 1281 signature=signature, 1282 join=join, 1283 exclude_dims=exclude_dims, 1284 keep_attrs=keep_attrs, 1285 ) 1286 # feed Variables directly through apply_variable_ufunc 1287 elif any(isinstance(a, Variable) for a in args): File ~/.conda/envs/geo_jupyter_env/lib/python3.12/site-packages/xarray/core/computation.py:320, in apply_dataarray_vfunc(func, signature, join, exclude_dims, keep_attrs, *args) 315 result_coords, result_indexes = build_output_coords_and_indexes( 316 args, signature, exclude_dims, combine_attrs=keep_attrs 317 ) 319 data_vars = [getattr(a, "variable", a) for a in args] -> 320 result_var = func(*data_vars) 322 out: tuple[DataArray, ...] | DataArray 323 if signature.num_outputs > 1: File ~/.conda/envs/geo_jupyter_env/lib/python3.12/site-packages/xarray/core/computation.py:831, in apply_variable_ufunc(func, signature, exclude_dims, dask, output_dtypes, vectorize, keep_attrs, dask_gufunc_kwargs, *args) 826 if vectorize: 827 func = _vectorize( 828 func, signature, output_dtypes=output_dtypes, exclude_dims=exclude_dims 829 ) -> 831 result_data = func(*input_data) 833 if signature.num_outputs == 1: 834 result_data = (result_data,) File ~/.conda/envs/geo_jupyter_env/lib/python3.12/site-packages/xarray/core/computation.py:808, in apply_variable_ufunc.<locals>.func(*arrays) 807 def func(*arrays): -> 808 res = chunkmanager.apply_gufunc( 809 numpy_func, 810 signature.to_gufunc_string(exclude_dims), 811 *arrays, 812 vectorize=vectorize, 813 output_dtypes=output_dtypes, 814 **dask_gufunc_kwargs, 815 ) 817 return res File ~/.conda/envs/geo_jupyter_env/lib/python3.12/site-packages/xarray/namedarray/daskmanager.py:155, in DaskManager.apply_gufunc(self, func, signature, axes, axis, keepdims, output_dtypes, output_sizes, vectorize, allow_rechunk, meta, *args, **kwargs) 138 def apply_gufunc( 139 self, 140 func: Callable[..., Any], (...) 151 **kwargs: Any, 152 ) -> Any: 153 from dask.array.gufunc import apply_gufunc -> 155 return apply_gufunc( 156 func, 157 signature, 158 *args, 159 axes=axes, 160 axis=axis, 161 keepdims=keepdims, 162 output_dtypes=output_dtypes, 163 output_sizes=output_sizes, 164 vectorize=vectorize, 165 allow_rechunk=allow_rechunk, 166 meta=meta, 167 **kwargs, 168 ) File ~/.conda/envs/geo_jupyter_env/lib/python3.12/site-packages/dask/array/gufunc.py:431, in apply_gufunc(func, signature, axes, axis, keepdims, output_dtypes, output_sizes, vectorize, allow_rechunk, meta, *args, **kwargs) 428 for dim, sizes in dimsizess.items(): 429 #### Check that the arrays have same length for same dimensions or dimension `1` 430 if set(sizes) | {1} != {1, max(sizes)}: -> 431 raise ValueError(f"Dimension `'{dim}'` with different lengths in arrays") 432 if not allow_rechunk: 433 chunksizes = chunksizess[dim] ValueError: Dimension `'__loopdim1__'` with different lengths in arrays
解决方案
错误根源是np.digitize的bins参数作为普通NumPy数组传入时,Dask无法正确识别其维度,导致分块处理时维度不匹配。以下两种方法可解决问题:
方法一:将bins包装为xarray.DataArray
明确bins的维度信息,让apply_ufunc正确解析输入结构:
import numpy as np import xarray as xr precip_bins = np.arange(0, 120, .1) # 将bins转为带维度的DataArray,指定虚拟维度名'bin' precip_bins_da = xr.DataArray(precip_bins, dims=['bin']) ds = xr.open_mfdataset(['/home/data/2022_rain.nc', '/home/data/2023_rain.nc'], chunks={'time': -1, 'latitude': 500, 'longitude': 700}) hourly_precip_ds = ds['precipitation'] bin_idxs = xr.apply_ufunc( np.digitize, hourly_precip_ds, precip_bins_da, input_core_dims=[['time'], ['bin']], output_core_dims=[['time']], dask="parallelized", dask_gufunc_kwargs={'allow_rechunk': True} ) - 1
方法二:启用vectorize参数
开启自动向量化模式,让apply_ufunc自动适配一维输入的广播逻辑:
import numpy as np import xarray as xr precip_bins = np.arange(0, 120, .1) ds = xr.open_mfdataset(['/home/data/2022_rain.nc', '/home/data/2023_rain.nc'], chunks={'time': -1, 'latitude': 500, 'longitude': 700}) hourly_precip_ds = ds['precipitation'] bin_idxs = xr.apply_ufunc( np.digitize, hourly_precip_ds, precip_bins, input_core_dims=[['time'], []], output_core_dims=[['time']], dask="parallelized", vectorize=True, # 新增vectorize参数 dask_gufunc_kwargs={'allow_rechunk': True} ) - 1
补充说明
- 方法一通过显式定义维度,避免Dask自动推断出错,逻辑更清晰;
- 方法二适合简单场景,无需修改输入结构,自动处理维度适配;
- 两种方案均能解决维度不匹配问题,可根据实际代码复杂度选择。
内容的提问来源于stack exchange,提问作者Innocuous Rift
相关产品推荐
相关产品推荐

