【问题标题】:Why does Tensorflow Reshape tf.reshape() break the flow of gradients?为什么 Tensorflow Reshape tf.reshape() 会破坏梯度的流动?
【发布时间】:2017-12-03 20:19:49
【问题描述】:

我正在创建一个tf.Variable(),然后使用该变量创建一个简单的函数,然后使用tf.reshape() 展平原始变量,然后在函数和展平变量之间使用tf.gradients()。为什么返回[None]

var = tf.Variable(np.ones((5,5)), dtype = tf.float32)
f = tf.reduce_sum(tf.reduce_sum(tf.square(var)))
var_f = tf.reshape(var, [-1])
print tf.gradients(f,var_f)

上述代码块执行时返回[None]。这是一个错误吗?请帮忙!

【问题讨论】:

  • 您必须在session 中运行它,如basic TF tutorials 所示。
  • @jkschin 在这种情况下不是这样。代码没有在计算图中执行任何东西,它只是定义计算图。亲自尝试一下——sn-p 在有和没有会话的情况下工作方式相同。

标签: python tensorflow


【解决方案1】:

您正在寻找f 相对于var_f 的导数,但f 不是var_f 的函数,而是var。这就是为什么你得到[无]。现在,如果您将代码更改为:

 var = tf.Variable(np.ones((5,5)), dtype = tf.float32)
 var_f = tf.reshape(var, [-1])
 f = tf.reduce_sum(tf.reduce_sum(tf.square(var_f)))
 grad = tf.gradients(f,var_f)
 print(grad)

您的渐变将被定义:

tf.Tensor 'gradients_28/Square_32_grad/mul_1:0' shape=(25,) dtype=float32>

以下代码的图表可视化如下:

 var = tf.Variable(np.ones((5,5)), dtype = tf.float32, name='var')
 f = tf.reduce_sum(tf.reduce_sum(tf.square(var)), name='f')
 var_f = tf.reshape(var, [-1], name='var_f')
 grad_1 = tf.gradients(f,var_f, name='grad_1')
 grad_2 = tf.gradients(f,var, name='grad_2')

grad_1 的导数未定义,而grad_2 的导数已定义。显示了两个梯度的反向传播图(梯度图)。

【讨论】:

  • 这个答案既简单又好,但我仍然感到惊讶的是渐变不会自动出现在重塑的变量上。张量的形状存储在包装张量的底层数组和其他数据的对象上。当.reshape 被调用时,底层数组和(一些?全部?)其他数据被重用或重新计算。这就是重塑速度很快的原因。因此,我认为期望函数的张量依赖关系通过重塑操作被重用(或至少重新计算)的数据来跟踪是合理的。但显然不是!我真的很想知道为什么。
  • 这是一个有趣的问题,这是我的理解:reshape() 是对的,数据被重用(未重新计算)但 var_f 仍将是图中的不同节点。所以当你调用tf.gradients() 时,它会构建一个反向传播图,在这种情况下,节点f 找不到节点var_f 的路径。
  • 这是一种简洁的表达方式:“node f 找不到指向节点var_f的路径”。谢谢你。作为一个使用 reshape 接口的程序员,我仍然对找不到路径感到惊讶。我想知道为什么 tensorflow 开发人员决定在计算图中创建一个新节点以进行重塑。最重要的是,我想知道为什么新的var_f 节点没有比f 更有优势。
  • 我说的是反向传播(梯度)图,它没有到相关节点的依赖路径,因为它的构建是为了实现链式规则。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2013-04-09
  • 1970-01-01
  • 1970-01-01
  • 2015-03-25
  • 1970-01-01
相关资源
最近更新 更多