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

Numpy indexing与broadcast应用:高维张量指定位置元素批量选取方法

解法

你可以通过构造三个可广播对齐的索引数组,一次性完成取值操作,核心写法如下:

import numpy as np

# 示例张量a构造
a = np.array([
    [[-1.054,  0.068, -0.572,  1.535,  1.746],
     [-0.115,  0.356,  0.222, -0.391,  0.367],
     [-0.53 , -0.856,  0.58 ,  1.099,  0.605]],
    [[ 0.31 ,  0.037, -0.85 , -0.054, -0.75 ],
     [-0.097, -1.707, -0.702,  0.658,  0.548],
     [ 1.727, -0.326, -1.525, -0.656,  0.349]]
])

# 一次性索引写法
res = a[
    [[0], [1]], 
    range(3), 
    [[0,2,4], [1,3,2]]
]
print(res)

运行输出结果和分开计算的结果完全一致:

[[-1.054  0.222  0.605]
 [ 0.037  0.658 -1.525]]

原理说明

numpy的高级索引要求各维度的索引数组形状可广播对齐:

  • 第一维度的索引[[0],[1]]形状为(2, 1),对应要取的第0、第1个3×5子张量
  • 第二维度的索引range(3)(即[0,1,2])形状为(3,),对应每个子张量的3行
  • 第三维度的索引[[0,2,4], [1,3,2]]形状为(2, 3),对应每行要取的列位置

三者广播后形状都对齐为(2, 3),每个位置的三组索引刚好对应目标元素位置。

如果需要适配更高维度的批量场景,可以用更通用的写法生成索引:

batch_idx = np.arange(a.shape[0])[:, np.newaxis]  # 自动生成形状为(batch_size, 1)的批次索引
row_idx = np.arange(a.shape[1])
col_idx = np.array([[0,2,4], [1,3,2]])
res = a[batch_idx, row_idx, col_idx]

内容的提问来源于stack exchange,提问作者Sam-gege

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 04:36:03