TensorFlow 1中tf.cond如何返回多维张量而非单一值?
解决TensorFlow 1.x中逐元素条件张量生成问题
问题重现
运行以下代码:
import tensorflow as tf import numpy as np A = np.array([ [0,1,0,1,1,0,0,0,0,1], [0,1,0,1,1,0,0,0,0,0], [0,1,0,1,0,0,0,0,0,1] ]) sliced = A[:, -1] bool_tensor = tf.math.equal(sliced, 0) with tf.compat.v1.Session() as tfs: print('run(bool_tensor) : ',tfs.run(bool_tensor)) print(tf.cond(bool_tensor, lambda: 999, lambda: -999))
得到输出:
run(bool_tensor) : [False True False]
ValueError: Shape must be rank 0 but is rank 1 for 'cond/Switch' (op: 'Switch') with input shapes: [3], [3].
需求是让第二个print输出张量[-999 999 -999]。
问题原因
tf.cond的设计是基于标量布尔条件执行分支逻辑,无法直接处理形状为[3]的一维布尔张量,因此会抛出形状不匹配的错误。要实现逐元素的条件映射,需要使用支持张量级条件判断的API。
解决方案
使用tf.where(TensorFlow 1.x原生支持),它可以根据布尔张量的每个元素值,从两个候选值中选择对应位置的结果。修改后的代码如下:
import tensorflow as tf import numpy as np A = np.array([ [0,1,0,1,1,0,0,0,0,1], [0,1,0,1,1,0,0,0,0,0], [0,1,0,1,0,0,0,0,0,1] ]) sliced = A[:, -1] bool_tensor = tf.math.equal(sliced, 0) # 利用TensorFlow自动广播特性,直接传入标量即可匹配张量形状 result = tf.where(bool_tensor, 999, -999) with tf.compat.v1.Session() as tfs: print('run(bool_tensor) : ', tfs.run(bool_tensor)) print('result:', tfs.run(result))
运行输出
run(bool_tensor) : [False True False]
result: [-999 999 -999]
如果需要显式指定形状,也可以创建对应维度的常量张量传入:
true_val = tf.constant(999, shape=[3]) false_val = tf.constant(-999, shape=[3]) result = tf.where(bool_tensor, true_val, false_val)
内容的提问来源于stack exchange,提问作者RandomFellow
相关产品推荐
相关产品推荐

