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

如何用索引数组对Numpy数组第二维度切片得到指定结果?

解决NumPy三维数组按指定索引切片的问题

首先,我们先明确原数组的结构:

import numpy as np
a = np.arange(24).reshape(4,3,2)
index_dim2 = np.array([0,1,2,2])

a的形状是(4,3,2),可以拆分为4个形状为(3,2)的子数组:

  • a[0] = [[0,1], [2,3], [4,5]]
  • a[1] = [[6,7], [8,9], [10,11]]
  • a[2] = [[12,13], [14,15], [16,17]]
  • a[3] = [[18,19], [20,21], [22,23]]

我们的目标是对每个第一维度的子数组(共4个),提取第二维度中对应index_dim2索引的元素,并且保持结果为(4,1,2)的三维结构(和目标输出一致)。

正确的切片表达式

你需要的切片表达式是:

a[np.arange(4), index_dim2[:, np.newaxis], :]

或者更简洁的写法(None等价于np.newaxis):

a[range(4), index_dim2[:, None], :]

为什么这样写?

  1. np.arange(4):遍历第一维度的所有索引(0到3),对应4个子数组。
  2. index_dim2[:, np.newaxis]:把原本形状为(4,)的索引数组转换为(4,1),这样在索引时会和第一维度的索引广播,让每个第一维度的元素只提取第二维度中指定的单个索引,同时保留一个长度为1的第二维度,确保结果是三维的。
  3. ::提取第三维度的所有元素(每个子元素的两个数值)。

验证结果

运行这段代码:

import numpy as np
a = np.arange(24).reshape(4,3,2)
index_dim2 = np.array([0,1,2,2])
b = a[np.arange(4), index_dim2[:, np.newaxis], :]
print(b)

输出结果完全符合你的目标:

array([[[ 0,  1]],
       [[ 8,  9]],
       [[16, 17]],
       [[22, 23]]])

另外,如果你先得到(4,2)的二维数组再扩展维度,也可以达到同样效果,比如:

b = a[np.arange(4), index_dim2, :][:, np.newaxis]

但第一种方式更直接,一步到位完成索引和维度保留。

内容的提问来源于stack exchange,提问作者cass

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 03:55:44