tf.argmax函数使用报错求助:Python3+TF1.3环境代码问题排查
问题分析与解决
你这段代码的问题出在tf.argmax()的axis参数上。在TensorFlow 1.x版本中,tf.argmax的第二个参数指定的是要沿着哪个维度寻找最大值的索引,而你传入的输入张量[1,0,0]是一个一维张量(只有维度0),但你却指定了axis=1,这超出了该张量的维度范围,所以会触发错误。
修正后的代码
import tensorflow as tf a = tf.argmax([1,0,0], 0) # 将axis改为0,或者用-1(表示最后一个维度) with tf.Session() as sess: print(sess.run(a))
补充说明
- 对于一维张量,可用的axis值只有
0或者-1(两者等价),运行修正后的代码会输出0,也就是最大值1所在的索引位置。 - 如果你的实际需求是处理更高维度的张量(比如二维矩阵),那
axis=1才是合理的——比如对于形状为[batch_size, num_classes]的张量,axis=1会沿着每个样本的类别维度取最大值索引。
内容的提问来源于stack exchange,提问作者David_tut
相关产品推荐
相关产品推荐

