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

三维数组中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的情况:

  1. 指定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。

  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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 17:17:14