如何在xarray中实现类似NumPy的非正交高级索引?
在xarray中实现类似NumPy的高级索引(选取非正交点)
你希望从xarray的DataArray中选取若干非正交的点,实现类似NumPy的高级索引效果(如arr_np[:, x_idxs, y_idxs]),但xarray默认采用正交索引,导致常规方法无法得到预期结果。
示例场景
构建示例数组
from string import ascii_lowercase, ascii_uppercase import xarray as xr import numpy as np sizes = {"band": 4, "x": 5, "y": 6} shape = tuple(sizes.values()) dims = tuple(sizes.keys()) arr_np = np.arange(np.prod(shape)).reshape(shape) arr_xr = xr.DataArray( data=arr_np, dims=dims, coords={ "band": np.arange(100, 100 + sizes["band"]), "x": list(ascii_lowercase[: sizes["x"]]), # abcde... "y": list(ascii_uppercase[: sizes["y"]]), # ABCDE... }, )
目标采样点
points = [("a", "C"), ("d", "D"), ("e", "A")] points_idxs = [(0, 2), (3, 3), (4, 0)] xs, ys = map(list, zip(*points)) # ['a', 'd', 'e'] and ['C', 'D', 'A'] x_idxs, y_idxs = map(list, zip(*points_idxs)) # [0, 3, 4] and [2, 3, 0]
NumPy中的预期结果
expected = arr_np[:, x_idxs, y_idxs] assert expected.shape == (sizes["band"], len(points))
输出:
array([[ 2, 21, 24], [ 32, 51, 54], [ 62, 81, 84], [ 92, 111, 114]])
xarray常规方法的问题
使用arr_xr.loc[:, xs, ys]、arr_xr.sel(x=xs, y=ys)或arr_xr[:, x_idxs, y_idxs]会得到正交索引的结果,形状为(band:4, x:3, y:3),而非预期的(band:4, point:3)。
解决方案
方法1:堆叠维度后选择点(推荐)
将x和y维度堆叠为一个新的point维度,直接选择目标点元组,这是最直观的xarray原生方法:
# 堆叠x和y维度为point维度 stacked = arr_xr.stack(point=("x", "y")) # 选择目标点 result = stacked.sel(point=points) # 验证结果 assert result.shape == (sizes["band"], len(points)) print(result)
输出:
<xarray.DataArray (band: 4, point: 3)> array([[ 2, 21, 24], [ 32, 51, 54], [ 62, 81, 84], [ 92, 111, 114]]) Coordinates: * band (band) int64 100 101 102 103 * point (point) object ('a', 'C') ('d', 'D') ('e', 'A')
方法2:利用NumPy索引后包装为DataArray
直接使用NumPy的高级索引获取数据,再重新包装为xarray DataArray,保留原始坐标信息:
# 用NumPy索引获取数据 data = arr_np[:, x_idxs, y_idxs] # 包装为DataArray result = xr.DataArray( data=data, dims=["band", "point"], coords={ "band": arr_xr.band, "point": points } ) # 验证结果 assert result.shape == (sizes["band"], len(points))
方法3:使用xr.DataArray的高级索引(传递匹配形状的索引)
通过将索引数组转换为带有对应维度的DataArray,实现点式索引:
# 创建点维度的索引数组 point_dim = xr.DataArray(range(len(points)), dims="point") # 用x_idxs和y_idxs构建对应点的索引 x_idx_arr = xr.DataArray(x_idxs, dims="point") y_idx_arr = xr.DataArray(y_idxs, dims="point") # 执行索引 result = arr_xr.isel(x=x_idx_arr, y=y_idx_arr) # 验证结果 assert result.shape == (sizes["band"], len(points))
内容的提问来源于stack exchange,提问作者paime
相关产品推荐
相关产品推荐

