多层LSTM中LSTMCell的weight_ih梯度与PyTorch不符问题排查
多层LSTM梯度匹配问题
我用numpy开发AI工具,单层LSTMCell的输出与梯度已通过PyTorch验证完全一致。但搭建两层及以上LSTM时,仅将前一层的h传入下一层,模型输出和输出梯度仍正确,然而weight_ih的梯度(即使是最后一层)与PyTorch结果不再匹配。推测多层场景下每个时间步的梯度存在更多依赖关系,但无法确定具体问题。
LSTMCell实现
class LSTMCell(Layer): def __init__(self, input_size, hidden_size): super().__init__() self.input_size = input_size self.hidden_size = hidden_size self.layer_type = 'r' weight_ih = self.xavier_init((4*hidden_size, input_size)) weight_hh = self.xavier_init((4*hidden_size,hidden_size)) bias_ih = np.zeros((4*hidden_size)) bias_hh = np.zeros((4*hidden_size)) self.weights = [weight_ih, weight_hh, bias_ih, bias_hh] self.numweights = len(self.weights) self.input_gates = [] self.forget_gates = [] self.cell_gates = [] self.output_gates = [] self.c_t = [] self.h_t = [] self.inputs = [] self.timesteps = 0 def __call__(self, x, hc): h, c = hc if self.timesteps == 0: self.h_t.append(h) self.c_t.append(c) new_h, new_c = self.forward(x, h, c) #save cell and hidden states of each timestep self.c_t.append(new_c) self.h_t.append(new_h) self.inputs.append(x) self.timesteps += 1 return (new_h, new_c) def forward(self, x, h, c): batch_size = x.shape[0] if len(x.shape) == 2 else 1 weight_ih, weight_hh, bias_ih, bias_hh = self.weights # Linear transformations gates = x @ weight_ih.T + bias_ih + h @ weight_hh.T + bias_hh # Split into 4 separate tensors i, f, g, o = np.split(gates, 4, axis=1) # Apply gate activation functions i = self.sigmoid(i) f = self.sigmoid(f) g = np.tanh(g) o = self.sigmoid(o) #store gates for quicker backward method self.input_gates.append(i) self.forget_gates.append(f) self.cell_gates.append(g) self.output_gates.append(o) # Calculate new cell state and hidden state new_c = f * c + i * g new_h = o * np.tanh(new_c) return (new_h, new_c) def calculate_gradients(self, dh): #initialize gradients to zero so we can sum over timesteps self.gradients = [np.zeros_like(weight) for weight in self.weights] dc = dh * self.output_gates[-1] * (1 - np.square(np.tanh(self.c_t[-1]))) while self.timesteps > 0: dx, dh, dc = self.backward_timestep(dh, dc) return dx def backward_timestep(self, dh, dc): weight_ih, weight_hh, bias_ih, bias_hh = self.weights grad_weight_ih, grad_weight_hh, grad_bias_ih, grad_bias_hh = self.gradients #collection of stored variables we need for this timestep c = self.c_t[-1] c_tanh = np.tanh(c) cprev = self.c_t[-2] hprev = self.h_t[-2] x = self.inputs[-1] #Gates i = self.input_gates[-1] f = self.forget_gates[-1] g = self.cell_gates[-1] o = self.output_gates[-1] #Derivative of gates di = i * (1 - i) df = f * (1 - f) dg = 1 - np.square(g) do = o * (1 - o) #gradient w.r.t. gates grad_i = dc * g * di grad_f = dc * cprev * df grad_g = dc * i * dg grad_o = dh * c_tanh * do #gradient w.r.t. weights dgates = np.concatenate([grad_i,grad_f,grad_g,grad_o], axis = 1) grad_weight_ih += dgates.T @ x grad_weight_hh += dgates.T @ hprev grad_bias = np.sum(dgates, axis = 0) grad_bias_ih += grad_bias grad_bias_hh += grad_bias #gradient w.r.t. previous hidden state and cell state dhprev = dgates @ weight_hh dcprev = None if self.timesteps > 1: #gradient w.r.t. previous cell state oprev = self.output_gates[-2] dcprev = f * dc + dhprev * oprev * (1 - np.square(np.tanh(cprev))) #gradient w.r.t. input dx = dgates @ weight_ih #remove the last timestep from the saved states self.c_t.pop() self.h_t.pop() self.inputs.pop() self.input_gates.pop() self.forget_gates.pop() self.cell_gates.pop() self.output_gates.pop() self.timesteps -= 1 self.gradients = [grad_weight_ih, grad_weight_hh, grad_bias_ih, grad_bias_hh] return dx, dhprev, dcprev
多层LSTM实现及PyTorch对比代码
多层LSTM类实现
class LSTM(Model): def __init__(self, input_size, hidden_size, num_layers, batch_first = False): super().__init__() self.input_size = input_size self.hidden_size = hidden_size self.num_layers = num_layers self.batch_first = batch_first self.layer_type = 'r' self.layers = [LSTMCell(input_size, hidden_size)] for i in range(1,num_layers): self.layers.append(LSTMCell(hidden_size, hidden_size)) def __call__(self, x, hc): h, c = hc has_batch = len(x.shape)==3 if self.batch_first and has_batch: for t in range(x.shape[1]): for i,layer in enumerate(self.layers): if i == 0: h[:,i,:], c[:,i,:] = layer(x[:,t,:], (h[:,i,:],c[:,i,:])) else: h[:,i,:], c[:,i,:] = layer(h[:,i-1,:], (h[:,i,:],c[:,i,:])) self.output.append(h[:,-1,:].copy()) else: for t in range(x.shape[0]): for i,layer in enumerate(self.layers): if i == 0: h[i], c[i] = layer(x[t], (h[i],c[i])) else: h[i], c[i] = layer(h[i-1], (h[i], c[i])) self.output.append(h[-1].copy()) self.output = np.stack(self.output,axis = 0) return self.output, (h,c) def backward(self, dh): for layer in reversed(self.layers): dh = layer.backward(dh)
对比测试代码
batch_size = 2 input_size = 3 hidden_size = 5 time_steps = 2 num_layers = 1 #Pytorch Model torchrnn = nn.LSTM(input_size, hidden_size, num_layers) #My Model myrnn = LSTM(input_size, hidden_size, num_layers) myrnn.torch_weighted(torchrnn) #Inputs torchin = torch.randn(time_steps, batch_size, input_size) htorch = torch.randn(num_layers, batch_size, hidden_size) ctorch = torch.randn(num_layers, batch_size, hidden_size) npin = torchin.detach().numpy() hnp = htorch.detach().numpy() cnp = ctorch.detach().numpy() #targets torchtrue = torch.randint(hidden_size,(batch_size,)) nptrue = torchtrue.detach().numpy() #Outputs torchout, (htorch, ctorch) = torchrnn(torchin,(htorch, ctorch)) torchout = torchout[-1] torchout.retain_grad() npout, (hnp, cnp) = myrnn(npin, (hnp, cnp)) npout = npout[-1] printdif("Outputs",torchout,npout) #Loss torch_loss_fn = nn.CrossEntropyLoss() torchloss = torch_loss_fn(torchout,torchtrue) np_loss_fn = CrossEntropyLoss() nploss = np_loss_fn(npout, nptrue) printdif("Loss", torchloss, nploss) #Calculate gradients torchloss.backward(retain_graph=True) myrnn.backward(np_loss_fn.gradient) #Gradients printdif("Output gradient", torchout.grad, np_loss_fn.gradient) for l in reversed(range(num_layers)): printdif(f"layer {l} weight_ih gradient", torchrnn.all_weights[l][0].grad, myrnn.layers[l].gradients[0]) printdif(f"layer {l} weight_hh gradient", torchrnn.all_weights[l][1].grad, myrnn.layers[l].gradients[1]) printdif(f"layer {l} bias_ih gradient", torchrnn.all_weights[l][2].grad, myrnn.layers[l].gradients[2]) printdif(f"layer {l} bias_hh gradient", torchrnn.all_weights[l][3].grad, myrnn.layers[l].gradients[3])
内容的提问来源于stack exchange,提问作者Liam
相关产品推荐
相关产品推荐

