如何在Xarray中对分块数据集并行应用输出尺寸不同的函数?
解决xarray分块数据集上输入输出维度不同的并行像素处理问题
问题核心
你有维度为(time, x, y)的分块数据集,需对每个(x,y)像素的完整时间序列应用插值函数,输出更短的时间序列。尝试u_funcs和apply_along_axis时因输入输出尺寸不匹配失败,且无法利用分块并行能力。
解决方案
方法1:用xarray.apply_ufunc实现并行处理
apply_ufunc支持自定义输入输出维度,结合numpy.vectorize可将单像素函数向量化,同时指定输出的新时间维度信息,适配分块数据集的并行计算。
import xarray as xr import numpy as np from scipy.interpolate import interp1d import pandas as pd # 创建示例分块数据集 dates = pd.date_range(start="2023-01-01", end="2023-05-01") data = np.random.rand(len(dates), 10, 10) dataset = xr.Dataset( {"v": (["time", "x", "y"], data)}, coords={"time": dates, "x": range(10), "y": range(10)}, ).chunk({'x':5, 'y':5}) # 仅对空间维度分块,保留完整时间序列 # 单像素插值函数 def interpolate_5_days(pixel): original_indices = np.arange(len(pixel)) interpolated_indices = np.arange(0, len(pixel), 5) interpolator = interp1d(original_indices, pixel, kind='linear') return interpolator(interpolated_indices) # 向量化函数,声明输入输出维度格式 vectorized_interp = np.vectorize( interpolate_5_days, signature='(n)->(m)' # 输入n维数组,输出m维数组 ) # 计算输出时间轴的长度和坐标 output_time_len = len(np.arange(0, len(dates), 5)) output_time_coords = dates[::5] # 用apply_ufunc并行处理 result = xr.apply_ufunc( vectorized_interp, dataset['v'], input_core_dims=[['time']], # 指定输入核心维度为time output_core_dims=[['new_time']], # 指定输出的新核心维度 output_sizes={'new_time': output_time_len}, # 定义输出维度长度 vectorize=True, dask='parallelized', # 启用dask并行计算 output_dtypes=[dataset['v'].dtype] ) # 给结果添加坐标并重命名维度 result = result.rename({'new_time': 'time'}) result = result.assign_coords(time=output_time_coords) print(result)
方法2:用xarray.map_blocks处理分块数据
map_blocks允许对每个数据块应用自定义函数,适合维度变化场景,需明确每个块的输出结构。
import xarray as xr import numpy as np from scipy.interpolate import interp1d import pandas as pd # 创建示例分块数据集 dates = pd.date_range(start="2023-01-01", end="2023-05-01") data = np.random.rand(len(dates), 10, 10) dataset = xr.Dataset( {"v": (["time", "x", "y"], data)}, coords={"time": dates, "x": range(10), "y": range(10)}, ).chunk({'x':5, 'y':5}) # 单像素插值函数 def interpolate_5_days(pixel): original_indices = np.arange(len(pixel)) interpolated_indices = np.arange(0, len(pixel), 5) interpolator = interp1d(original_indices, pixel, kind='linear') return interpolator(interpolated_indices) # 定义单个数据块的处理函数 def process_block(block): # 对块内所有(x,y)像素应用插值 interpolated_data = np.apply_along_axis(interpolate_5_days, axis=0, arr=block['v'].values) # 生成输出时间坐标 output_time = dates[::5] # 返回新的数据集块 return xr.Dataset( {"v": (["time", "x", "y"], interpolated_data)}, coords={"time": output_time, "x": block['x'], "y": block['y']} ) # 用map_blocks处理所有分块 result = dataset.map_blocks(process_block).compute() print(result)
关键说明
- 维度匹配:通过
input_core_dims和output_core_dims明确输入输出的核心维度,output_sizes定义新维度长度,解决输入输出尺寸不匹配问题。 - 分块策略:仅对
x、y空间维度分块,保留time维度完整性,确保每个像素能获取完整时间序列插值。 - 并行能力:两种方法均支持dask并行,自动利用分块数据进行多进程/线程计算,适配大规模数据集。
内容的提问来源于stack exchange,提问作者Nihilum
相关产品推荐
相关产品推荐

