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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 21:07:32