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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 09:30:23