TensorFlow中如何计算数据集所有特征的全局最大值
实现方案
你可以通过以下两种方式获取所有labels张量(即你提到的j)的全局最大值:
方法1:遍历数据集手动计算(适合小数据集)
直接迭代数据集的每一组数据,实时更新全局最大值即可,代码示例:
import tensorflow as tf input_slice=3 labels_slice=2 def split_window(x): inputs = tf.slice(x,[0], [input_slice]) labels = tf.slice(x,[input_slice], [labels_slice]) return inputs, labels dataset = tf.data.Dataset.range(1, 25 + 1).batch(5).map(split_window) # 初始化全局最大值,注意和张量数据类型匹配 global_max = tf.constant(-1, dtype=tf.int64) for i, j in dataset: current_max = tf.reduce_max(j) if current_max > global_max: global_max = current_max print(global_max.numpy()) # 输出结果为25
方法2:使用Dataset的reduce方法(更高效,适合大数据集)
利用TensorFlow数据集自带的规约API,不需要手动写循环,性能更优,代码示例:
import tensorflow as tf input_slice=3 labels_slice=2 def split_window(x): inputs = tf.slice(x,[0], [input_slice]) labels = tf.slice(x,[input_slice], [labels_slice]) return inputs, labels dataset = tf.data.Dataset.range(1, 25 + 1).batch(5).map(split_window) # 初始状态设为极小值,类型和数据集元素匹配 global_max = dataset.reduce( initial_state=tf.constant(-1, dtype=tf.int64), reduce_func=lambda state, batch: tf.maximum(state, tf.reduce_max(batch[1])) ) print(global_max.numpy()) # 输出结果为25
注意:如果你的实际场景中特征值可能出现负数,只需要把初始最大值的初始值调整为对应数据类型的最小值即可,比如tf.constant(tf.int64.min, dtype=tf.int64)。
内容的提问来源于stack exchange,提问作者freak11
相关产品推荐
相关产品推荐

