PyTorch torch.nn.LSTM Function
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 layersc_n: The last cell state of all layers
Usage Examples
Example 1: Basic Usage
Create and use LSTM:
Example
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.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.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.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.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.
Other Extensions