tf.unique仅支持1D张量,如何获取2D张量的唯一值索引?
在2D张量中获取行内唯一值的索引
因为tf.unique仅支持1D张量,要实现2D张量每行内的唯一值索引映射,可以通过逐行处理的方式实现——对每一行单独调用tf.unique,再将结果拼接回2D张量。
实现代码
import tensorflow as tf # 定义输入张量 input_tensor = tf.constant([[0,0,1,1,1,8], [2,2,5,5,5,2], [6,6,6,8,8,9]]) # 定义逐行处理的函数:对单行执行tf.unique并返回索引 def get_row_unique_indices(row): _, indices = tf.unique(row) return indices # 用tf.map_fn遍历每一行处理,再堆叠成2D结果 result = tf.map_fn(get_row_unique_indices, input_tensor, dtype=tf.int32) # 验证结果 print(result.numpy())
输出结果
运行后会得到:
[[0 0 1 1 1 2] [0 0 1 1 1 0] [0 0 0 1 1 2]]
完全符合期望的输出格式。
原理说明
tf.map_fn会遍历输入2D张量的每一行,将每行传入处理函数;- 对单行调用
tf.unique,返回的indices就是该行元素对应的唯一值索引(首次出现的元素对应0,第二次出现的新元素对应1,以此类推); - 最后
tf.map_fn会自动将所有行的索引结果堆叠成2D张量,得到最终输出。
内容的提问来源于stack exchange,提问作者Ocxs
相关产品推荐
相关产品推荐

