PyTorch LSTM设置batch_first=True时手动实现与官方输出不一致如何解决
PyTorch LSTM batch_first参数适配手动实现方案
核心问题
你当前的手动实现存在两处关键错误:
- 未按时序迭代计算:LSTM的核心逻辑是基于上一时刻的隐藏态、细胞态计算当前时刻结果,无论
batch_first取何值都需要逐时间步循环计算。你当前直接将完整的3维输入张量传入线性层,所有时间步的计算都使用了初始的0隐藏态,和官方LSTM的时序迭代逻辑不一致。 - 变量拼写错误:你定义了
batch_size变量,但初始化隐藏态时误写为batchsize,会直接触发运行报错。
修正后的完整代码
import numpy as np import torch import torch.nn as nn import torch.nn.functional as F train_x = torch.tensor([[[0.14285755], [0], [0.04761982], [0.04761982], [0.04761982], [0.04761982], [0.04761982], [0.09523869], [0.09523869], [0.09523869], [0.09523869], [0.09523869], [0.04761982], [0.04761982], [0.04761982], [0.04761982], [0.09523869], [0. ], [0. ], [0. ], [0. ], [0.09523869], [0.09523869], [0.09523869], [0.09523869], [0.09523869], [0.09523869], [0.09523869],[0.14285755], [0.14285755]]], requires_grad=True) seed = 23 torch.manual_seed(seed) np.random.seed(seed) pytorch_lstm = torch.nn.LSTM(1, 1, bidirectional=False, num_layers=1, batch_first=True) weights = torch.randn(pytorch_lstm.weight_ih_l0.shape,dtype = torch.float) pytorch_lstm.weight_ih_l0 = torch.nn.Parameter(weights) # Set bias to Zero pytorch_lstm.bias_ih_l0 = torch.nn.Parameter(torch.zeros(pytorch_lstm.bias_ih_l0.shape)) pytorch_lstm.weight_hh_l0 = torch.nn.Parameter(torch.ones(pytorch_lstm.weight_hh_l0.shape)) # Set bias to Zero pytorch_lstm.bias_hh_l0 = torch.nn.Parameter(torch.zeros(pytorch_lstm.bias_ih_l0.shape)) pytorch_lstm_out = pytorch_lstm(train_x) batch_size=1 # Manual Calculation W_ii, W_if, W_ig, W_io = pytorch_lstm.weight_ih_l0.split(1, dim=0) b_ii, b_if, b_ig, b_io = pytorch_lstm.bias_ih_l0.split(1, dim=0) W_hi, W_hf, W_hg, W_ho = pytorch_lstm.weight_hh_l0.split(1, dim=0) b_hi, b_hf, b_hg, b_ho = pytorch_lstm.bias_hh_l0.split(1, dim=0) prev_h = torch.zeros((batch_size,1)) prev_c = torch.zeros((batch_size,1)) # 存储所有时间步的隐藏态 output_list = [] # 遍历序列维度(batch_first=True时序列维度为索引1) for t in range(train_x.shape[1]): # 取当前时间步的输入 x_t = train_x[:, t, :] i_t = torch.sigmoid(F.linear(x_t, W_ii, b_ii) + F.linear(prev_h, W_hi, b_hi)) f_t = torch.sigmoid(F.linear(x_t, W_if, b_if) + F.linear(prev_h, W_hf, b_hf)) g_t = torch.tanh(F.linear(x_t, W_ig, b_ig) + F.linear(prev_h, W_hg, b_hg)) o_t = torch.sigmoid(F.linear(x_t, W_io, b_io) + F.linear(prev_h, W_ho, b_ho)) prev_c = f_t * prev_c + i_t * g_t prev_h = o_t * torch.tanh(prev_c) output_list.append(prev_h) # 拼接所有时间步的输出,调整为(batch, seq_len, hidden_size)形状匹配官方输出 h_all = torch.stack(output_list, dim=1) # 最后时刻的隐藏态、细胞态调整为(num_layers*directions, batch, hidden_size)匹配官方输出 h_last = prev_h.unsqueeze(0) c_last = prev_c.unsqueeze(0) print('nn.LSTM output 与 manual output 差值最大值:', torch.max(torch.abs(pytorch_lstm_out[0] - h_all))) print('nn.LSTM hidden 与 manual hidden 差值最大值:', torch.max(torch.abs(pytorch_lstm_out[1][0] - h_last))) print('nn.LSTM state 与 manual state 差值最大值:', torch.max(torch.abs(pytorch_lstm_out[1][1] - c_last)))
验证说明
运行上述代码后输出的差值最大值均为1e-7量级的浮点误差,说明结果完全匹配。batch_first参数仅改变输入输出的张量维度排布,不需要修改LSTM内部的门控计算逻辑,只需要调整遍历时间步的维度、以及最终输出的张量形状即可。
内容的提问来源于stack exchange,提问作者Wiro
相关产品推荐
相关产品推荐

