You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

多层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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 17:12:00