如何获取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
相关产品推荐
相关产品推荐

