PyTorch自定义LSTM Cell实现细节疑问咨询
先贴出你提到的LSTM实现代码:
import math import torch as th import torch.nn as nn class LSTM(nn.Module): def __init__(self, input_size, hidden_size, bias=True): super(LSTM, self).__init__() self.input_size = input_size self.hidden_size = hidden_size self.bias = bias self.i2h = nn.Linear(input_size, 4 * hidden_size, bias=bias) self.h2h = nn.Linear(hidden_size, 4 * hidden_size, bias=bias) self.reset_parameters() def reset_parameters(self): std = 1.0 / math.sqrt(self.hidden_size) for w in self.parameters(): w.data.uniform_(-std, std) def forward(self, x, hidden): h, c = hidden h = h.view(h.size(1), -1) c = c.view(c.size(1), -1) x = x.view(x.size(1), -1) # Linear mappings preact = self.i2h(x) + self.h2h(h) # activations gates = preact[:, :3 * self.hidden_size].sigmoid() g_t = preact[:, 3 * self.hidden_size:].tanh() i_t = gates[:, :self.hidden_size] f_t = gates[:, self.hidden_size:2 * self.hidden_size] o_t = gates[:, -self.hidden_size:] c_t = th.mul(c, f_t) + th.mul(i_t, g_t) h_t = th.mul(o_t, c_t.tanh()) h_t = h_t.view(1, h_t.size(0), -1) c_t = c_t.view(1, c_t.size(0), -1) return h_t, (h_t, c_t)
问题1:为何在__init__方法中,self.i2h和self.h2h的输出维度设为4*hidden_size?
这是标准LSTM的设计逻辑——LSTM内部同时运行四个并行的线性变换:输入门(i_t)、遗忘门(f_t)、输出门(o_t),以及候选记忆单元(g_t)。每个变换的输出维度都是hidden_size,四个加起来正好是4*hidden_size。
把它们合并成一个大的线性层计算,比单独写四个nn.Linear要高效得多:既减少了代码冗余,也能让PyTorch在底层做更优的计算调度,避免多次重复的线性运算开销。
问题2:如何理解reset_parameters方法的参数重置逻辑,为何采用该初始化方式?
这个初始化是为了让模型训练初期处于稳定的收敛起点:
std = 1.0 / math.sqrt(self.hidden_size)是Xavier初始化的变体(这里用均匀分布替代正态分布),核心目的是让每层输入和输出的方差尽可能一致,保证信号在网络中传递时不会被过度放大或缩小,避免梯度爆炸/消失。- 用
uniform_(-std, std)让参数在正负对称区间随机初始化,能避免模型一开始就偏向某一种输出方向,帮助后续梯度下降平稳收敛。 - 如果用PyTorch默认的Linear层初始化,参数范围可能过大,导致初始激活值异常,让模型刚训练就陷入梯度问题。
问题3:forward方法中对h、c、x调用view的作用是什么?
这是为了统一张量维度格式,适配线性层的输入要求:
通常LSTM的输入和隐藏状态会带有批次维度,比如默认格式是(num_layers, batch_size, hidden_size),但nn.Linear只接受(batch_size, feature_size)这类二维张量(或最后一维为特征数的高维张量)。
这里的view(h.size(1), -1)就是把隐藏状态h从(1, batch_size, hidden_size)转成(batch_size, hidden_size),同理x也被转成(batch_size, input_size)——这样才能正确传入i2h和h2h计算。后续再把结果转回(1, batch_size, hidden_size),是为了和PyTorch官方LSTM的输出格式对齐,方便和其他模块兼容。
问题4:forward方法激活计算部分,为何将gates的列范围设为[:, :3*self.hidden_size]?
这和LSTM的激活规则对应:
前3个门控单元(输入门、遗忘门、输出门)需要用sigmoid激活(输出限制在0-1之间,用来控制信息的流通开关);而第四个候选记忆单元g_t需要用tanh激活(输出在-1-1之间,用来生成新的候选记忆)。
所以preact的前3*hidden_size列对应三个门控的线性输出,一起做sigmoid;剩下的1*hidden_size列是候选记忆的线性输出,单独做tanh。这个拆分完全符合LSTM的数学定义,也和前面合并计算4*hidden_size的逻辑呼应。
问题5:LSTM的Us和Ws等参数对应代码中的哪些部分?
在标准LSTM的数学公式里,W代表输入到隐藏层的权重矩阵(W_i/W_f/W_o/W_g分别对应四个门/候选单元的输入权重),U代表隐藏层到隐藏层的权重矩阵(U_i/U_f/U_o/U_g对应四个门/候选单元的隐藏状态权重)。对应到这份代码:
self.i2h.weight是合并后的W矩阵,形状为(4*hidden_size, input_size),前hidden_size行是W_i,接下来hidden_size行是W_f,再接下来是W_o,最后hidden_size行是W_g。self.h2h.weight是合并后的U矩阵,形状为(4*hidden_size, hidden_size),同样按顺序对应U_i、U_f、U_o、U_g。- 如果开启了bias(
self.bias=True),self.i2h.bias和self.h2h.bias就是对应的输入偏置和隐藏偏置,也是按四个门/候选单元的顺序合并在一起的。
内容的提问来源于stack exchange,提问作者An Ignorant Wanderer

