如何在TensorFlow的tf.while_loop()中使用向量条件?
如何让tf.while_loop的向量条件全部为真时才执行循环?
你遇到的问题核心在于:tf.while_loop要求条件函数返回单个标量布尔值,但tf.greater返回的是布尔向量,而Python原生的all()无法在TensorFlow图模式下迭代张量(因为此时张量还只是计算图节点,没有实际数值),所以才会抛出那个TypeError。
解决这个问题的关键是用TensorFlow内置的图兼容函数来判断布尔向量的所有元素是否为真,推荐用tf.reduce_all()——它会对布尔张量的所有元素执行逻辑与操作,最终返回一个标量布尔值,完美符合tf.while_loop的要求。
修改后的代码如下:
import tensorflow as tf i = tf.constant(0) # 使用tf.reduce_all将向量条件转换为单个布尔标量 c = lambda i: tf.reduce_all(tf.greater([10, 10], [i, i])) b = lambda i: tf.add(i, 1) r = tf.while_loop(c, b, [i])
简单解释下:tf.reduce_all(tf.greater([10,10], [i,i]))会先比较两个向量的对应元素,得到一个布尔向量,再把这个向量里的所有值做逻辑与,只有当所有元素都是True时,才返回True,循环继续;只要有一个元素是False,就返回False,循环终止。
另外补充个小知识点:如果你的需求是只要向量中有任意一个元素为真就继续循环,可以把tf.reduce_all()换成tf.reduce_any(),它会执行逻辑或操作,满足任意为真的判断。
内容的提问来源于stack exchange,提问作者Mencia
相关产品推荐
相关产品推荐

