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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 17:39:03