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

如何获取Tensor中argmax索引对应的维度实际值(保持输出形状一致)

获取TensorFlow中argmax索引对应的值

嘿,这个需求我之前也碰到过,给你两种实用的实现方式,按需选择就行:

方法一:基于已有的argmax索引提取(适合需要保留索引的场景)

如果你已经计算出了argmax索引,想要基于这些索引提取对应的值,可以用tf.gather_nd来实现。关键是要构造出每个索引对应的三维坐标:

import tensorflow as tf

# 初始化示例Tensor
tensor = tf.random.uniform((60, 128, 30000))
# 计算axis=2上的argmax索引
argmax_indices = tf.argmax(tensor, axis=2)

# 生成对应第一、第二维度的坐标网格
i, j = tf.meshgrid(tf.range(tensor.shape[0]), tf.range(tensor.shape[1]), indexing='ij')
# 将网格坐标和argmax索引组合成三维坐标数组(形状为(60,128,3))
coordinates = tf.stack([i, j, argmax_indices], axis=-1)
# 根据坐标提取对应的值
max_values = tf.gather_nd(tensor, coordinates)

# 验证形状
print(max_values.shape)  # 输出 (60, 128)

这里用indexing='ij'是为了让生成的网格坐标和原Tensor的维度顺序一致,确保每个坐标(i,j,argmax_indices[i][j])能精准定位到原Tensor中的对应元素。

方法二:直接提取最大值(更高效,无需先算argmax)

如果你的最终目标只是获取axis=2上的最大值,不需要保留argmax索引,那直接用tf.reduce_max会更简单高效,一步到位:

import tensorflow as tf

tensor = tf.random.uniform((60, 128, 30000))
# 直接提取axis=2上的最大值
max_values = tf.reduce_max(tensor, axis=2)

print(max_values.shape)  # 输出 (60, 128)

这个方法的结果和方法一是完全一致的,因为argmax对应的就是最大值的位置,reduce_max会直接计算该维度的最大值,省去了构造坐标和索引提取的步骤,性能更好。


内容的提问来源于stack exchange,提问作者OK 400

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.01 00:48:12