【问题标题】:trying to use if in tensorflow's map_fn尝试在张量流的 map_fn 中使用 if
【发布时间】:2018-04-04 10:54:08
【问题描述】:

我正在尝试使用 tensorflow 中的 map_fn 对列向量进行转换,但它不起作用。

对于以下列向量:

elems = np.array([[1.0], [2.0], [3.0]])

当我这样做时:

tf_m = tf.map_fn(lambda x: x + 1.0, elems)
with tf.Session() as sess:
    res = sess.run(tf_m)
    print(str(res))

我得到了我期望的结果,即这个列向量:

[[2.]
 [3.]
 [4.]]

但是,当我这样做时:

tf_m2 = tf.map_fn(lambda x: x+1 if x % 2 > 0 else x, elems)
with tf.Session() as sess:
    res = sess.run(tf_m2)
    print(str(res))

代码失败,出现以下异常:

TypeError:不允许将tf.Tensor 用作Python bool。使用if t is not None: 而不是if t: 来测试是否定义了张量,并使用 TensorFlow 操作(例如 tf.cond)来执行以张量值为条件的子图。

我试过打印 x 的类型,它是一个形状为 (1,) 的张量。因此,看起来正在发生的是,这些值不是作为标量值传递给 lambda,而是作为具有形状 (1,) 的张量; % 被广播,产生另一个形状为 (1,) 的张量,但该张量不能应用 >= 运算符。

有没有办法让它工作?有没有办法获得一个我可以应用 >= 运算符的 actual 标量?如果没有,我可以使用 map_fn 的有效替代方法吗?

(我查看了 tf.cond,在这种情况下如何使用它并不明显。据我了解,tf.cond 产生一个操作,而不是可调用的,所以我将如何使用它来自 map_fn 应用的 lambda?)

【问题讨论】:

  • 我刚刚意识到您的条件是 x % 2 >= 0... 基本上,对于 x 的任何值,这不是真的吗?
  • 糟糕!那是我的一个愚蠢的错误,对不起。我将编辑问题。

标签: python tensorflow


【解决方案1】:

您可以像这样使用tf.map_fntf.cond 做到这一点:

elems_shape = tf.shape(elems)
elems_flat = tf.reshape(elems, [-1])
tf_m2_flat = tf.map_fn(lambda x: tf.cond(x % 2 > 0, lambda: x + 1, lambda: x), elems_flat)
tf_m2 = tf.reshape(tf_m2_flat, elems_shape)

但您也可以像这样简单地使用tf.where

tf_m2 = tf.where(elems % 2 > 0, elems + 1, elems)

【讨论】:

  • 这行得通,非常感谢!至少对我来说,第二个和第三个选项需要是可调用的,所以我必须这样做:tf.map_fn(lambda x: tf.cond(x % 2 > 0, lambda: x+1, lambda: x), tf_elems)。这不适用于列向量,它抱怨形状是 (1,) 而不是 ()。但是,如果我将其展平,应用 map_fn,然后将其重新整形,它就可以工作。我需要担心这样做的性能开销吗?
  • @DavidBolding 对不起,答案不是很好,我忘记了lambda: 的事情,我没有考虑到elems 是二维的,感谢您指出错误。我自己没有对其进行基准测试,但原则上扁平化和重塑应该几乎没有计算成本。我还使用tf.where 添加了一个更简单的解决方案。
  • 非常感谢 tf.where 示例,我认为这将在一个紧凑的表达式中处理我需要的所有内容。
  • 用户究竟应该如何意识到这一点?如果你接受一个 lambda,那么……你接受一个 lambda!!对张量流来说太糟糕了。
猜你喜欢
  • 2016-06-06
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2018-06-09
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多