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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 20:36:02