求PyTorch 0.4.0中nn.LayerNorm结合nn.LSTMCell的实现示例
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.LayerNormis 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

