TensorFlow自定义梯度op的grad参数是否始终为全1矩阵?
关于TensorFlow自定义梯度中
grad参数的解惑 嘿,这个问题问得挺深入的!其实答案是否定的——自定义梯度操作里的grad参数并不是始终是全1矩阵,它的取值完全取决于你后续的计算需求,也就是常说的「上游梯度」。
你之所以会觉得它总是全1,大概率是因为测试时都是直接对自定义操作的输出张量求和,再求这个总和对输入的梯度。这种情况下,求和操作对输出张量的每个元素的梯度都是1,所以传递给自定义梯度函数的grad就是全1矩阵/张量。但只要改变后续的计算逻辑,grad就会跟着变化。
结合你给出的两种自定义方式分析
1. @tf.RegisterGradient 方式
你写的示例代码是:
@tf.RegisterGradient("CustomGrad") def _custom_grad(op, grad): return grad
这里的grad就是上游传来的梯度。比如换个测试场景:如果自定义操作的输出是y,我们计算z = tf.reduce_sum(y * 3),此时求z对输入的梯度,z对y的梯度就是全3的张量,那么_custom_grad里的grad就会是全3,而非全1。
2. @function.Defun 方式
你的示例代码是:
@function.Defun(tf.float32, tf.float32) def bprop(op, grad): return grad @function.Defun(tf.float32, grad_func=bprop) def fprop(W): W = tf.sign(W) return W
同样的,如果你直接计算tf.gradients(fprop(W), W),默认是计算sum(fprop(W))对W的梯度,此时上游梯度是全1。但如果我们计算out = fprop(W) * 5,再求sum(out)对W的梯度,那么bprop里的grad就会变成全5的张量。
直观代码示例验证
import tensorflow as tf from tensorflow.python.framework import function # 测试RegisterGradient @tf.RegisterGradient("ScaledSignGrad") def _scaled_sign_grad(op, grad): # 打印上游梯度的值 print("上游梯度值:", grad.numpy()) return grad * 0.5 def custom_scaled_sign(x): g = tf.get_default_graph() with g.gradient_override_map({"Identity": "ScaledSignGrad"}): return tf.identity(x) # 场景1:直接对输出求和 x1 = tf.Variable([1.0, -2.0, 3.0]) with tf.GradientTape() as tape: y1 = custom_scaled_sign(x1) sum_y1 = tf.reduce_sum(y1) grad1 = tape.gradient(sum_y1, x1) print("场景1的最终梯度:", grad1.numpy()) # 输出[0.5, 0.5, 0.5],对应上游梯度全1 # 场景2:输出乘2后求和 x2 = tf.Variable([1.0, -2.0, 3.0]) with tf.GradientTape() as tape: y2 = custom_scaled_sign(x2) sum_y2 = tf.reduce_sum(y2 * 2) grad2 = tape.gradient(sum_y2, x2) print("场景2的最终梯度:", grad2.numpy()) # 输出[1.0, 1.0, 1.0],对应上游梯度全2
运行这段代码就能看到,两种场景下的上游梯度grad分别是全1和全2,完全不是固定值。
总结
自定义梯度函数里的grad参数是动态变化的,它代表后续计算对当前操作输出的梯度需求:
- 当你直接对输出张量求和时,
grad是全1; - 当输出被缩放、加权或者参与其他复杂计算时,
grad会对应变成相应的数值; - 如果最终输出是标量损失,初始的上游梯度是1,但如果输出是多维张量,上游梯度的形状会和输出一致,值取决于后续操作。
内容的提问来源于stack exchange,提问作者Khoi Tran
相关产品推荐
相关产品推荐

