基于Numpy手动实现RNN时梯度检查失败问题排查
Gradient Check Failures in RNN BPTT - Fixed
我最近从零开始用Numpy构建RNN,就是为了彻底搞懂RNN的内部工作机制,结果在实现时间反向传播(BPTT)的时候碰了个钉子——dLdU、dLdW和dLdB_state这几个参数的梯度检查始终通不过。我反复核对数学推导,确认输入输出张量的形状(X是(batch_size, seq_length, input_dim),Y是(batch_size, seq_length, output_dim)),甚至检查了前向传播缓存里的Z_states、States等张量的维度,都没发现问题,折腾了好久才找到根源。
问题定位
最后发现是BPTT内层循环里的一个细节错误:在更新dLdZ_state的时候,我错误地用了已激活的隐藏状态States来计算激活函数的导数,而实际上应该用激活前的隐藏层输入Z_states。
错误代码片段
# 错误写法:误用了激活后的States dLdZ_state = np.multiply(np.dot(dLdZ_state, self.W.T), self.hidden_activation_function_prime(States[:,t_prev-1,:]))
修复后的代码片段
# 正确写法:改用激活前的Z_states dLdZ_state = np.multiply(np.dot(dLdZ_state, self.W.T), self.hidden_activation_function_prime(Z_states[:,t_prev-1,:]))
为什么这个错误会导致梯度检查失败?
回忆RNN的前向传播逻辑:
$Z_{state_t} = X_t \cdot U + State_{t-1} \cdot W + B_{state}$
$State_t = f(Z_{state_t})$ (f是隐藏层激活函数)
反向传播计算梯度时,我们需要的是激活函数$f$对**激活前的输入$Z_{state}$**的导数,也就是$f'(Z_{state_t})$。如果用States来计算,相当于求$f'(f(Z_{state_t}))$,完全偏离了正确的梯度传导路径,自然会让梯度检查结果和数值梯度对不上。
把这个细节修正后,梯度检查就顺利通过了,后续的参数更新也能正常工作了。
内容的提问来源于stack exchange,提问作者pickandroll3
相关产品推荐
相关产品推荐

