【问题标题】:Tensorflow: gradient_override_map cannot override op tf.stack 's backward gradientTensorflow:gradient_override_map 无法覆盖 op tf.stack 的后向梯度
【发布时间】:2017-10-30 10:56:11
【问题描述】:

我试图用tf.RegisterGradienttf.gradient_override_map编辑tf.stack op的后向梯度计算机制,这是我的代码:

import tensorflow as tf

class SynthGradBuilder(object):
    def __init__(self):
        self.num_calls = 0

    def __call__(self, x, l=1.0):
        op_name = "SynthGrad%d" % self.num_calls
        @tf.RegisterGradient(op_name)
        def _grad_synth(op, grad):
            return grad[0]

        g = tf.get_default_graph()
        with g.gradient_override_map({"stack": op_name}):
            y = tf.stack([x,x])

        self.num_calls += 1
        return y

GradSys = SynthGradBuilder()

在另一个脚本中,我写了

import tensorflow as tf
from gradient_synthesizer import GradSys

x = tf.Variable([1,2])
y = GradSys(x, l=1)
z = tf.stack([x,x])


grad = tf.gradients(y, x, grad_ys=[[tf.convert_to_tensor([3, 4]), 
                              tf.convert_to_tensor([6, 8])]])
grad_stack = tf.gradients(z, x, grad_ys=[[tf.convert_to_tensor([3, 4]), 
                              tf.convert_to_tensor([6, 8])]])

with tf.Session() as sess:
    sess.run(tf.global_variables_initializer())

    print "grad bp: ", sess.run(grad)
    print "grad_stack: ", sess.run(grad_stack)
    print "y: ", sess.run(y)

预期的输出应该是:

grad bp: [3,4];
grad_stack: [3+6, 4+8] = [9, 12];
y: [[1,2], [1,2]];

我实际上从代码中得到的是:

表示tf.stack的后向梯度根本没有被替换,这与我的预期相反。

不知道是不是误用了“stack”作为操作tf.stack的类型字符串导致了这样的差异,我做了如下实验:

描述张量y的第一项,“stack:0”表明optf.stack的注册名称是“stack”,也是它的类型字符串。所以看起来这不是“堆栈”的错。

我不知道我的代码问题的原因。我想知道是否有人可以帮助我。

【问题讨论】:

  • 这真的很奇怪。似乎 def _grad_synth(op, grad) 实际上从未被调用过。很想知道您是否可以找出原因并解决它。我也会从我身边尝试。如果有什么事情会通知你。

标签: python tensorflow


【解决方案1】:

Tl;dr:正确的代码应该是:

@tf.RegisterGradient(op_name)
def _grad_synth(op, grad):
  x, y = tf.unstack(grad)
  return [x, tf.zeros_like(y)]

g = tf.get_default_graph()
with g.gradient_override_map({"Pack": op_name}):
  y = tf.stack([x, x])

因为这是一个很常见的问题,所以我想稍微解释一下:

您的原始代码中有两个主要问题:

  1. gradient_override_map的错误用法:

tf.stack 的实际 OP 名称是 Pack(不是 Stack),因此您需要覆盖 Pack 而不是 Stack

`g.gradient_override_map({"Pack": op_name})`.

您可能想知道我如何知道实际的 OP 名称?好吧,一个简单的方法是通过运行以下代码来探测 GraphDef:

with tf.Graph().as_default():
  x = tf.constant(0)
  y = tf.stack([x, x])
  print(tf.get_default_graph().as_graph_def())
  1. 梯度函数错误:

Pack 的原始渐变是一个简单的Unpack (official code)。在您的情况下,您仍然需要先解包渐变,但只传播第一部分:

@tf.RegisterGradient(op_name)
def _grad_synth(op, grad):
  x, y = tf.unstack(grad)
  return [x, tf.zeros_like(y)]

请注意,此代码非常适合您的情况。但是,如果您想支持任意长度的堆栈,您可以使用稍微复杂一点的版本:

@tf.RegisterGradient(op_name)
def _grad_synth(op, grad):
  x_list = tf.unstack(grad)
  for i in range(1, len(x_list)):
    x_list[i] = tf.zeros_like(x_list[i])
  return x_list

【讨论】:

  • 感谢您的帮助!你的方法很适合我的问题!仍然对打印操作得到的“stack:0”感到好奇,它是什么?当我使用名称“stack”通过tf.Graph.get_operation_by_name 检索操作时,它会正确返回操作。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 2012-11-25
  • 1970-01-01
  • 2019-01-13
  • 2019-10-27
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多