首页 > 代码库 > TF-tf.arg_max 介绍
TF-tf.arg_max 介绍
定义为
def arg_max(input, dimension, name=None)
作用是取行或者列的最大值的位置。
input:
类型为 float32, float64, int64, int32, uint8, uint16, int16, int8, complex64, complex128, qint8, quint8, qint32, half的tensor
dimension:
必须为int32, int64. int32,取值为0或1.
name:
名字
returns:
返回一个tensor
例如以下测试
import tensorflow as tf list_a = [[1,2,3,4,5], [3,3,3,1,6], [5,1,2,1,1]] sess = tf.InteractiveSession() argmax0 = tf.arg_max(list_a, 0) print("argmax 0={}".format(argmax0.eval())) argmax1 = tf.arg_max(list_a, 1) print("argmax 1={}".format(argmax1.eval()))
结果为
argmax 0=[2 1 0 0 1]
argmax 1=[4 4 0]
TF-tf.arg_max 介绍
声明:以上内容来自用户投稿及互联网公开渠道收集整理发布,本网站不拥有所有权,未作人工编辑处理,也不承担相关法律责任,若内容有误或涉及侵权可进行投诉: 投诉/举报 工作人员会在5个工作日内联系你,一经查实,本站将立刻删除涉嫌侵权内容。