如何用TensorFlow的top_k和scatter_nd实现特征图仅保留Top-K值
实现TensorFlow中每个空间位置的Top-K特征保留操作
我来帮你用tf.math.top_k和tf.scatter_nd实现这个需求——仅保留每个空间位置(x,y)对应的channels维度中的Top-K值,其余位置置0,输出和输入形状完全一致。
实现思路
核心思路是先定位每个需要保留的值的完整坐标,再把这些值"放回"原形状的张量中:
- 第一步:将输入张量重塑,把batch和空间维度合并,方便在channels维度上计算Top-K
- 第二步:获取每个空间位置的Top-K值及其在channels维度的索引
- 第三步:构造这些值对应的四维完整坐标(batch, x, y, channel)
- 第四步:用
scatter_nd将Top-K值填充到对应位置,其余位置自动为0
完整代码实现
import tensorflow as tf def top_k_feature_propagation(input_tensor, k=1): # 获取输入张量的静态形状(如果已知的话) batch_size, height, width, channels = input_tensor.get_shape().as_list() # 重塑张量:将[batch, height, width, channels]转为[batch*height*width, channels] # 这样每一行对应一个空间位置的所有channels值 reshaped_input = tf.reshape(input_tensor, [-1, channels]) # 计算每个空间位置的Top-K值和对应的channels索引 top_k_vals, top_k_indices = tf.math.top_k(reshaped_input, k=k) # 生成每个Top-K值对应的行索引(对应batch*height*width中的每个空间位置) row_indices = tf.tile(tf.expand_dims(tf.range(tf.shape(reshaped_input)[0]), 1), [1, k]) # 将行索引拆分为batch、x、y维度的索引 flat_row_indices = tf.reshape(row_indices, [-1]) batch_idx = flat_row_indices // (height * width) spatial_flat_idx = flat_row_indices % (height * width) x_idx = spatial_flat_idx // width y_idx = spatial_flat_idx % width # 展平Top-K的channels索引 channel_idx = tf.reshape(top_k_indices, [-1]) # 构造四维完整索引:[batch, x, y, channel] full_indices = tf.stack([batch_idx, x_idx, y_idx, channel_idx], axis=-1) # 展平Top-K的值,用于scatter操作 flat_top_k_vals = tf.reshape(top_k_vals, [-1]) # 使用scatter_nd构造输出张量,未指定的位置自动填充0 output_tensor = tf.scatter_nd(full_indices, flat_top_k_vals, shape=tf.shape(input_tensor)) return output_tensor # 测试你的示例输入 if __name__ == "__main__": input_np = [[[[6.4, 1.4, 1.3], [2.1, 6.5, 4.8]], [[2.3, 9.2, 2.8], [7.9, 5.1, 0.6]]]] input_tensor = tf.convert_to_tensor(input_np, dtype=tf.float32) # 调用函数,k=1 output_tensor = top_k_feature_propagation(input_tensor, k=1) print("输出结果:") print(output_tensor.numpy())
代码解释
- 张量重塑:把四维张量转成二维,让每个空间位置的channels值变成一行,简化Top-K的计算。
- Top-K计算:
tf.math.top_k直接给出每个空间位置的前k个最大值和它们在channels中的位置。 - 坐标构造:通过行索引拆分出batch、x、y的位置,再和channels索引组合成完整的四维坐标,确保每个值都能对应到原张量的正确位置。
- Scatter填充:
tf.scatter_nd会根据我们提供的坐标,把Top-K值放到对应位置,其余位置默认填充0,完美匹配你的需求。
运行代码后,输出结果和你给出的示例完全一致:
[[[[6.4 0. 0. ] [0. 6.5 0. ]] [[0. 9.2 0. ] [7.9 0. 0. ]]]]
内容的提问来源于stack exchange,提问作者YAO
相关产品推荐
相关产品推荐

