TensorFlow中如何基于条件筛选张量的指定行
纯TensorFlow环境下按条件筛选张量行的实现
针对筛选「张量中第一个元素取值为0的行」的需求,直接调用TensorFlow内置的tf.boolean_mask接口即可实现,全程不需要引入其他第三方依赖。
完整实现代码
import tensorflow as tf # 定义示例张量x x = tf.constant([ [0, 1, 2], [1, 1, 2], [0, 1, 4] ], dtype=tf.int32) # 第一步:生成行级布尔掩码:判断每行第一个元素(索引为0)是否等于0 row_mask = x[:, 0] == 0 # 此时row_mask的值为 [True, False, True],对应第1、3行满足筛选条件 # 第二步:通过掩码过滤原张量,得到全新的结果张量 filtered_x = tf.boolean_mask(x, row_mask)
运行结果
执行代码后得到的filtered_x张量如下,完全符合筛选要求:
<tf.Tensor: shape=(2, 3), dtype=int32, numpy= array([[0, 1, 2], [0, 1, 4]], dtype=int32)>
补充说明
- 该方法返回的是独立的新张量,不会修改原始张量
x的内容 - 如果需要调整筛选规则,只需要修改
row_mask的生成逻辑即可,比如要筛选第三列值大于3的行,直接写row_mask = x[:, 2] > 3就能适配 - 所有逻辑均使用TensorFlow原生API实现,可以直接在图模式、eager模式下正常运行
内容的提问来源于stack exchange,提问作者Dav Bhaji
相关产品推荐
相关产品推荐

