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

PyTorch自定义LSTM Cell实现细节疑问咨询

自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 19:12:37