You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 18:37:12