TensorFlow中对常量张量执行Argmax函数无法输出预期值
解决TensorFlow中tf.argmax输出张量信息而非预期数值的问题
嘿,我来帮你理清这个问题~你遇到的核心问题是混淆了TensorFlow的计算图节点(张量对象)和实际运行后的数值结果,咱们一步步拆解:
问题原因分析
你的代码里,先通过sess.run(vector)拿到了numpy数组v,接着用tf.argmax(v,1)创建了一个新的张量argm——但这个argm只是TensorFlow计算图里的一个待执行节点,你没有在会话中运行它,直接打印的话,输出自然是这个张量的元信息(名称、形状、数据类型),而不是你想要的最大值索引结果。
另外补充个小细节:你预期的[4,7,8]是子数组的最大值本身,但tf.argmax返回的是最大值的索引(从0开始),所以正确的索引结果应该是[3,3,1],对应每个子数组里的4、7、9的位置~
两种修正方案
方案一:直接在TensorFlow张量上操作,统一在会话中运行
既然已经用TensorFlow构建计算图,没必要先把张量转成numpy数组,直接在张量上调用tf.argmax,然后在会话中运行这个操作就能得到结果:
import tensorflow as tf vector = tf.constant([[1,2,3,4],[4,5,6,7],[8,9,1,2]],tf.int32,name="vector") argm = tf.argmax(vector, 1) # 直接对TensorFlow张量操作 with tf.Session() as sess: print(sess.run(argm)) # 在会话中运行argm节点,得到实际数值
方案二:用numpy处理已获取的数组
如果你已经通过sess.run拿到了numpy数组v,其实可以直接用numpy的argmax函数来计算,不用再创建TensorFlow张量:
import tensorflow as tf import numpy as np vector = tf.constant([[1,2,3,4],[4,5,6,7],[8,9,1,2]],tf.int32,name="vector") with tf.Session() as sess: v = sess.run(vector) argm = np.argmax(v, axis=1) # 使用numpy的argmax处理数组 print(argm)
额外小提示
你用的是TensorFlow 1.x的计算图模式,所有操作都需要在tf.Session()中运行才能得到实际数值;如果是TensorFlow 2.x,默认开启eager执行模式,不需要会话就能直接得到结果,写法会更简洁。
内容的提问来源于stack exchange,提问作者Jpmarulandas
相关产品推荐
相关产品推荐

