PyTorch torch.nn.LSTM Function

PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual


torch.nn.LSTMIt is a module in PyTorch for Long Short-Term Memory networks.

LSTM is a special type of recurrent neural network that can learn long-term dependencies and is widely used in sequence modeling tasks.

Function Definition

torch.nn.LSTM(input_size, hidden_size, num_layers=1, bias=True, batch_first=True, dropout=0, bidirectional=False)

Parameter Description:

  • input_size(int): Input feature dimension.
  • hidden_size(int): Hidden state dimension.
  • num_layers(int): Number of LSTM layers. Default is 1.
  • bias(bool): Whether to use bias. Default is True.
  • batch_first(bool): If True, the input and output shape is (batch, seq, feature). Default is True.
  • dropout(float): Dropout used for non-last layers. Default is 0.
  • bidirectional(bool): Whether to use bidirectional LSTM. Default is False.

Inputs and Outputs

Inputs:

  • input: Tensor of shape (batch, seq_len, input_size)
  • h_0: Initial hidden state, shape (num_layers * num_directions, batch, hidden_size)
  • c_0: Initial cell state, shape (num_layers * num_directions, batch, hidden_size)

Outputs:

  • output: The output of the last hidden layer, shape (batch, seq_len, num_directions * hidden_size)
  • h_n: The last hidden state of all layers
  • c_n: The last cell state of all layers

Usage Examples

Example 1: Basic Usage

Create and use LSTM:

Example

import torch
import torch.nn as nn

# Create LSTM: input dim 256, hidden dim 512, 2 layers
lstm = nn.LSTM(input_size=256, hidden_size=512, num_layers=2, batch_first=True)

# Create input: batch=4, sequence length=10, input dim=256
input_tensor = torch.randn(4, 10, 256)

# Forward pass
output, (h_n, c_n) = lstm(input_tensor)

print("Input shape:", input_tensor.shape)
print("Output shape:", output.shape)      # (4, 10, 512)
print("Hidden state shape:", h_n.shape)     # (2, 4, 512)
print("Cell state shape:", c_n.shape)     # (2, 4, 512)

Example 2: Bidirectional LSTM

Use bidirectional LSTM to capture bidirectional context:

Example

import torch
import torch.nn as nn

# Bidirectional LSTM
bilstm = nn.LSTM(input_size=256, hidden_size=256, num_layers=2, batch_first=True, bidirectional=True)

input_tensor = torch.randn(4, 10, 256)
output, (h_n, c_n) = bilstm(input_tensor)

print("Bidirectional LSTM output shape:", output.shape)   # (4, 10, 512) = 256*2
print("Hidden state shape:", h_n.shape)            # (4, 4, 256) = 2 layers * 2 directions
print("Last layer hidden state:", h_n[-2:, :, :].shape)  # Forward and backward

Example 3: Initializing Hidden State

Manually initialize hidden state:

Example

import torch
import torch.nn as nn

lstm = nn.LSTM(input_size=256, hidden_size=512, batch_first=True)

# Manually create initial hidden state
batch_size = 4
num_layers = 2
hidden_size = 512

h_0 = torch.zeros(num_layers, batch_size, hidden_size)
c_0 = torch.zeros(num_layers, batch_size, hidden_size)

# Pass in initial state
input_tensor = torch.randn(4, 10, 256)
output, (h_n, c_n) = lstm(input_tensor, (h_0, c_0))

print("Completed with custom initial state")
print("Output shape:", output.shape)

Example 4: Complete Sentiment Classification Model

Text classification based on LSTM:

Example

import torch
import torch.nn as nn

class LSTMClassifier(nn.Module):
    def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes=2):
        super(LSTMClassifier, self).__init__()
        # Embedding layer
        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
        # LSTM layer
        self.lstm = nn.LSTM(
            input_size=embed_dim,
            hidden_size=hidden_dim,
            num_layers=2,
            batch_first=True,
            bidirectional=True,
            dropout=0.3
        )
        # Fully connected classification layer
        self.fc = nn.Linear(hidden_dim * 2, num_classes)

    def forward(self, x):
        # x: (batch, seq_len)
        embedded = self.embedding(x)  # (batch, seq_len, embed_dim)

        # LSTM output
        output, (hidden, cell) = self.lstm(embedded)

        # Concatenate hidden states of the last bidirectional layer
        # hidden: (4, batch, hidden_dim) - 2 layers * 2 directions
        hidden = torch.cat([hidden[-2], hidden[-1]], dim=1)  # (batch, hidden_dim*2)

        # Classification
        logits = self.fc(hidden)
        return logits

# Instantiate the model
vocab_size = 10000
model = LSTMClassifier(vocab_size=vocab_size, embed_dim=128, hidden_dim=128, num_classes=2)

# Test input: batch=8, sequence length=50
input_ids = torch.randint(1, vocab_size, (8, 50))
output = model(input_ids)

print("Model structure:")
print(model)
print("nInput shape:", input_ids.shape)
print("Output shape:", output.shape)  # (8, 2)

Example 5: Multi-layer Stacked LSTM

Deep LSTM network:

Example

import torch
import torch.nn as nn

# 4-layer stacked LSTM with dropout
deep_lstm = nn.LSTM(
    input_size=256,
    hidden_size=512,
    num_layers=4,
    batch_first=True,
    dropout=0.4  # Dropout between layers
)

input_tensor = torch.randn(2, 20, 256)
output, (h_n, c_n) = deep_lstm(input_tensor)

print("4-layer LSTM output shape:", output.shape)
print("Hidden state shape (4 layers):", h_n.shape)
print("Cell state shape (4 layers):", c_n.shape)

Concept of LSTM Gates

LSTM controls the flow of information through three gates:

  • Forget gate: Determines how much information from the previous time step to retain
  • Input gate: Determines how much new information to add
  • Output gate: Determines how much information to output

FAQ

Q1: What does batch_first=True mean?

The first dimension of the input and output tensors is batch_size. If False, the first dimension is the sequence length.

Q2: When should bidirectional LSTM be used?

Tasks that require bidirectional context, such as sequence labeling and sentiment analysis. Machine translation often uses the encoder-decoder architecture.

Q3: How to choose the hidden layer size?

Usually 128-512, adjusted according to task complexity and data volume. Too small leads to underfitting, too large leads to overfitting.


Use Cases

nn.LSTMMain application scenarios include:

  • Natural language processing: Text classification, named entity recognition
  • Time series prediction: Stock prediction, speech recognition
  • Sequence-to-sequence tasks: Machine translation, text generation

Tip: When using bidirectional=True, the output dimension becomes hidden_size * 2.


PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual

Other Extensions