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

能否用PyTorch的LSTMCell实现多层LSTM?正确连接方式是什么?

当然可以用LSTMCell实现多层LSTM!

其实核心逻辑很简单:多层LSTM就是把多个LSTMCell按层串联,上一层的输出(隐藏状态h)作为下一层的输入。因为LSTMCell本身是单时间步、单层的单元,所以我们需要手动管理层与层之间的输入传递,以及每层的隐藏状态(h)和细胞状态(c)。

下面我会一步步讲清楚实现思路,再给出可直接运行的代码,甚至可以和PyTorch官方的nn.LSTM对齐效果。

核心原理

多层LSTM的运作逻辑是:

  1. 第一层的输入是原始序列的每个时间步数据;
  2. 每一层的LSTMCell处理完当前时间步的输入后,输出的隐藏状态h会作为下一层LSTMCell的输入;
  3. 每层都有自己独立的隐藏状态h和细胞状态c,这些状态只在当前层的时间步之间传递,不会跨层共享;
  4. 最终的输出是最后一层所有时间步的h,而返回的隐藏状态是所有层最后一个时间步的h和c。

实现代码

这里我写了一个可复用的MultiLayerLSTM类,完全用LSTMCell实现,并且和官方nn.LSTM的输入输出形状完全对齐:

import torch
import torch.nn as nn

class MultiLayerLSTM(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, bias=True, batch_first=False):
        super().__init__()
        self.num_layers = num_layers
        self.hidden_size = hidden_size
        self.batch_first = batch_first
        
        # 为每层创建独立的LSTMCell
        self.layers = nn.ModuleList()
        for layer_idx in range(num_layers):
            # 第一层输入维度是input_size,后续层是上一层的hidden_size
            current_input_size = input_size if layer_idx == 0 else hidden_size
            self.layers.append(nn.LSTMCell(current_input_size, hidden_size, bias=bias))
    
    def forward(self, x, hidden=None):
        # 处理batch_first的情况,统一转为(seq_len, batch_size, input_size)
        if self.batch_first:
            x = x.transpose(0, 1)
        seq_len, batch_size, _ = x.shape
        device = x.device
        
        # 初始化各层的h和c
        if hidden is None:
            h_states = [torch.zeros(batch_size, self.hidden_size, device=device) for _ in range(self.num_layers)]
            c_states = [torch.zeros(batch_size, self.hidden_size, device=device) for _ in range(self.num_layers)]
        else:
            # 传入的hidden是(num_layers, batch_size, hidden_size)的张量,拆分到每层
            h_all, c_all = hidden
            h_states = [h_all[layer_idx] for layer_idx in range(self.num_layers)]
            c_states = [c_all[layer_idx] for layer_idx in range(self.num_layers)]
        
        # 存储最后一层的所有时间步输出
        final_outputs = []
        
        # 遍历每个时间步处理
        for t in range(seq_len):
            current_input = x[t]  # 当前时间步输入:(batch_size, input_size)
            
            # 逐层传递输入
            for layer_idx in range(self.num_layers):
                # 当前层的LSTMCell计算
                h_states[layer_idx], c_states[layer_idx] = self.layers[layer_idx](
                    current_input, (h_states[layer_idx], c_states[layer_idx])
                )
                # 下一层的输入是当前层的h
                current_input = h_states[layer_idx]
            
            # 记录最后一层的输出
            final_outputs.append(h_states[-1])
        
        # 整理输出和隐藏状态的形状,和官方LSTM对齐
        final_outputs = torch.stack(final_outputs, dim=0)  # (seq_len, batch_size, hidden_size)
        if self.batch_first:
            final_outputs = final_outputs.transpose(0, 1)  # 转回batch_first格式
        
        # 把各层的h和c堆叠成(num_layers, batch_size, hidden_size)
        final_h = torch.stack(h_states, dim=0)
        final_c = torch.stack(c_states, dim=0)
        
        return final_outputs, (final_h, final_c)

测试与验证

为了确保我们的实现和官方nn.LSTM一致,可以做个简单的测试:

# 配置参数
input_size = 16
hidden_size = 32
num_layers = 2
batch_size = 4
seq_len = 8
batch_first = True

# 创建模型
custom_lstm = MultiLayerLSTM(input_size, hidden_size, num_layers, batch_first=batch_first)
official_lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=batch_first)

# 统一初始化权重(保证测试结果一致)
def init_weights(module):
    if isinstance(module, (nn.LSTMCell, nn.LSTM)):
        nn.init.orthogonal_(module.weight_ih_l0 if isinstance(module, nn.LSTM) else module.weight_ih)
        nn.init.orthogonal_(module.weight_hh_l0 if isinstance(module, nn.LSTM) else module.weight_hh)
        nn.init.zeros_(module.bias_ih_l0 if isinstance(module, nn.LSTM) else module.bias_ih)
        nn.init.zeros_(module.bias_hh_l0 if isinstance(module, nn.LSTM) else module.bias_hh)

custom_lstm.apply(init_weights)
official_lstm.apply(init_weights)

# 生成测试输入
x = torch.randn(batch_size, seq_len, input_size)

# 前向传播
custom_out, custom_hidden = custom_lstm(x)
official_out, official_hidden = official_lstm(x)

# 检查结果是否一致(浮点精度允许微小误差)
print("输出是否一致:", torch.allclose(custom_out, official_out, atol=1e-6))
print("隐藏状态h是否一致:", torch.allclose(custom_hidden[0], official_hidden[0], atol=1e-6))
print("细胞状态c是否一致:", torch.allclose(custom_hidden[1], official_hidden[1], atol=1e-6))

运行这段代码,你会看到三个True,说明我们的实现和官方LSTM完全等价。

关键注意点

  • 每层独立参数:必须为每层创建独立的LSTMCell实例,不能共享参数,否则就不是真正的多层LSTM了。
  • 状态传递逻辑:细胞状态c只在当前层的时间步之间传递,跨层传递的是隐藏状态h。
  • 形状对齐:如果需要和官方nn.LSTM无缝替换,一定要保证输入输出的形状完全一致,尤其是batch_first的处理。

内容的提问来源于stack exchange,提问作者Amuoeba

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 07:03:01