You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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())

代码解释

  1. 张量重塑:把四维张量转成二维,让每个空间位置的channels值变成一行,简化Top-K的计算。
  2. Top-K计算:tf.math.top_k直接给出每个空间位置的前k个最大值和它们在channels中的位置。
  3. 坐标构造:通过行索引拆分出batch、x、y的位置,再和channels索引组合成完整的四维坐标,确保每个值都能对应到原张量的正确位置。
  4. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 07:48:07