TensorFlow中按指定二维值筛选对应张量的高效实现方法
更高效的TensorFlow张量筛选方案
针对你的需求,这里提供两种更简洁高效的实现思路,均优于你当前的方法:
方法1:利用类别索引快速筛选
你的y_true是one-hot编码格式,[1,0]对应类别0,[0,1]对应类别1。可以先将y_true转换为类别索引,再通过布尔索引直接筛选y_pred:
import tensorflow as tf y_true = tf.constant([[1,0], [0,1], [1,0], [1,0], [0,1], [0,1], [1,0], [0,1], [1,0], [0,1]]) y_pred = tf.constant([[0.6,0.4], [0.3,0.7], [0.8,0.2], [0.8,0.2], [0.3,0.7],[0.1,0.9],[0.9, 0.1],[0.4,0.6],[0.6,0.4],[0.2,0.8]]) # 筛选对应[1,0]的y_pred(类别0) class_true = tf.argmax(y_true, axis=1) mask_0 = tf.equal(class_true, 0) zeros = y_pred[mask_0] # 筛选对应[0,1]的y_pred(类别1) mask_1 = tf.equal(class_true, 1) ones = y_pred[mask_1]
优势:tf.argmax是TensorFlow高度优化的内置操作,布尔索引y_pred[mask]比tf.gather_nd+tf.where的组合更简洁,底层实现也减少了中间张量的创建开销。
方法2:简化Mask生成逻辑
如果不想转换类别索引,可直接用tf.reduce_all一次性判断整行是否等于目标one-hot向量,替代手动逐元素判断+逻辑与的步骤:
# 筛选对应[1,0]的y_pred mask_0 = tf.reduce_all(tf.equal(y_true, [1, 0]), axis=1) zeros = y_pred[mask_0] # 筛选对应[0,1]的y_pred mask_1 = tf.reduce_all(tf.equal(y_true, [0, 1]), axis=1) ones = y_pred[mask_1]
优势:将原方法中ind_zero和index_zero两个中间张量的运算合并为一步,代码更简洁,同时减少了不必要的计算步骤,效率更高。
内容的提问来源于stack exchange,提问作者Armin Amini
相关产品推荐
相关产品推荐

