TensorFlow中按索引张量排序唯一元素新张量的代码问题排查
问题分析与修复方案
你的代码逻辑其实是对的,但有两个小细节问题导致你看不到期望的输出:
- 没有捕获计算结果:在TensorFlow 1.x中,
sess.run(d)只是执行了张量的计算,但你没有把结果赋值给变量,直接打印d的话,输出的是张量的定义对象,而不是实际计算出来的数值。 - Python语法错误:Python 3里
print是函数,必须加括号写成print(d),否则会触发语法报错。
先拆解下你的代码逻辑(是正确的):
tf.unique([[1,2], [3,4], [1,2], [3,4], [3,4]])返回的a是原张量的唯一子张量列表:[[1,2], [3,4]]tf.unique([1,0,1,0,0])返回的b是原序列的唯一元素:[1, 0]a[b, :]会根据b的索引从a中取元素,也就是a[1, :](即[3,4])和a[0, :](即[1,2]),正好是你期望的[[3,4], [1,2]]
修正后的代码:
import tensorflow as tf a, _ = tf.unique([[1, 2], [3, 4], [1, 2], [3, 4], [3, 4]]) b, _ = tf.unique([1, 0, 1, 0, 0]) d = a[b, :] with tf.Session() as sess: # 这里不需要初始化全局变量,因为我们没有定义任何可训练变量 output = sess.run(d) print(output)
运行这段代码就会输出你想要的结果:
[[3 4] [1 2]]
额外补充(TensorFlow 2.x 版本写法):
如果使用TF2.x,不需要手动创建Session,直接用 eager 执行模式会更简洁:
import tensorflow as tf a, _ = tf.unique([[1, 2], [3, 4], [1, 2], [3, 4], [3, 4]]) b, _ = tf.unique([1, 0, 1, 0, 0]) d = a[b, :] print(d.numpy())
内容的提问来源于stack exchange,提问作者nairouz mrabah
相关产品推荐
相关产品推荐

