如何在NumPy中利用argmax结果矩阵索引概率矩阵?
用索引矩阵提取概率矩阵对应元素的NumPy实现
给定形状为(2, 3)的索引矩阵m,以及形状为(2, 3, 2)的概率矩阵p,我们需要提取p中每个位置对应m指定索引的概率值,得到与m同形状的结果矩阵。
方法一:使用np.take_along_axis(推荐)
这是NumPy专门为按轴索引提取场景设计的函数,代码简洁直观:
import numpy as np # 定义输入矩阵 m = np.array( [ [0, 1, 0], [1, 0, 1] ] ) p = np.array( [ [ [0.6, 0.4], [0.3, 0.7], [0.8, 0.2] ], [ [0.35, 0.65], [0.7, 0.3], [0.1, 0.9] ], ] ) # 核心操作:为索引矩阵增加维度以匹配概率矩阵的轴,提取后压缩多余维度 result = np.take_along_axis(p, m[..., np.newaxis], axis=2).squeeze(axis=2) print(result)
运行结果:
[[0.6 0.7 0.8] [0.65 0.7 0.9]]
解释:
m[..., np.newaxis]将m从(2,3)转换为(2,3,1),确保与p的前两维维度对齐axis=2指定在概率矩阵的第三个维度(即每个位置的概率列表维度)上按索引提取squeeze(axis=2)移除提取后多余的单维度,得到与m形状一致的结果
方法二:使用广播式三维索引
通过生成行列索引并利用NumPy的广播特性,直接定位到目标元素:
import numpy as np # 定义输入矩阵(同方法一) m = np.array([[0,1,0],[1,0,1]]) p = np.array([[[0.6,0.4],[0.3,0.7],[0.8,0.2]],[[0.35,0.65],[0.7,0.3],[0.1,0.9]]]) # 生成行列索引矩阵 rows = np.arange(m.shape[0])[:, np.newaxis] # 形状(2,1) cols = np.arange(m.shape[1]) # 形状(3,) # 广播后三维索引提取 result = p[rows, cols, m] print(result)
解释:
rows和cols通过广播会自动扩展为(2,3)的形状,与m完全匹配p[rows, cols, m]直接定位到每个(行,列)位置下,m指定索引的概率值
内容的提问来源于stack exchange,提问作者Denis Logvinenko
相关产品推荐
相关产品推荐

