如何在tf.math.bincount中用权重的最大/最小值替代权重求和?
实现按values分组取权重的最大/最小值
TensorFlow的tf.math.bincount仅支持对相同values对应的weights求和,若要实现取最大或最小值,可以用分段聚合函数替代,具体实现如下:
取最大值的实现
import tensorflow as tf values = tf.constant([1,1,2,3,2,4,4,5]) weights = tf.constant([1,5,0,1,0,5,4,5]) # 确定结果长度(与bincount行为一致,为max(values)+1) num_segments = tf.reduce_max(values) + 1 # 按values分组,对每组weights取最大值 max_result = tf.math.unsorted_segment_max(weights, values, num_segments) # 将未出现的索引(对应值为int类型最小值/浮点类型-inf)替换为0 if weights.dtype.is_integer: max_result = tf.where(max_result == tf.constant(tf.int32.min, dtype=weights.dtype), tf.zeros_like(max_result), max_result) else: max_result = tf.where(tf.math.is_inf(max_result), tf.zeros_like(max_result), max_result) print(max_result.numpy()) # 输出:[0 5 0 1 5 5]
取最小值的实现
只需将上述代码中的tf.math.unsorted_segment_max替换为tf.math.unsorted_segment_min,并对应处理未出现索引的初始值:
import tensorflow as tf values = tf.constant([1,1,2,3,2,4,4,5]) weights = tf.constant([1,5,0,1,0,5,4,5]) num_segments = tf.reduce_max(values) + 1 min_result = tf.math.unsorted_segment_min(weights, values, num_segments) # 将未出现的索引(对应值为int类型最大值/浮点类型+inf)替换为0 if weights.dtype.is_integer: min_result = tf.where(min_result == tf.constant(tf.int32.max, dtype=weights.dtype), tf.zeros_like(min_result), min_result) else: min_result = tf.where(tf.math.is_inf(min_result), tf.zeros_like(min_result), min_result) print(min_result.numpy()) # 输出:[0 1 0 1 4 5]
关键说明
tf.math.unsorted_segment_max/min会根据segment_ids(即这里的values)对data(即weights)进行分组聚合,返回每个分组的最大/最小值- 未出现的索引会被填充对应数据类型的极值(整数类型为最大/最小值,浮点为±inf),需要手动替换为0,与
tf.math.bincount的默认行为保持一致
内容的提问来源于stack exchange,提问作者Le_Coeur
相关产品推荐
相关产品推荐

