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

Numpy 3D张量掩码索引维度丢失问题求解

Numpy 3D张量掩码索引维度丢失问题解决方案

你当前使用的二维掩码直接索引3D张量时,Numpy会将所有True对应的行扁平化提取,导致维度丢失。要在保留目标维度结构的同时实现需求,且不使用reshape或手动行索引,可以通过以下基于掩码的高级索引方式实现:

代码实现

import numpy as np

# 生成原始3D张量
array = np.repeat(np.arange(15).reshape(3,5)[None,:], 3, axis=0)
# 定义掩码
mask = np.array([[False, True, True],
                 [True, False, True],
                 [False, True, True]])

# 核心操作:结合广播的高级索引
result = array[np.arange(3)[:, None], mask]

输出验证

print(result)
# 输出:
# array([[[ 5,  6,  7,  8,  9],
#         [10, 11, 12, 13, 14]],
# 
#        [[ 0,  1,  2,  3,  4],
#         [10, 11, 12, 13, 14]],
# 
#        [[ 0,  1,  2,  3,  4],
#         [ 5,  6,  7,  8,  9]]])

原理说明

  • np.arange(3)[:, None]生成形状为(3,1)的索引数组,对应3D张量的第一个维度(3个独立矩阵)。
  • 掩码mask为(3,3),和上述索引数组广播后,会为每个矩阵精准筛选出需要保留的2行,最终得到(3,2,5)的目标形状。
  • 该方案完全基于掩码规则实现,性能和纯布尔掩码索引一致,同时规避了维度丢失问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 22:35:14