如何在TensorFlow中将二维数组每行最大值替换为1其余值置为0
问题原因
你的原代码将二维数组通过tf.reshape(x, (-1,))展平为一维数组,因此计算的是全局最大值,自然只能将全局最大值所在的位置置为1。
解决代码
你可以直接按行计算最大值再逐元素对比,以下是两种可行实现:
方法1:使用tf.reduce_max(逻辑最直观)
import tensorflow as tf x = tf.constant( [[0, 4], [2,3], [6,7], [9,2]]) # 计算每行最大值,keepdims保持二维结构方便后续广播对比 row_max = tf.reduce_max(x, axis=1, keepdims=True) # 逐元素判断是否等于该行最大值,转换为float64格式 output = tf.cast(tf.equal(x, row_max), tf.float64)
运行后output的输出为:
[[0. 1.] [0. 1.] [0. 1.] [1. 0.]]
方法2:使用tf.nn.top_k(适配仅需保留第一个最大值的场景)
如果一行存在多个相同最大值,且你只需要把第一个出现的最大值置为1,其余都为0,可以用该方法配合one_hot实现:
import tensorflow as tf x = tf.constant( [[0, 4], [2,3], [6,7], [9,2]]) # 取每行排名第一的数值和对应的索引 top_values, top_indices = tf.nn.top_k(x, k=1) # 将索引转为one_hot编码,即对应位置为1其余为0 output = tf.one_hot(tf.squeeze(top_indices, axis=1), depth=x.shape[1], dtype=tf.float64)
内容的提问来源于stack exchange,提问作者Hufsa
相关产品推荐
相关产品推荐

