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

求PyTorch 0.4.0中nn.LayerNorm结合nn.LSTMCell的实现示例

Using nn.LayerNorm with LSTM via LSTMCell in PyTorch 0.4.0

I get it—integrating LayerNorm into an LSTM can feel tricky when the built-in nn.LSTM doesn’t support it directly, especially in older PyTorch versions like 0.4.0. Since contributors mentioned you need to use nn.LSTMCell, here’s a complete, working example that implements an LSTM with layer normalization applied to cell and hidden states at every time step:

Custom LayerNorm-LSTM Implementation

This module replicates the behavior of the standard nn.LSTM but adds LayerNorm after each LSTMCell operation:

import torch
import torch.nn as nn

class LayerNormLSTM(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers=1, dropout=0.0):
        super(LayerNormLSTM, self).__init__()
        self.hidden_size = hidden_size
        self.num_layers = num_layers
        
        # Initialize LSTM cells and corresponding LayerNorm layers
        self.lstm_cells = nn.ModuleList()
        self.layer_norms = nn.ModuleList()
        
        for layer_idx in range(num_layers):
            # Input size varies for the first layer vs. subsequent layers
            current_input_size = input_size if layer_idx == 0 else hidden_size
            self.lstm_cells.append(nn.LSTMCell(current_input_size, hidden_size))
            # Add LayerNorm for both hidden and cell states
            self.layer_norms.append(nn.LayerNorm(hidden_size))
            self.layer_norms.append(nn.LayerNorm(hidden_size))
        
        self.dropout = nn.Dropout(dropout) if dropout > 0 else None

    def forward(self, x, hidden=None):
        # Input shape: (seq_len, batch_size, input_size)
        seq_len, batch_size, _ = x.size()
        
        # Initialize hidden/cell states if not provided
        if hidden is None:
            device = x.device
            h = [torch.zeros(batch_size, self.hidden_size, device=device) for _ in range(self.num_layers)]
            c = [torch.zeros(batch_size, self.hidden_size, device=device) for _ in range(self.num_layers)]
        else:
            h, c = hidden
            # Unpack stacked hidden states into a list for per-layer processing
            if isinstance(h, torch.Tensor):
                h = list(h.unbind(0))
            if isinstance(c, torch.Tensor):
                c = list(c.unbind(0))
        
        outputs = []
        for t in range(seq_len):
            current_x = x[t]
            for layer_idx in range(self.num_layers):
                # Forward pass through LSTM cell
                h[layer_idx], c[layer_idx] = self.lstm_cells[layer_idx](current_x, (h[layer_idx], c[layer_idx]))
                # Apply LayerNorm to hidden and cell states
                h[layer_idx] = self.layer_norms[2*layer_idx](h[layer_idx])
                c[layer_idx] = self.layer_norms[2*layer_idx + 1](c[layer_idx])
                # Apply dropout between layers (skip for final layer)
                if self.dropout is not None and layer_idx != self.num_layers - 1:
                    h[layer_idx] = self.dropout(h[layer_idx])
                current_x = h[layer_idx]
            outputs.append(current_x)
        
        # Stack outputs back to (seq_len, batch_size, hidden_size)
        outputs = torch.stack(outputs, dim=0)
        # Repack hidden/cell states to (num_layers, batch_size, hidden_size)
        h = torch.stack(h, dim=0)
        c = torch.stack(c, dim=0)
        return outputs, (h, c)

How to Use the Module

You can drop this into your existing code just like a standard nn.LSTM:

# Hyperparameters
input_size = 128
hidden_size = 256
num_layers = 2
seq_len = 10
batch_size = 8

# Initialize the model
model = LayerNormLSTM(input_size, hidden_size, num_layers, dropout=0.1)

# Dummy input tensor (seq_len, batch_size, input_size)
x = torch.randn(seq_len, batch_size, input_size)

# Forward pass
output, (final_hidden, final_cell) = model(x)

print(f"Output shape: {output.shape}")  # Expected: (10, 8, 256)
print(f"Final hidden state shape: {final_hidden.shape}")  # Expected: (2, 8, 256)

Key Details

  • We apply LayerNorm to both hidden (h) and cell (c) states to stabilize training, a common practice for normalized RNNs.
  • The module maintains the same input/output shape as nn.LSTM, so you won’t need to rewrite downstream code.
  • Dropout is only applied between layers (not after the final layer) to avoid losing critical output information.

Note for PyTorch 0.4.0: Double-check that nn.LayerNorm is available—this module was officially added in 0.4.0, so you shouldn’t run into import issues if you’re on the correct version.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:35:08