Numpy如何根据类别标签筛选矩阵对应列提取目标子矩阵
实现方案
场景1:标签与列的映射规则可灵活调整
这种写法适配性强,就算后续标签对应列的规则变了,只要改映射字典即可:
import numpy as np # 已有数据 arr1 = np.array([[1,2,3,4,5,6,7,8,9],[10,11,12,13,14,15,16,17,18],[19,20,21,22,23,24,25,26,27]]) arr2 = np.array([["A"],["B"],["C"]]) # 定义标签到对应列索引的映射 label_col_map = { "A": [0, 1, 2], "B": [3, 4, 5], "C": [6, 7, 8] } # 逐行匹配列索引,构造形状为(3,3)的列索引数组 col_idx = np.array([label_col_map[label] for label in arr2.ravel()]) # 构造形状为(3,1)的行索引数组,利用广播机制匹配列索引形状 row_idx = np.arange(len(arr1))[:, None] # 二维高级索引取值 result = arr1[row_idx, col_idx]
运行后result的输出就是你需要的结果:
array([[ 1, 2, 3], [13, 14, 15], [25, 26, 27]])
场景2:标签对应列的规则固定(按每3列分组)
如果你的规则固定是A对应第一组3列、B对应第二组、C对应第三组,可以不用写映射字典,写法更简洁:
# 将arr1变形为 (行数, 分组数, 每组元素数) 的结构,此处为 (3,3,3) reshaped = arr1.reshape(arr1.shape[0], 3, 3) # 将标签转换为分组索引:A→0、B→1、C→2 group_idx = np.vectorize(lambda x: ord(x) - ord('A'))(arr2.ravel()) # 按行提取对应分组 result = reshaped[np.arange(arr1.shape[0]), group_idx]
报错原因说明
你之前写的arr[[0,1,2],[3,4,5],[6,7,8]]是针对三维数组的索引语法,你的arr1是二维数组,只能接收两个维度的索引参数,因此会触发形状不匹配的报错。二维高级索引需要传入分别对应行、列两个维度的索引数组,两个数组形状需要满足广播规则。
内容的提问来源于stack exchange,提问作者zachvac
相关产品推荐
相关产品推荐

