【发布时间】:2017-10-30 10:56:11
【问题描述】:
我试图用tf.RegisterGradient和tf.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