三维数组中argmax函数的工作原理?附TensorFlow示例与结果
argmax在三维数组中的工作原理(结合TensorFlow示例)
tf.math.argmax的核心作用是找出指定轴(axis)上最大值对应的索引。对于三维数组,关键是明确你指定的axis对应哪一个维度——沿着这个维度遍历,对每个对应位置的元素比较大小,最终返回最大值所在的索引位置。
你的示例分析
先看输入张量x的结构,它的shape是(3, 2, 3),可以拆解为三个维度:
- axis=0:最外层的3个"子矩阵",每个子矩阵是2行3列的结构
- axis=1:每个子矩阵内部的2行
- axis=2:每行内部的3列
你的代码调用了tf.math.argmax(x, axis=0),也就是沿着axis=0的方向(即3个子矩阵的维度)做比较:
我们逐个对应位置看:
- 子矩阵的(0,0)位置元素:1(第一个子矩阵)、7(第二个)、13(第三个)→ 最大值13,在axis=0上的索引是2
- 子矩阵的(0,1)位置元素:2、8、14 → 最大值14,索引2
- 子矩阵的(0,2)位置元素:3、9、15 → 最大值15,索引2
- 子矩阵的(1,0)位置元素:4、10、16 → 最大值16,索引2
- 子矩阵的(1,1)位置元素:5、11、17 → 最大值17,索引2
- 子矩阵的(1,2)位置元素:6、12、18 → 最大值18,索引2
所以最终得到的结果是一个shape为(2, 3)的张量,所有元素都是2,和你给出的运行结果完全一致。
补充不同axis的示例
为了更清晰理解,再看两个不同axis的情况:
- 指定
axis=1(沿子矩阵的行维度比较):
print(tf.math.argmax(x, axis=1))
输出结果:
tf.Tensor( [[1 1 1] [1 1 1] [1 1 1]], shape=(3, 3), dtype=int64)
每个子矩阵中第二行的元素都比第一行大,所以索引都是1。
- 指定
axis=2(沿每行的列维度比较):
print(tf.math.argmax(x, axis=2))
输出结果:
tf.Tensor( [[2 2] [2 2] [2 2]], shape=(3, 2), dtype=int64)
每行的第三列元素都是最大的,所以索引都是2。
内容的提问来源于stack exchange,提问作者Shiva kumar das
相关产品推荐
相关产品推荐

