TensorFlow实现直通估计器时最后层加tf.stop_gradient无梯度报错如何解决
错误原因
你将tf.stop_gradient()放在模型最后一层时,整个模型输出的梯度被完全截断,反向传播过程中所有前面的可训练参数都无法获取到梯度,因此抛出无梯度的报错。而当你把它放在中间层时,该层之后还有可训练层,梯度可以从损失端传递到该层之前的参数,所以可以正常运行。
直通估计器的正确实现
直通估计器的核心逻辑是:前向传播使用你需要的截断/离散化后的值,反向传播保留原始输出的梯度,不能直接对最终输出做梯度截断,正确写法遵循以下通用公式:前向输出 = 原始输出 + tf.stop_gradient(你需要的前向计算结果 - 原始输出)
修改后代码示例
假设你需要对最后一层Dense输出的sigmoid值做二值化的直通逻辑,代码修改如下:
def generator(): model = Sequential([ Embedding(21, 100, input_shape = (50, )), Conv1D(100, 3,2, padding = "same", activation = "relu"), Conv1D(100, 3,2, padding = "same", activation = "relu"), Conv1D(21, 3,1, padding = "same", activation = "relu"), Flatten(), Dense(1,"sigmoid", trainable = False), # 替换原来的纯梯度截断层为直通估计器实现 Lambda(lambda x: x + tf.stop_gradient(tf.cast(x > 0.5, tf.float32) - x)), ]) return model gen = generator() def dis_inputs(): ind = np.random.randint(0, len(x_train), batch_size * 2) real = x_train[ind] out = np.ones((batch_size * 2, )) out[: batch_size ] = 0 return real, out batch_size = 128 x, y = dis_inputs() gen.compile("adam", "binary_crossentropy") gen.fit(x, y, epochs = 100)
逻辑说明
- 前向传播时,
x和-x抵消,最终输出的是你需要的二值化结果,和你的预期前向逻辑一致 - 反向传播时,
tf.stop_gradient包裹的部分不会传递梯度,整个表达式的梯度完全等于原始输出x的梯度,可以正常回传给前面的Embedding、Conv1D层,不会出现无梯度报错 - 如果你需要的前向逻辑不是二值化,只需要把
tf.cast(x > 0.5, tf.float32)替换成你自己的前向计算逻辑即可,通用公式保持不变
内容的提问来源于stack exchange,提问作者100rabh
相关产品推荐
相关产品推荐

