tensorflow:将tensor中满足某一条件的数值取出组成新的tensor

本文介绍了如何在TensorFlow中利用tf.where获取满足条件的索引,然后使用tf.gather或tf.gather_nd提取相应值。通过示例代码展示了当条件为数值大于0.5时,如何构建新Tensor。区别在于,tf.gather适用于一维索引,而tf.gather_nd能处理多维索引。
摘要由CSDN通过智能技术生成

首先使用tf.where()将满足条件的数值索引取出来,在numpy中,可以直接用矩阵引用索引将满足条件的数值取出来,但是在tensorflow中这样是不行的。所幸,tensorflow提供了tf.gather()和tf.gather_nd()函数。
看下面这一段代码:

import tensorflow as tf
sess = tf.Session()
def get_tensor():
    x = tf.random_uniform((5, 4))
    ind = tf.where(x>0.5)
    y = tf.gather_nd(x, ind)
    return x, ind, y

在上述代码中,输出分别是原始的tensor x,x中满足特定条件(此处为>0.5)的数值的索引,以及x中满足特定条件的数值。执行以下步骤,观察三个tensor对应的数值:

x, ind, y = get_tensor()
x_, ind_, y_ = sess.run([x, ind, y])

可以得到如下结果:
tensor x的值
tensor ind的值
tensor y的值
可以看

评论 2
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值